1use super::{ModelEvent, Source};
3use crate::model::{ByteSource, DataSpan, ModelDescription, ObjectId, ObjectInfo};
4use anyhow::{Context, Result, ensure};
5use cherenkov_model_data::{ContainerFormat, Inventory, read_safetensors};
6use reqwest::{
7 StatusCode,
8 blocking::Client,
9 header::{CONTENT_RANGE, RANGE},
10};
11use serde::Deserialize;
12use serde_json::{Value, json};
13use std::{
14 collections::BTreeSet,
15 io::{Read, Write},
16 time::Duration,
17};
18
19mod headers;
20
21const METADATA_LIMIT: u64 = 64 * 1024 * 1024;
22
23#[derive(Deserialize)]
24struct Repo {
25 sha: String,
26 siblings: Vec<Sibling>,
27}
28#[derive(Deserialize)]
29struct Sibling {
30 rfilename: String,
31}
32
33pub(super) fn endpoint() -> String {
34 std::env::var("HF_ENDPOINT")
35 .unwrap_or_else(|_| "https://huggingface.co".into())
36 .trim_end_matches('/')
37 .into()
38}
39
40pub(super) fn inspect(
41 repo: &str,
42 revision: &str,
43 token: Option<&str>,
44 events: Option<&mut dyn FnMut(ModelEvent)>,
45) -> Result<(Source, ModelDescription)> {
46 let token = token
47 .map(|value| Ok(Some(value.to_owned())))
48 .unwrap_or_else(ambient_token)?;
49 let hub = Hub {
50 client: Client::builder()
51 .timeout(Duration::from_secs(120))
52 .build()?,
53 endpoint: endpoint(),
54 token,
55 };
56
57 hub.inspect(repo, revision, events.unwrap_or(&mut |_| {}))
58}
59
60struct Hub {
61 client: Client,
62 endpoint: String,
63 token: Option<String>,
64}
65
66impl Hub {
67 fn get(&self, url: reqwest::Url) -> reqwest::blocking::RequestBuilder {
68 let request = self.client.get(url);
69
70 match &self.token {
71 Some(token) => request.bearer_auth(token),
72 None => request,
73 }
74 }
75
76 fn url(&self, components: &[&str]) -> Result<reqwest::Url> {
77 let mut url = reqwest::Url::parse(&self.endpoint)?;
78
79 ensure!(
80 matches!(url.scheme(), "http" | "https")
81 && url.username().is_empty()
82 && url.password().is_none(),
83 "HF endpoint must be an HTTP(S) URL without credentials"
84 );
85 url.path_segments_mut()
86 .map_err(|_| anyhow::anyhow!("invalid HF endpoint"))?
87 .pop_if_empty()
88 .extend(components);
89
90 Ok(url)
91 }
92
93 fn json(&self, components: &[&str]) -> Result<Value> {
94 let response = self.get(self.url(components)?).send()?.error_for_status()?;
95 let mut bytes = Vec::new();
96
97 response.take(METADATA_LIMIT + 1).read_to_end(&mut bytes)?;
98 ensure!(
99 bytes.len() as u64 <= METADATA_LIMIT,
100 "HF metadata exceeds 64 MiB"
101 );
102
103 Ok(serde_json::from_slice(&bytes)?)
104 }
105
106 fn inspect(
107 &self,
108 repo: &str,
109 revision: &str,
110 events: &mut dyn FnMut(ModelEvent),
111 ) -> Result<(Source, ModelDescription)> {
112 events(ModelEvent::Resolving {
113 source: format!("hf://{repo}@{revision}"),
114 });
115
116 let (owner, name) = repo
117 .split_once('/')
118 .context("HF source must be owner/repo")?;
119 let info: Repo = serde_json::from_value(
120 self.json(&["api", "models", owner, name, "revision", revision])?,
121 )?;
122
123 ensure!(
124 info.sha.len() == 40 && info.sha.bytes().all(|b| b.is_ascii_hexdigit()),
125 "HF did not return a full commit"
126 );
127
128 let config = self.json(&[owner, name, "resolve", &info.sha, "config.json"])?;
129 let indexed = info
130 .siblings
131 .iter()
132 .any(|s| s.rfilename == "model.safetensors.index.json");
133 let index = if indexed {
134 self.json(&[
135 owner,
136 name,
137 "resolve",
138 &info.sha,
139 "model.safetensors.index.json",
140 ])?
141 } else {
142 Value::Null
143 };
144 let files: BTreeSet<String> = if indexed {
145 index["weight_map"]
146 .as_object()
147 .context("missing safetensors weight_map")?
148 .values()
149 .map(|v| {
150 v.as_str()
151 .context("invalid shard filename")
152 .map(str::to_owned)
153 })
154 .collect::<Result<_>>()?
155 } else {
156 info.siblings
157 .into_iter()
158 .filter_map(|s| s.rfilename.ends_with(".safetensors").then_some(s.rfilename))
159 .collect()
160 };
161
162 ensure!(
163 !files.is_empty(),
164 "HF registration requires a safetensors checkpoint"
165 );
166
167 let files: Vec<_> = files.into_iter().collect();
168 let shards = self.headers(&[owner, name, "resolve", &info.sha], &files, events)?;
169 let mut objects = Vec::with_capacity(shards.len());
170 let mut tensors = Vec::new();
171 let mut names = BTreeSet::new();
172 let mut metadata = Vec::new();
173
174 for (file, (object, shard)) in files.iter().zip(shards) {
175 for tensor in &shard.tensors {
176 ensure!(
177 names.insert(tensor.name.clone()),
178 "duplicate tensor {}",
179 tensor.name
180 );
181
182 if indexed {
183 ensure!(
184 index["weight_map"][&tensor.name].as_str() == Some(file),
185 "shard index mismatch"
186 );
187 }
188 }
189
190 objects.push(object);
191 metadata.push(shard.metadata);
192 tensors.extend(shard.tensors);
193 }
194
195 if indexed {
196 ensure!(
197 names.len() == index["weight_map"].as_object().unwrap().len(),
198 "incomplete shard index"
199 );
200 }
201
202 let inventory = Inventory {
203 format: ContainerFormat::Safetensors,
204 metadata: json!({"config": config, "index": index, "shards": metadata}),
205 tensors,
206 };
207 let reader = RemoteReader {
208 hub: self,
209 objects: &objects,
210 };
211
212 events(ModelEvent::ReadingMetadata);
213
214 let description = crate::model::inspect::describe_inventory(&inventory, &reader, &config)?;
215
216 Ok((
217 Source::HuggingFace {
218 repo: repo.into(),
219 revision: info.sha,
220 endpoint: self.endpoint.clone(),
221 },
222 description,
223 ))
224 }
225
226 fn range(&self, mut url: reqwest::Url, offset: u64, length: u64) -> Result<(Vec<u8>, u64)> {
227 ensure!(
228 length > 0 && length <= METADATA_LIMIT,
229 "invalid metadata range length"
230 );
231
232 let end = offset.checked_add(length - 1).context("range overflow")?;
233
234 url.query_pairs_mut()
235 .append_pair("header_range", &format!("{offset}-{end}"));
236
237 let response = self
238 .get(url)
239 .header(RANGE, format!("bytes={offset}-{end}"))
240 .send()?
241 .error_for_status()?;
242
243 ensure!(
244 response.status() == StatusCode::PARTIAL_CONTENT,
245 "HF ignored the byte range; refusing a full shard read"
246 );
247
248 let range = response
249 .headers()
250 .get(CONTENT_RANGE)
251 .context("missing content range")?
252 .to_str()?;
253 let (bounds, total) = range.split_once('/').context("invalid content range")?;
254
255 ensure!(
256 bounds == format!("bytes {offset}-{end}"),
257 "incorrect content range"
258 );
259
260 let total = total.parse()?;
261 let mut bytes = Vec::new();
262
263 response.take(length + 1).read_to_end(&mut bytes)?;
264 ensure!(bytes.len() as u64 == length, "incomplete metadata range");
265
266 Ok((bytes, total))
267 }
268}
269
270struct RemoteObject {
271 id: ObjectId,
272 url: reqwest::Url,
273 prefix: Vec<u8>,
274 size: u64,
275}
276struct RemoteReader<'a> {
277 hub: &'a Hub,
278 objects: &'a [RemoteObject],
279}
280
281impl ByteSource for RemoteReader<'_> {
282 fn objects(&self) -> Vec<ObjectInfo> {
283 self.objects
284 .iter()
285 .map(|o| ObjectInfo {
286 id: o.id,
287 bytes: o.size,
288 })
289 .collect()
290 }
291 fn read(&self, span: &DataSpan, out: &mut dyn Write) -> Result<()> {
292 let object = self
293 .objects
294 .iter()
295 .find(|object| object.id == span.object)
296 .context("unknown HF shard")?;
297
298 ensure!(
299 span.offset
300 .checked_add(span.length)
301 .is_some_and(|end| end <= object.size),
302 "range exceeds HF shard"
303 );
304
305 if span.offset == 0 && span.length == 8 {
306 out.write_all(&object.prefix)?;
307
308 return Ok(());
309 }
310
311 let (bytes, size) = self
312 .hub
313 .range(object.url.clone(), span.offset, span.length)?;
314
315 ensure!(size == object.size, "HF shard changed during inspection");
316 out.write_all(&bytes)?;
317
318 Ok(())
319 }
320}
321
322fn ambient_token() -> Result<Option<String>> {
323 if std::env::var("HF_HUB_DISABLE_IMPLICIT_TOKEN")
324 .is_ok_and(|v| matches!(v.to_uppercase().as_str(), "1" | "ON" | "YES" | "TRUE"))
325 {
326 return Ok(None);
327 }
328
329 if let Ok(token) = std::env::var("HF_TOKEN") {
330 return Ok(Some(token));
331 }
332
333 let home = std::env::var_os("HF_HOME")
334 .map(std::path::PathBuf::from)
335 .or_else(|| dirs::home_dir().map(|h| h.join(".cache/huggingface")));
336 let path = std::env::var_os("HF_TOKEN_PATH")
337 .map(std::path::PathBuf::from)
338 .or_else(|| home.map(|h| h.join("token")));
339 let Some(path) = path else {
340 return Ok(None);
341 };
342
343 match std::fs::read_to_string(path) {
344 Ok(token) => Ok(Some(token.trim().to_owned())),
345 Err(e) if e.kind() == std::io::ErrorKind::NotFound => Ok(None),
346 Err(e) => Err(e.into()),
347 }
348}
349
350#[cfg(test)]
351#[path = "../../../tests/unit/model/index/hub.rs"]
352mod tests;