Skip to main content

cherenkov/model/index/
hub.rs

1//! Metadata-only HF discovery. Payload downloads use the existing HF/Xet client.
2use 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;