1use wasm_bindgen::prelude::{wasm_bindgen, JsError};
11
12#[wasm_bindgen]
17pub struct Encoding {
18 name: &'static str,
20 bpe: &'static tiktoken::CoreBpe,
22}
23
24#[wasm_bindgen]
25impl Encoding {
26 pub fn encode(&self, text: &str) -> Vec<u32> {
31 self.bpe.encode(text)
32 }
33
34 #[wasm_bindgen(js_name = encodeWithSpecialTokens)]
39 pub fn encode_with_special_tokens(&self, text: &str) -> Vec<u32> {
40 self.bpe.encode_with_special_tokens(text)
41 }
42
43 pub fn decode(&self, tokens: &[u32]) -> String {
47 let bytes = self.bpe.decode(tokens);
48 String::from_utf8_lossy(&bytes).into_owned()
49 }
50
51 pub fn count(&self, text: &str) -> usize {
55 self.bpe.count(text)
56 }
57
58 #[wasm_bindgen(js_name = countWithSpecialTokens)]
63 pub fn count_with_special_tokens(&self, text: &str) -> usize {
64 self.bpe.count_with_special_tokens(text)
65 }
66
67 #[wasm_bindgen(js_name = vocabSize, getter)]
69 pub fn vocab_size(&self) -> usize {
70 self.bpe.vocab_size()
71 }
72
73 #[wasm_bindgen(js_name = numSpecialTokens, getter)]
75 pub fn num_special_tokens(&self) -> usize {
76 self.bpe.num_special_tokens()
77 }
78
79 #[wasm_bindgen(getter)]
81 pub fn name(&self) -> String {
82 self.name.to_string()
83 }
84}
85
86#[wasm_bindgen(js_name = listEncodings)]
90pub fn list_encodings() -> Vec<String> {
91 tiktoken::list_encodings()
92 .iter()
93 .map(|s| s.to_string())
94 .collect()
95}
96
97#[wasm_bindgen(js_name = getEncoding)]
114pub fn get_encoding(name: &str) -> Result<Encoding, JsError> {
115 let static_name = tiktoken::list_encodings()
117 .iter()
118 .find(|&&n| n == name)
119 .ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
120 let bpe = tiktoken::get_encoding(name)
121 .ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
122 Ok(Encoding {
123 name: static_name,
124 bpe,
125 })
126}
127
128#[wasm_bindgen(js_name = encodingForModel)]
134pub fn encoding_for_model(model: &str) -> Result<Encoding, JsError> {
135 let name = tiktoken::model_to_encoding(model)
136 .ok_or_else(|| JsError::new(&format!("unknown model: {model}")))?;
137 let bpe = tiktoken::get_encoding(name)
138 .ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
139 Ok(Encoding { name, bpe })
140}
141
142#[wasm_bindgen(js_name = modelToEncoding)]
146pub fn model_to_encoding(model: &str) -> Option<String> {
147 tiktoken::model_to_encoding(model).map(|s| s.to_string())
148}
149
150#[wasm_bindgen(js_name = estimateCost)]
155pub fn estimate_cost(
156 model_id: &str,
157 input_tokens: u32,
158 output_tokens: u32,
159) -> Result<f64, JsError> {
160 tiktoken::pricing::estimate_cost(model_id, input_tokens as u64, output_tokens as u64)
161 .ok_or_else(|| JsError::new(&format!("unknown model: {model_id}")))
162}
163
164#[wasm_bindgen(js_name = getModelInfo)]
171pub fn get_model_info(model_id: &str) -> Result<ModelInfo, JsError> {
172 let model = tiktoken::pricing::get_model(model_id)
173 .ok_or_else(|| JsError::new(&format!("unknown model: {model_id}")))?;
174 Ok(convert_model(model))
175}
176
177#[wasm_bindgen(js_name = allModels)]
181pub fn all_models() -> Vec<ModelInfo> {
182 tiktoken::pricing::all_models()
183 .iter()
184 .map(convert_model)
185 .collect()
186}
187
188#[wasm_bindgen(js_name = modelsByProvider)]
193pub fn models_by_provider(provider: &str) -> Vec<ModelInfo> {
194 let Some(provider) = parse_provider(provider) else {
195 return Vec::new();
196 };
197
198 tiktoken::pricing::models_by_provider(provider)
199 .iter()
200 .map(|m| convert_model(m))
201 .collect()
202}
203
204fn convert_model(m: &tiktoken::pricing::Model) -> ModelInfo {
205 ModelInfo {
206 id: m.id,
207 provider: m.provider.to_string(),
208 input_per_1m: m.pricing.input_per_1m,
209 output_per_1m: m.pricing.output_per_1m,
210 cached_input_per_1m: m.pricing.cached_input_per_1m,
211 context_window: m.context_window,
212 max_output: m.max_output,
213 }
214}
215
216fn parse_provider(s: &str) -> Option<tiktoken::pricing::Provider> {
217 match s {
218 "OpenAI" => Some(tiktoken::pricing::Provider::OpenAI),
219 "Anthropic" => Some(tiktoken::pricing::Provider::Anthropic),
220 "Google" => Some(tiktoken::pricing::Provider::Google),
221 "Meta" => Some(tiktoken::pricing::Provider::Meta),
222 "DeepSeek" => Some(tiktoken::pricing::Provider::DeepSeek),
223 "Alibaba" => Some(tiktoken::pricing::Provider::Alibaba),
224 "Mistral" => Some(tiktoken::pricing::Provider::Mistral),
225 "Moonshot" => Some(tiktoken::pricing::Provider::Moonshot),
226 "Zhipu" => Some(tiktoken::pricing::Provider::Zhipu),
227 "MiniMax" => Some(tiktoken::pricing::Provider::MiniMax),
228 _ => None,
229 }
230}
231
232#[wasm_bindgen]
234#[derive(Clone)]
235pub struct ModelInfo {
236 id: &'static str,
237 provider: String,
238 input_per_1m: f64,
239 output_per_1m: f64,
240 cached_input_per_1m: Option<f64>,
241 context_window: u32,
242 max_output: u32,
243}
244
245#[wasm_bindgen]
246impl ModelInfo {
247 #[wasm_bindgen(getter)]
248 pub fn id(&self) -> String {
249 self.id.to_string()
250 }
251 #[wasm_bindgen(getter)]
252 pub fn provider(&self) -> String {
253 self.provider.clone()
254 }
255 #[wasm_bindgen(getter, js_name = inputPer1m)]
256 pub fn input_per_1m(&self) -> f64 {
257 self.input_per_1m
258 }
259 #[wasm_bindgen(getter, js_name = outputPer1m)]
260 pub fn output_per_1m(&self) -> f64 {
261 self.output_per_1m
262 }
263 #[wasm_bindgen(getter, js_name = cachedInputPer1m)]
264 pub fn cached_input_per_1m(&self) -> Option<f64> {
265 self.cached_input_per_1m
266 }
267 #[wasm_bindgen(getter, js_name = contextWindow)]
268 pub fn context_window(&self) -> u32 {
269 self.context_window
270 }
271 #[wasm_bindgen(getter, js_name = maxOutput)]
272 pub fn max_output(&self) -> u32 {
273 self.max_output
274 }
275}
276
277#[cfg(test)]
278mod tests {
279 use super::*;
280
281 #[test]
282 fn all_encodings_roundtrip() {
283 for &name in tiktoken::list_encodings() {
284 let enc = get_encoding(name).unwrap();
285 let text = "hello world 你好 🚀";
286 let tokens = enc.encode(text);
287 let decoded = enc.decode(&tokens);
288 assert_eq!(decoded, text, "roundtrip failed for {name}");
289 }
290 }
291
292 #[test]
293 fn encoding_for_known_models() {
294 let models = [
295 "gpt-4o", "gpt-4", "gpt-3.5-turbo", "llama-4", "deepseek-r1", "qwen3", "mistral-large",
296 ];
297 for model in models {
298 let enc = encoding_for_model(model);
299 assert!(enc.is_ok(), "encoding_for_model failed for {model}");
300 }
301 }
302
303 #[test]
304 fn list_encodings_count() {
305 let names = list_encodings();
308 assert_eq!(names.len(), tiktoken::list_encodings().len());
309 }
310
311 #[test]
312 fn all_models_count() {
313 let models = all_models();
314 assert_eq!(models.len(), tiktoken::pricing::all_models().len());
315 }
316
317 #[test]
318 fn models_by_valid_provider() {
319 let openai = models_by_provider("OpenAI");
320 assert!(!openai.is_empty());
321 for m in &openai {
322 assert_eq!(m.provider, "OpenAI");
323 }
324 }
325
326 #[test]
327 fn models_by_invalid_provider() {
328 let unknown = models_by_provider("NonExistent");
329 assert!(unknown.is_empty());
330 }
331
332 #[test]
333 fn estimate_cost_known_model() {
334 let cost = estimate_cost("gpt-4o", 1000, 1000).unwrap();
335 assert!(cost > 0.0);
336 }
337
338 #[test]
339 #[cfg(target_arch = "wasm32")] fn estimate_cost_unknown_model() {
341 assert!(estimate_cost("fake-model", 1000, 1000).is_err());
342 }
343
344 #[test]
345 fn get_model_info_known() {
346 let info = get_model_info("gpt-4o").unwrap();
347 assert_eq!(info.id(), "gpt-4o");
348 assert_eq!(info.provider(), "OpenAI");
349 assert!(info.context_window() > 0);
350 }
351
352 #[test]
353 #[cfg(target_arch = "wasm32")] fn get_model_info_unknown() {
355 assert!(get_model_info("fake-model").is_err());
356 }
357
358 #[test]
359 #[cfg(target_arch = "wasm32")] fn unknown_encoding_error() {
361 assert!(get_encoding("nonexistent").is_err());
362 }
363
364 #[test]
365 #[cfg(target_arch = "wasm32")] fn unknown_model_encoding_error() {
367 assert!(encoding_for_model("nonexistent-model-xyz").is_err());
368 }
369
370 #[test]
371 fn model_to_encoding_known() {
372 let name = model_to_encoding("gpt-4o");
373 assert_eq!(name.as_deref(), Some("o200k_base"));
374 }
375
376 #[test]
377 fn model_to_encoding_unknown() {
378 assert!(model_to_encoding("fake-model").is_none());
379 }
380
381 #[test]
382 fn parse_provider_all_variants() {
383 assert!(parse_provider("OpenAI").is_some());
384 assert!(parse_provider("Anthropic").is_some());
385 assert!(parse_provider("Google").is_some());
386 assert!(parse_provider("Meta").is_some());
387 assert!(parse_provider("DeepSeek").is_some());
388 assert!(parse_provider("Alibaba").is_some());
389 assert!(parse_provider("Mistral").is_some());
390 assert!(parse_provider("Moonshot").is_some());
391 assert!(parse_provider("Zhipu").is_some());
392 assert!(parse_provider("MiniMax").is_some());
393 assert!(parse_provider("Unknown").is_none());
394 }
395}