Merge pull request #272 from chopratejas/rust-stage-2.1-hf-hub

rust(stage 2.1): HfTokenizer::from_pretrained via hf-hub
This commit is contained in:
Tejas Chopra 2026-04-25 16:15:24 -07:00 committed by GitHub
commit c275567d94
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 269 additions and 7 deletions

162
Cargo.lock generated
View file

@ -2,6 +2,12 @@
# It is not intended for manual editing.
version = 3
[[package]]
name = "adler2"
version = "2.0.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
[[package]]
name = "ahash"
version = "0.8.12"
@ -414,6 +420,15 @@ dependencies = [
"libc",
]
[[package]]
name = "crc32fast"
version = "1.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511"
dependencies = [
"cfg-if",
]
[[package]]
name = "criterion"
version = "0.5.1"
@ -600,6 +615,27 @@ dependencies = [
"crypto-common",
]
[[package]]
name = "dirs"
version = "6.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c3e8aa94d75141228480295a7d0e7feb620b1a5ad9f12bc40be62411e38cce4e"
dependencies = [
"dirs-sys",
]
[[package]]
name = "dirs-sys"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e01a3366d27ee9890022452ee61b2b63a67e6f13f58900b651ff5665f0bb1fab"
dependencies = [
"libc",
"option-ext",
"redox_users",
"windows-sys 0.61.2",
]
[[package]]
name = "displaydoc"
version = "0.2.5"
@ -671,6 +707,16 @@ version = "0.1.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582"
[[package]]
name = "flate2"
version = "1.1.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "843fba2746e448b37e26a819579957415c8cef339bf08564fe8b7ddbd959573c"
dependencies = [
"crc32fast",
"miniz_oxide",
]
[[package]]
name = "fnv"
version = "1.0.7"
@ -881,6 +927,7 @@ version = "0.1.0"
dependencies = [
"bytes",
"criterion",
"hf-hub",
"proptest",
"serde",
"serde_json",
@ -956,6 +1003,26 @@ version = "0.5.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c"
[[package]]
name = "hf-hub"
version = "0.4.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "629d8f3bbeda9d148036d6b0de0a3ab947abd08ce90626327fc3547a49d59d97"
dependencies = [
"dirs",
"http",
"indicatif",
"libc",
"log",
"rand 0.9.4",
"reqwest",
"serde",
"serde_json",
"thiserror 2.0.18",
"ureq",
"windows-sys 0.60.2",
]
[[package]]
name = "http"
version = "1.4.0"
@ -1304,6 +1371,15 @@ version = "0.2.185"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "52ff2c0fe9bc6cb6b14a0592c2ff4fa9ceb83eea9db979b0487cd054946a2b8f"
[[package]]
name = "libredox"
version = "0.1.16"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e02f3bb43d335493c96bf3fd3a321600bf6bd07ed34bc64118e9293bdffea46c"
dependencies = [
"libc",
]
[[package]]
name = "linux-raw-sys"
version = "0.12.1"
@ -1386,6 +1462,16 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a"
[[package]]
name = "miniz_oxide"
version = "0.8.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1fa76a2c86f704bdb222d66965fb3d63269ce38518b83cb0575fca855ebb6316"
dependencies = [
"adler2",
"simd-adler32",
]
[[package]]
name = "mio"
version = "1.2.0"
@ -1503,6 +1589,12 @@ version = "11.1.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e"
[[package]]
name = "option-ext"
version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d"
[[package]]
name = "paste"
version = "1.0.15"
@ -1861,6 +1953,17 @@ dependencies = [
"crossbeam-utils",
]
[[package]]
name = "redox_users"
version = "0.5.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a4e608c6638b9c18977b00b475ac1f28d14e84b27d8d42f70e0bf1e3dec127ac"
dependencies = [
"getrandom 0.2.17",
"libredox",
"thiserror 2.0.18",
]
[[package]]
name = "regex"
version = "1.12.3"
@ -1977,6 +2080,7 @@ version = "0.23.39"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7c2c118cb077cca2822033836dfb1b975355dfb784b5e8da48f7b6c5db74e60e"
dependencies = [
"log",
"once_cell",
"ring",
"rustls-pki-types",
@ -2147,6 +2251,12 @@ dependencies = [
"libc",
]
[[package]]
name = "simd-adler32"
version = "0.3.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "703d5c7ef118737c72f1af64ad2f6f8c5e1921f818cdcb97b8fe6fc69bf66214"
[[package]]
name = "slab"
version = "0.4.12"
@ -2169,6 +2279,17 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "socks"
version = "0.3.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f0c3dbbd9ae980613c6dd8e28a9407b50509d3803b57624d5dfe8315218cd58b"
dependencies = [
"byteorder",
"libc",
"winapi",
]
[[package]]
name = "spm_precompiled"
version = "0.1.4"
@ -2677,6 +2798,25 @@ version = "0.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1"
[[package]]
name = "ureq"
version = "2.12.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "02d1a66277ed75f640d608235660df48c8e3c19f3b4edb6a263315626cc3c01d"
dependencies = [
"base64 0.22.1",
"flate2",
"log",
"once_cell",
"rustls",
"rustls-pki-types",
"serde",
"serde_json",
"socks",
"url",
"webpki-roots 0.26.11",
]
[[package]]
name = "url"
version = "2.5.8"
@ -2922,6 +3062,22 @@ dependencies = [
"rustls-pki-types",
]
[[package]]
name = "winapi"
version = "0.3.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419"
dependencies = [
"winapi-i686-pc-windows-gnu",
"winapi-x86_64-pc-windows-gnu",
]
[[package]]
name = "winapi-i686-pc-windows-gnu"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6"
[[package]]
name = "winapi-util"
version = "0.1.11"
@ -2931,6 +3087,12 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "winapi-x86_64-pc-windows-gnu"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f"
[[package]]
name = "windows-link"
version = "0.2.1"

View file

@ -18,6 +18,11 @@ tiktoken-rs = "0.11"
# pull in `onig` for the BPE pre-tokenizer regex; that vendors oniguruma so it
# builds without a system dep on macOS/Linux.
tokenizers = "0.21"
# `hf-hub` is the HuggingFace Hub client. We use the blocking `ureq` transport
# with `rustls` (no system OpenSSL dep — keeps the binary static-linkable for
# AWS deploys). `from_pretrained` is called once at startup, so blocking is
# fine; if a tokio caller needs it later we can wrap in `spawn_blocking`.
hf-hub = { version = "0.4", default-features = false, features = ["ureq", "rustls-tls"] }
[dev-dependencies]
proptest = "1"

View file

@ -8,14 +8,20 @@
//! `tokenizer.json` on the HuggingFace Hub and the `tokenizers` crate is a
//! pure-Rust loader, so we don't have to estimate.
//!
//! # Sources
//! - [`HfTokenizer::from_bytes`] — for tokenizers embedded via `include_bytes!`.
//! - [`HfTokenizer::from_file`] — for a local `tokenizer.json`.
//! - [`HfTokenizer::from_pretrained`] — pulls `tokenizer.json` from the
//! HuggingFace Hub via `hf-hub`. First call downloads to
//! `~/.cache/huggingface/hub`, subsequent calls hit the cache. Blocking;
//! call from `main()` or `tokio::task::spawn_blocking`. Gated repos
//! (Llama, Mistral) require an `HF_TOKEN` env var or a token in
//! `~/.cache/huggingface/token`.
//!
//! # What's NOT here
//! - **No HuggingFace Hub auto-download.** Callers pass bytes or a path. A
//! later stage can add `hf-hub` integration behind a Cargo feature; doing it
//! here would drag in ureq/rustls, a `~/.cache/huggingface` dependency, and
//! gated-model auth flows that don't belong in the core crate.
//! - **No tokenizer.json bundled in the binary.** Bundling Llama / Cohere
//! tokenizers would add several MB of binary bloat for code paths most users
//! don't hit.
//! don't hit. `from_pretrained` lazily downloads instead.
use std::path::Path;
use std::sync::Arc;
@ -35,6 +41,14 @@ pub enum HfTokenizerError {
#[source]
source: Box<dyn std::error::Error + Send + Sync>,
},
/// The HuggingFace Hub fetch failed: network error, 404 on the repo, or
/// 401 on a gated model without an `HF_TOKEN`.
#[error("failed to download `{repo}` from HuggingFace Hub: {source}")]
Hub {
repo: String,
#[source]
source: Box<dyn std::error::Error + Send + Sync>,
},
}
/// Token counter backed by a HuggingFace `tokenizer.json`.
@ -86,6 +100,37 @@ impl HfTokenizer {
})
}
/// Download (or fetch from cache) `tokenizer.json` for `repo` from the
/// HuggingFace Hub and load it. `repo` is the canonical Hub identifier,
/// e.g. `"CohereForAI/c4ai-command-r-v01"` or `"meta-llama/Meta-Llama-3-8B"`.
///
/// Uses the `main` revision. Blocking — calls into `ureq` synchronously.
/// First successful call writes the file to `~/.cache/huggingface/hub`
/// (or `$HF_HOME` if set); subsequent calls in the same or later
/// processes hit the on-disk cache.
///
/// Errors:
/// - [`HfTokenizerError::Hub`] for download failures (no network, 404,
/// 401 on a gated model without `HF_TOKEN`).
/// - [`HfTokenizerError::Load`] if the downloaded bytes don't parse as
/// a valid `tokenizer.json`. Should not happen for healthy HF repos.
pub fn from_pretrained(repo: &str) -> Result<Self, HfTokenizerError> {
let api = hf_hub::api::sync::Api::new().map_err(|e| HfTokenizerError::Hub {
repo: repo.to_string(),
source: Box::new(e),
})?;
let path = api
.model(repo.to_string())
.get("tokenizer.json")
.map_err(|e| HfTokenizerError::Hub {
repo: repo.to_string(),
source: Box::new(e),
})?;
// `get` returns the on-disk path; reuse `from_file` to keep the load
// path identical to user-supplied tokenizer.json files.
Self::from_file(repo, path)
}
/// The logical name this tokenizer was registered under (e.g.
/// `"command-r-plus"`). Used in logs and metrics.
pub fn name(&self) -> &str {
@ -237,4 +282,33 @@ mod tests {
let _ = std::fs::remove_dir_all(&dir);
}
/// Network-dependent: hits HuggingFace Hub. Run with
/// `cargo test -p headroom-core -- --ignored from_pretrained_downloads_real_tokenizer`.
/// `gpt2` is a small public unauthenticated repo (~1.4 MB tokenizer.json).
#[test]
#[ignore = "network-dependent: hits HuggingFace Hub"]
fn from_pretrained_downloads_real_tokenizer() {
let t = HfTokenizer::from_pretrained("gpt2").expect("download succeeds");
// GPT-2 BPE: "hello world" is 2 tokens. Locks in that we got a real
// BPE tokenizer (not a WhitespaceSplit fixture) from HF.
assert_eq!(t.count_text("hello world"), 2);
assert_eq!(t.name(), "gpt2");
assert_eq!(t.backend(), Backend::HuggingFace);
}
/// Negative path: a malformed repo name fails before any network call
/// (the Hub URL builder rejects it). No network required to run.
#[test]
fn from_pretrained_invalid_repo_returns_hub_error() {
// Empty repo name — hf-hub rejects this without making a network
// request. Locks in that we propagate Hub errors as `Hub`, not
// `Load`, so callers can distinguish "couldn't fetch" from
// "fetched but malformed".
let r = HfTokenizer::from_pretrained("");
assert!(
matches!(r, Err(HfTokenizerError::Hub { .. })),
"expected HfTokenizerError::Hub, got {r:?}"
);
}
}

