Skip to main content

tiktoken_wasm/
lib.rs

1//! WebAssembly bindings for the tiktoken BPE tokenizer.
2//!
3//! Provides browser-compatible wrappers around the core `tiktoken` crate,
4//! enabling high-performance token encoding, decoding, counting, and
5//! cost estimation directly in JavaScript/TypeScript applications.
6//!
7//! All encoding instances are cached globally via `OnceLock`, so repeated
8//! calls to `getEncoding()` with the same name return the same underlying data.
9
10use wasm_bindgen::prelude::{wasm_bindgen, JsError};
11
12/// WASM wrapper around a tiktoken encoding instance.
13///
14/// Created via [`get_encoding`] or [`encoding_for_model`].
15/// Call `.free()` when done to release WASM memory.
16#[wasm_bindgen]
17pub struct Encoding {
18    /// encoding name (e.g. "cl100k_base") — always a static string
19    name: &'static str,
20    /// reference to the globally cached CoreBpe instance
21    bpe: &'static tiktoken::CoreBpe,
22}
23
24#[wasm_bindgen]
25impl Encoding {
26    /// Encode text into token ids (returns `Uint32Array` in JS).
27    ///
28    /// Special tokens like `<|endoftext|>` are treated as ordinary text.
29    /// Use `encodeWithSpecialTokens()` to recognize them.
30    pub fn encode(&self, text: &str) -> Vec<u32> {
31        self.bpe.encode(text)
32    }
33
34    /// Encode text into token ids, recognizing special tokens.
35    ///
36    /// Special tokens (e.g. `<|endoftext|>`) are encoded as their designated ids
37    /// instead of being split into sub-word pieces.
38    #[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    /// Decode token ids back to a UTF-8 string.
44    ///
45    /// Uses lossy UTF-8 conversion — invalid byte sequences are replaced with U+FFFD.
46    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    /// Count tokens without building the full token id array.
52    ///
53    /// Faster than `encode(text).length` for cases where you only need the count.
54    pub fn count(&self, text: &str) -> usize {
55        self.bpe.count(text)
56    }
57
58    /// Count tokens, recognizing special tokens.
59    ///
60    /// Like `count()` but special tokens (e.g. `<|endoftext|>`) are counted
61    /// as single tokens instead of being split into sub-word pieces.
62    #[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    /// Get the number of regular (non-special) tokens in the vocabulary.
68    #[wasm_bindgen(js_name = vocabSize, getter)]
69    pub fn vocab_size(&self) -> usize {
70        self.bpe.vocab_size()
71    }
72
73    /// Get the number of special tokens in the vocabulary.
74    #[wasm_bindgen(js_name = numSpecialTokens, getter)]
75    pub fn num_special_tokens(&self) -> usize {
76        self.bpe.num_special_tokens()
77    }
78
79    /// Get the encoding name (e.g. `"cl100k_base"`).
80    #[wasm_bindgen(getter)]
81    pub fn name(&self) -> String {
82        self.name.to_string()
83    }
84}
85
86/// List all available encoding names.
87///
88/// Returns an array of strings: `["cl100k_base", "o200k_base", ...]`
89#[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/// Get an encoding by name.
98///
99/// Supported encodings:
100/// - `"cl100k_base"` — GPT-4, GPT-3.5-turbo
101/// - `"o200k_base"` — GPT-4o, GPT-4.1, o1, o3
102/// - `"o200k_harmony"` — gpt-oss (harmony chat format)
103/// - `"p50k_base"` — text-davinci-002/003
104/// - `"p50k_edit"` — text-davinci-edit
105/// - `"r50k_base"` — GPT-3 (davinci, curie, etc.)
106/// - `"llama3"` — Meta Llama 3/4
107/// - `"deepseek_v3"` — DeepSeek V3/R1
108/// - `"qwen2"` — Qwen 2/2.5/3
109/// - `"mistral_v3"` — Mistral/Codestral/Pixtral
110///
111/// Throws `Error` for unknown encoding names.
112#[wasm_bindgen(js_name = getEncoding)]
113pub fn get_encoding(name: &str) -> Result<Encoding, JsError> {
114    // look up the static name from tiktoken's canonical list (single source of truth)
115    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/// Get an encoding for a model name (e.g. `"gpt-4o"`, `"o3-mini"`, `"llama-4"`, `"deepseek-r1"`).
128///
129/// Supports models from OpenAI, Meta, DeepSeek, Qwen, and Mistral.
130/// Automatically resolves the model name to the correct encoding.
131/// Throws `Error` for unknown model names.
132#[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/// Map a model name to its encoding name without loading the encoding.
142///
143/// Returns the encoding name string (e.g. `"o200k_base"`) or `null` for unknown models.
144#[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/// Estimate cost in USD for a given model, input token count, and output token count.
150///
151/// Supports OpenAI, Anthropic Claude, Google Gemini, Meta Llama, DeepSeek, Qwen, and Mistral models.
152/// Throws `Error` for unknown model ids.
153#[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/// Get model pricing and metadata.
164///
165/// Returns a typed object with: `id`, `provider`, `inputPer1m`, `outputPer1m`,
166/// `cachedInputPer1m`, `contextWindow`, `maxOutput`.
167///
168/// Throws `Error` for unknown model ids.
169#[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/// List all supported models with pricing info.
177///
178/// Returns an array of `ModelInfo` objects.
179#[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/// List models filtered by provider name.
188///
189/// Provider names: `"OpenAI"`, `"Anthropic"`, `"Google"`, `"Meta"`, `"DeepSeek"`, `"Alibaba"`, `"Mistral"`.
190/// Returns an empty array for unknown providers.
191#[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/// Model pricing and metadata.
229#[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}