use wasm_bindgen::prelude::{wasm_bindgen, JsError};
#[wasm_bindgen]
pub struct Encoding {
name: &'static str,
bpe: &'static tiktoken::CoreBpe,
}
#[wasm_bindgen]
impl Encoding {
pub fn encode(&self, text: &str) -> Vec<u32> {
self.bpe.encode(text)
}
#[wasm_bindgen(js_name = encodeWithSpecialTokens)]
pub fn encode_with_special_tokens(&self, text: &str) -> Vec<u32> {
self.bpe.encode_with_special_tokens(text)
}
pub fn decode(&self, tokens: &[u32]) -> String {
let bytes = self.bpe.decode(tokens);
String::from_utf8_lossy(&bytes).into_owned()
}
pub fn count(&self, text: &str) -> usize {
self.bpe.count(text)
}
#[wasm_bindgen(js_name = countWithSpecialTokens)]
pub fn count_with_special_tokens(&self, text: &str) -> usize {
self.bpe.count_with_special_tokens(text)
}
#[wasm_bindgen(js_name = vocabSize, getter)]
pub fn vocab_size(&self) -> usize {
self.bpe.vocab_size()
}
#[wasm_bindgen(js_name = numSpecialTokens, getter)]
pub fn num_special_tokens(&self) -> usize {
self.bpe.num_special_tokens()
}
#[wasm_bindgen(getter)]
pub fn name(&self) -> String {
self.name.to_string()
}
}
#[wasm_bindgen(js_name = listEncodings)]
pub fn list_encodings() -> Vec<String> {
tiktoken::list_encodings()
.iter()
.map(|s| s.to_string())
.collect()
}
#[wasm_bindgen(js_name = getEncoding)]
pub fn get_encoding(name: &str) -> Result<Encoding, JsError> {
let static_name = tiktoken::list_encodings()
.iter()
.find(|&&n| n == name)
.ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
let bpe = tiktoken::get_encoding(name)
.ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
Ok(Encoding {
name: static_name,
bpe,
})
}
#[wasm_bindgen(js_name = encodingForModel)]
pub fn encoding_for_model(model: &str) -> Result<Encoding, JsError> {
let name = tiktoken::model_to_encoding(model)
.ok_or_else(|| JsError::new(&format!("unknown model: {model}")))?;
let bpe = tiktoken::get_encoding(name)
.ok_or_else(|| JsError::new(&format!("unknown encoding: {name}")))?;
Ok(Encoding { name, bpe })
}
#[wasm_bindgen(js_name = modelToEncoding)]
pub fn model_to_encoding(model: &str) -> Option<String> {
tiktoken::model_to_encoding(model).map(|s| s.to_string())
}
#[wasm_bindgen(js_name = estimateCost)]
pub fn estimate_cost(
model_id: &str,
input_tokens: u32,
output_tokens: u32,
) -> Result<f64, JsError> {
tiktoken::pricing::estimate_cost(model_id, input_tokens as u64, output_tokens as u64)
.ok_or_else(|| JsError::new(&format!("unknown model: {model_id}")))
}
#[wasm_bindgen(js_name = getModelInfo)]
pub fn get_model_info(model_id: &str) -> Result<ModelInfo, JsError> {
let model = tiktoken::pricing::get_model(model_id)
.ok_or_else(|| JsError::new(&format!("unknown model: {model_id}")))?;
Ok(convert_model(model))
}
#[wasm_bindgen(js_name = allModels)]
pub fn all_models() -> Vec<ModelInfo> {
tiktoken::pricing::all_models()
.iter()
.map(convert_model)
.collect()
}
#[wasm_bindgen(js_name = modelsByProvider)]
pub fn models_by_provider(provider: &str) -> Vec<ModelInfo> {
let Some(provider) = parse_provider(provider) else {
return Vec::new();
};
tiktoken::pricing::models_by_provider(provider)
.iter()
.map(|m| convert_model(m))
.collect()
}
fn convert_model(m: &tiktoken::pricing::Model) -> ModelInfo {
ModelInfo {
id: m.id,
provider: m.provider.to_string(),
input_per_1m: m.pricing.input_per_1m,
output_per_1m: m.pricing.output_per_1m,
cached_input_per_1m: m.pricing.cached_input_per_1m,
context_window: m.context_window,
max_output: m.max_output,
}
}
fn parse_provider(s: &str) -> Option<tiktoken::pricing::Provider> {
match s {
"OpenAI" => Some(tiktoken::pricing::Provider::OpenAI),
"Anthropic" => Some(tiktoken::pricing::Provider::Anthropic),
"Google" => Some(tiktoken::pricing::Provider::Google),
"Meta" => Some(tiktoken::pricing::Provider::Meta),
"DeepSeek" => Some(tiktoken::pricing::Provider::DeepSeek),
"Alibaba" => Some(tiktoken::pricing::Provider::Alibaba),
"Mistral" => Some(tiktoken::pricing::Provider::Mistral),
_ => None,
}
}
#[wasm_bindgen]
#[derive(Clone)]
pub struct ModelInfo {
id: &'static str,
provider: String,
input_per_1m: f64,
output_per_1m: f64,
cached_input_per_1m: Option<f64>,
context_window: u32,
max_output: u32,
}
#[wasm_bindgen]
impl ModelInfo {
#[wasm_bindgen(getter)]
pub fn id(&self) -> String {
self.id.to_string()
}
#[wasm_bindgen(getter)]
pub fn provider(&self) -> String {
self.provider.clone()
}
#[wasm_bindgen(getter, js_name = inputPer1m)]
pub fn input_per_1m(&self) -> f64 {
self.input_per_1m
}
#[wasm_bindgen(getter, js_name = outputPer1m)]
pub fn output_per_1m(&self) -> f64 {
self.output_per_1m
}
#[wasm_bindgen(getter, js_name = cachedInputPer1m)]
pub fn cached_input_per_1m(&self) -> Option<f64> {
self.cached_input_per_1m
}
#[wasm_bindgen(getter, js_name = contextWindow)]
pub fn context_window(&self) -> u32 {
self.context_window
}
#[wasm_bindgen(getter, js_name = maxOutput)]
pub fn max_output(&self) -> u32 {
self.max_output
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn all_encodings_roundtrip() {
for &name in tiktoken::list_encodings() {
let enc = get_encoding(name).unwrap();
let text = "hello world 你好 🚀";
let tokens = enc.encode(text);
let decoded = enc.decode(&tokens);
assert_eq!(decoded, text, "roundtrip failed for {name}");
}
}
#[test]
fn encoding_for_known_models() {
let models = [
"gpt-4o", "gpt-4", "gpt-3.5-turbo", "llama-4", "deepseek-r1", "qwen3", "mistral-large",
];
for model in models {
let enc = encoding_for_model(model);
assert!(enc.is_ok(), "encoding_for_model failed for {model}");
}
}
#[test]
fn list_encodings_count() {
let names = list_encodings();
assert_eq!(names.len(), 9);
}
#[test]
fn all_models_count() {
let models = all_models();
assert_eq!(models.len(), tiktoken::pricing::all_models().len());
}
#[test]
fn models_by_valid_provider() {
let openai = models_by_provider("OpenAI");
assert!(!openai.is_empty());
for m in &openai {
assert_eq!(m.provider, "OpenAI");
}
}
#[test]
fn models_by_invalid_provider() {
let unknown = models_by_provider("NonExistent");
assert!(unknown.is_empty());
}
#[test]
fn estimate_cost_known_model() {
let cost = estimate_cost("gpt-4o", 1000, 1000).unwrap();
assert!(cost > 0.0);
}
#[test]
fn estimate_cost_unknown_model() {
assert!(estimate_cost("fake-model", 1000, 1000).is_err());
}
#[test]
fn get_model_info_known() {
let info = get_model_info("gpt-4o").unwrap();
assert_eq!(info.id(), "gpt-4o");
assert_eq!(info.provider(), "OpenAI");
assert!(info.context_window() > 0);
}
#[test]
fn get_model_info_unknown() {
assert!(get_model_info("fake-model").is_err());
}
#[test]
fn unknown_encoding_error() {
assert!(get_encoding("nonexistent").is_err());
}
#[test]
fn unknown_model_encoding_error() {
assert!(encoding_for_model("nonexistent-model-xyz").is_err());
}
#[test]
fn model_to_encoding_known() {
let name = model_to_encoding("gpt-4o");
assert_eq!(name.as_deref(), Some("o200k_base"));
}
#[test]
fn model_to_encoding_unknown() {
assert!(model_to_encoding("fake-model").is_none());
}
#[test]
fn parse_provider_all_variants() {
assert!(parse_provider("OpenAI").is_some());
assert!(parse_provider("Anthropic").is_some());
assert!(parse_provider("Google").is_some());
assert!(parse_provider("Meta").is_some());
assert!(parse_provider("DeepSeek").is_some());
assert!(parse_provider("Alibaba").is_some());
assert!(parse_provider("Mistral").is_some());
assert!(parse_provider("Unknown").is_none());
}
}