View file

@ -36,7 +36,9 @@ mod tiktoken_impl;
pub use estimator::EstimatingCounter;
pub use hf_impl::{HfTokenizer, HfTokenizerError};
pub use registry::{clear_hf_registrations, detect_backend, get_tokenizer, register_hf, Backend};
pub use registry::{
clear_hf_registrations, detect_backend, get_tokenizer, register_hf, try_register_hf, Backend,
};
pub use tiktoken_impl::{TiktokenCounter, TiktokenError};
/// Counts tokens. Implementations must be thread-safe (`Send + Sync`).

View file

@ -17,7 +17,7 @@
use std::collections::HashMap;
use std::sync::{OnceLock, RwLock};
use super::{EstimatingCounter, HfTokenizer, TiktokenCounter, Tokenizer};
use super::{EstimatingCounter, HfTokenizer, HfTokenizerError, TiktokenCounter, Tokenizer};
/// Which family of tokenizer was selected for a model.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@ -124,6 +124,25 @@ pub fn clear_hf_registrations() {
hf_table().write().expect("hf registry poisoned").clear();
}
/// Convenience: download `tokenizer.json` for `repo` from the HuggingFace
/// Hub and register it under `prefix`. One-line glue around
/// [`HfTokenizer::from_pretrained`] + [`register_hf`].
///
/// Useful for proxy startup code that wants real tokenizers for the major
/// non-OpenAI families. Each call is independent — failure for one model
/// (e.g. a gated Llama repo without `HF_TOKEN`) does not affect others.
///
/// ```no_run
/// use headroom_core::tokenizer::try_register_hf;
/// let _ = try_register_hf("command-", "CohereForAI/c4ai-command-r-v01");
/// let _ = try_register_hf("mistral-", "mistralai/Mistral-7B-v0.1");
/// ```
pub fn try_register_hf(prefix: &str, repo: &str) -> Result<(), HfTokenizerError> {
let t = HfTokenizer::from_pretrained(repo)?;
register_hf(prefix, t);
Ok(())
}
fn lookup_hf(model: &str) -> Option<HfTokenizer> {
let m = model.to_ascii_lowercase();
let table = hf_table().read().expect("hf registry poisoned");