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)]
113pub fn get_encoding(name: &str) -> Result<Encoding, JsError> {
114 let static_name = tiktoken::list_encodings()
116 .iter()
117 .find(|&&n| n == name)
118 .ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
119 let bpe = tiktoken::get_encoding(name)
120 .ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
121 Ok(Encoding {
122 name: static_name,
123 bpe,
124 })
125}
126
127#[wasm_bindgen(js_name = encodingForModel)]
133pub fn encoding_for_model(model: &str) -> Result<Encoding, JsError> {
134 let name = tiktoken::model_to_encoding(model)
135 .ok_or_else(|| JsError::new(&format!("unknown model: {model}")))?;
136 let bpe = tiktoken::get_encoding(name)
137 .ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
138 Ok(Encoding { name, bpe })
139}
140
141#[wasm_bindgen(js_name = modelToEncoding)]
145pub fn model_to_encoding(model: &str) -> Option<String> {
146 tiktoken::model_to_encoding(model).map(|s| s.to_string())
147}
148
149#[wasm_bindgen(js_name = estimateCost)]
154pub fn estimate_cost(
155 model_id: &str,
156 input_tokens: u32,
157 output_tokens: u32,
158) -> Result<f64, JsError> {
159 tiktoken::pricing::estimate_cost(model_id, input_tokens as u64, output_tokens as u64)
160 .ok_or_else(|| JsError::new(&format!("unknown model: {model_id}")))
161}
162
163#[wasm_bindgen(js_name = getModelInfo)]
170pub fn get_model_info(model_id: &str) -> Result<ModelInfo, JsError> {
171 let model = tiktoken::pricing::get_model(model_id)
172 .ok_or_else(|| JsError::new(&format!("unknown model: {model_id}")))?;
173 Ok(convert_model(model))
174}
175
176#[wasm_bindgen(js_name = allModels)]
180pub fn all_models() -> Vec<ModelInfo> {
181 tiktoken::pricing::all_models()
182 .iter()
183 .map(convert_model)
184 .collect()
185}
186
187#[wasm_bindgen(js_name = modelsByProvider)]
192pub fn models_by_provider(provider: &str) -> Vec<ModelInfo> {
193 let Some(provider) = parse_provider(provider) else {
194 return Vec::new();
195 };
196
197 tiktoken::pricing::models_by_provider(provider)
198 .iter()
199 .map(|m| convert_model(m))
200 .collect()
201}
202
203fn convert_model(m: &tiktoken::pricing::Model) -> ModelInfo {
204 ModelInfo {
205 id: m.id,
206 provider: m.provider.to_string(),
207 input_per_1m: m.pricing.input_per_1m,
208 output_per_1m: m.pricing.output_per_1m,
209 cached_input_per_1m: m.pricing.cached_input_per_1m,
210 context_window: m.context_window,
211 max_output: m.max_output,
212 }
213}
214
215fn parse_provider(s: &str) -> Option<tiktoken::pricing::Provider> {
216 match s {
217 "OpenAI" => Some(tiktoken::pricing::Provider::OpenAI),
218 "Anthropic" => Some(tiktoken::pricing::Provider::Anthropic),
219 "Google" => Some(tiktoken::pricing::Provider::Google),
220 "Meta" => Some(tiktoken::pricing::Provider::Meta),
221 "DeepSeek" => Some(tiktoken::pricing::Provider::DeepSeek),
222 "Alibaba" => Some(tiktoken::pricing::Provider::Alibaba),
223 "Mistral" => Some(tiktoken::pricing::Provider::Mistral),
224 _ => None,
225 }
226}
227
228#[wasm_bindgen]
230#[derive(Clone)]
231pub struct ModelInfo {
232 id: &'static str,
233 provider: String,
234 input_per_1m: f64,
235 output_per_1m: f64,
236 cached_input_per_1m: Option<f64>,
237 context_window: u32,
238 max_output: u32,
239}
240
241#[wasm_bindgen]
242impl ModelInfo {
243 #[wasm_bindgen(getter)]
244 pub fn id(&self) -> String {
245 self.id.to_string()
246 }
247 #[wasm_bindgen(getter)]
248 pub fn provider(&self) -> String {
249 self.provider.clone()
250 }
251 #[wasm_bindgen(getter, js_name = inputPer1m)]
252 pub fn input_per_1m(&self) -> f64 {
253 self.input_per_1m
254 }
255 #[wasm_bindgen(getter, js_name = outputPer1m)]
256 pub fn output_per_1m(&self) -> f64 {
257 self.output_per_1m
258 }
259 #[wasm_bindgen(getter, js_name = cachedInputPer1m)]
260 pub fn cached_input_per_1m(&self) -> Option<f64> {
261 self.cached_input_per_1m
262 }
263 #[wasm_bindgen(getter, js_name = contextWindow)]
264 pub fn context_window(&self) -> u32 {
265 self.context_window
266 }
267 #[wasm_bindgen(getter, js_name = maxOutput)]
268 pub fn max_output(&self) -> u32 {
269 self.max_output
270 }
271}
272
273#[cfg(test)]
274mod tests {
275 use super::*;
276
277 #[test]
278 fn all_encodings_roundtrip() {
279 for &name in tiktoken::list_encodings() {
280 let enc = get_encoding(name).unwrap();
281 let text = "hello world 你好 🚀";
282 let tokens = enc.encode(text);
283 let decoded = enc.decode(&tokens);
284 assert_eq!(decoded, text, "roundtrip failed for {name}");
285 }
286 }
287
288 #[test]
289 fn encoding_for_known_models() {
290 let models = [
291 "gpt-4o", "gpt-4", "gpt-3.5-turbo", "llama-4", "deepseek-r1", "qwen3", "mistral-large",
292 ];
293 for model in models {
294 let enc = encoding_for_model(model);
295 assert!(enc.is_ok(), "encoding_for_model failed for {model}");
296 }
297 }
298
299 #[test]
300 fn list_encodings_count() {
301 let names = list_encodings();
302 assert_eq!(names.len(), 9);
303 }
304
305 #[test]
306 fn all_models_count() {
307 let models = all_models();
308 assert_eq!(models.len(), tiktoken::pricing::all_models().len());
309 }
310
311 #[test]
312 fn models_by_valid_provider() {
313 let openai = models_by_provider("OpenAI");
314 assert!(!openai.is_empty());
315 for m in &openai {
316 assert_eq!(m.provider, "OpenAI");
317 }
318 }
319
320 #[test]
321 fn models_by_invalid_provider() {
322 let unknown = models_by_provider("NonExistent");
323 assert!(unknown.is_empty());
324 }
325
326 #[test]
327 fn estimate_cost_known_model() {
328 let cost = estimate_cost("gpt-4o", 1000, 1000).unwrap();
329 assert!(cost > 0.0);
330 }
331
332 #[test]
333 fn estimate_cost_unknown_model() {
334 assert!(estimate_cost("fake-model", 1000, 1000).is_err());
335 }
336
337 #[test]
338 fn get_model_info_known() {
339 let info = get_model_info("gpt-4o").unwrap();
340 assert_eq!(info.id(), "gpt-4o");
341 assert_eq!(info.provider(), "OpenAI");
342 assert!(info.context_window() > 0);
343 }
344
345 #[test]
346 fn get_model_info_unknown() {
347 assert!(get_model_info("fake-model").is_err());
348 }
349
350 #[test]
351 fn unknown_encoding_error() {
352 assert!(get_encoding("nonexistent").is_err());
353 }
354
355 #[test]
356 fn unknown_model_encoding_error() {
357 assert!(encoding_for_model("nonexistent-model-xyz").is_err());
358 }
359
360 #[test]
361 fn model_to_encoding_known() {
362 let name = model_to_encoding("gpt-4o");
363 assert_eq!(name.as_deref(), Some("o200k_base"));
364 }
365
366 #[test]
367 fn model_to_encoding_unknown() {
368 assert!(model_to_encoding("fake-model").is_none());
369 }
370
371 #[test]
372 fn parse_provider_all_variants() {
373 assert!(parse_provider("OpenAI").is_some());
374 assert!(parse_provider("Anthropic").is_some());
375 assert!(parse_provider("Google").is_some());
376 assert!(parse_provider("Meta").is_some());
377 assert!(parse_provider("DeepSeek").is_some());
378 assert!(parse_provider("Alibaba").is_some());
379 assert!(parse_provider("Mistral").is_some());
380 assert!(parse_provider("Unknown").is_none());
381 }
382}