Skip to main content

datui_lib/
model_files.rs

1//! SafeTensors and GGUF model files, read as a table of their tensors.
2//!
3//! Only the header is read: a 70 GB checkpoint opens as fast as a 7 MB one, because
4//! the tensor data is never touched. Both parsers are hand-written over a reader that
5//! knows where the header must end, and every length the file states is checked
6//! against that before anything is allocated or skipped, so a corrupt or hostile
7//! header is an error rather than a multi-gigabyte allocation or a panic.
8//!
9//! The rows are a small eager `DataFrame` (one per tensor) made lazy, as ORC and Excel
10//! are. What is not a row — the file's metadata and the totals — is a
11//! [`ModelSummary`], carried to the dataset for the Info panel's Model tab.
12
13use std::io::Read;
14use std::path::{Path, PathBuf};
15
16use color_eyre::Result;
17use color_eyre::eyre::eyre;
18use polars::prelude::*;
19
20use crate::FileFormat;
21use crate::error_display::FileError;
22
23/// What datui does with a SafeTensors file: see [`crate::readers`].
24pub(crate) const SAFETENSORS: crate::readers::Reader = crate::readers::Reader {
25    scan,
26    signatures: &[crate::readers::Signature {
27        says: |head, _| looks_like_safetensors(head),
28        kind: crate::readers::Kind::Magic,
29        trusted: crate::readers::EVERYWHERE,
30    }],
31    ..crate::readers::BASE
32};
33
34/// What datui does with a GGUF file: see [`crate::readers`].
35pub(crate) const GGUF: crate::readers::Reader = crate::readers::Reader {
36    scan,
37    signatures: &[crate::readers::Signature {
38        says: |head, _| looks_like_gguf(head),
39        kind: crate::readers::Kind::Magic,
40        trusted: crate::readers::EVERYWHERE,
41    }],
42    ..crate::readers::BASE
43};
44
45/// The largest SafeTensors header read: the limit the reference implementation sets.
46pub const MAX_SAFETENSORS_HEADER: u64 = 100_000_000;
47/// The largest `model.safetensors.index.json` read. Real ones are a few hundred KB.
48const MAX_INDEX_JSON: u64 = 64 * 1024 * 1024;
49/// Where a GGUF header must have ended. Real headers, vocabulary and merges included,
50/// are tens of MB; a length that reaches past this is corruption.
51const MAX_GGUF_HEADER: u64 = 1024 * 1024 * 1024;
52/// The longest single GGUF string kept. A chat template is a few KB.
53const MAX_GGUF_STRING: u64 = 16 * 1024 * 1024;
54/// The most dimensions a tensor may have. GGML uses four.
55const MAX_DIMS: usize = 8;
56/// The most tensors or key/value pairs one GGUF file may declare. Real files have a
57/// few thousand tensors and a few dozen pairs; each kept tensor costs about 150 bytes.
58const MAX_GGUF_COUNT: u64 = 1 << 20;
59/// How deep arrays of arrays are followed.
60const MAX_ARRAY_DEPTH: u32 = 4;
61/// An array is listed in full up to this many items, and summarized by length after.
62const LIST_ITEMS_SHOWN: u64 = 16;
63/// A string inside a listed array is cut to this many characters.
64const LIST_ITEM_CHARS: usize = 120;
65/// The most shard files an index or a directory may name.
66const MAX_SHARDS: usize = 100_000;
67
68/// Which of the two formats a model file is.
69#[derive(Debug, Clone, Copy, PartialEq, Eq)]
70pub enum ModelKind {
71    SafeTensors,
72    Gguf { version: u32 },
73}
74
75impl ModelKind {
76    pub fn label(self) -> String {
77        match self {
78            ModelKind::SafeTensors => "SafeTensors".to_string(),
79            ModelKind::Gguf { version } => format!("GGUF v{version}"),
80        }
81    }
82}
83
84/// One metadata value, as the Model tab shows it.
85#[derive(Debug, Clone, PartialEq)]
86pub enum MetaValue {
87    /// A scalar, or a string kept whole (a chat template is shown in full).
88    Text(String),
89    /// An array: listed when short, otherwise only its length and what it holds.
90    List {
91        /// What the items are, plural: `strings`, `integers`.
92        of: &'static str,
93        len: u64,
94        /// Every item, when the array is short enough to list; empty otherwise.
95        items: Vec<String>,
96    },
97}
98
99/// What a model's header says beyond its tensors, and their totals.
100#[derive(Debug, Clone, PartialEq)]
101pub struct ModelSummary {
102    pub kind: ModelKind,
103    /// The files the tensors come from.
104    pub files: usize,
105    pub tensors: usize,
106    /// The sum of every tensor's parameter count.
107    pub params: u64,
108    /// The sum of every tensor's bytes, where they are known.
109    pub bytes: u64,
110    /// Each dtype or quantization type, its tensors and its parameters, most
111    /// parameters first.
112    pub types: Vec<TypeShare>,
113    /// Key and value, in the order the file has them. Across several files the first
114    /// file to name a key gives its value.
115    pub metadata: Vec<(String, MetaValue)>,
116}
117
118/// One dtype's share of a model.
119#[derive(Debug, Clone, PartialEq)]
120pub struct TypeShare {
121    pub name: String,
122    pub tensors: usize,
123    pub params: u64,
124}
125
126/// One tensor, as a header describes it.
127#[derive(Debug, Clone, PartialEq)]
128pub struct Tensor {
129    pub name: String,
130    /// The SafeTensors dtype or the GGML type name.
131    pub dtype: String,
132    pub shape: Vec<u64>,
133    /// `None` when the product of the shape overflows.
134    pub params: Option<u64>,
135    /// `None` for a GGML type datui does not know the size of.
136    pub bytes: Option<u64>,
137    /// SafeTensors: the start of `data_offsets`. GGUF: the offset as written, from the
138    /// start of the tensor data.
139    pub offset: u64,
140    /// SafeTensors only: the end of `data_offsets`.
141    pub offset_end: Option<u64>,
142}
143
144/// Metadata as key and value, in the order the file has them.
145pub type Metadata = Vec<(String, MetaValue)>;
146
147/// One file's header.
148#[derive(Debug, Clone, PartialEq)]
149pub struct Header {
150    pub kind: ModelKind,
151    pub tensors: Vec<Tensor>,
152    pub metadata: Vec<(String, MetaValue)>,
153}
154
155/// A reader that knows where the header must end and refuses any length past it.
156struct Bounded<R> {
157    inner: R,
158    pos: u64,
159    end: u64,
160    big_endian: bool,
161}
162
163impl<R: Read> Bounded<R> {
164    fn left(&self) -> u64 {
165        self.end.saturating_sub(self.pos)
166    }
167
168    /// Fail unless `n` more bytes fit before the end.
169    fn need(&self, n: u64, what: &str) -> Result<()> {
170        if n > self.left() {
171            return Err(eyre!(
172                "{what} runs past the end of the GGUF header ({n} bytes, {} left)",
173                self.left()
174            ));
175        }
176        Ok(())
177    }
178
179    fn fill<const N: usize>(&mut self, what: &str) -> Result<[u8; N]> {
180        self.need(N as u64, what)?;
181        let mut buf = [0u8; N];
182        self.inner
183            .read_exact(&mut buf)
184            .map_err(|e| eyre!("cannot read {what} in the GGUF header: {e}"))?;
185        self.pos += N as u64;
186        Ok(buf)
187    }
188
189    fn u8(&mut self, what: &str) -> Result<u8> {
190        Ok(self.fill::<1>(what)?[0])
191    }
192
193    fn u16(&mut self, what: &str) -> Result<u16> {
194        let b = self.fill::<2>(what)?;
195        Ok(if self.big_endian {
196            u16::from_be_bytes(b)
197        } else {
198            u16::from_le_bytes(b)
199        })
200    }
201
202    fn u32(&mut self, what: &str) -> Result<u32> {
203        let b = self.fill::<4>(what)?;
204        Ok(if self.big_endian {
205            u32::from_be_bytes(b)
206        } else {
207            u32::from_le_bytes(b)
208        })
209    }
210
211    fn u64(&mut self, what: &str) -> Result<u64> {
212        let b = self.fill::<8>(what)?;
213        Ok(if self.big_endian {
214            u64::from_be_bytes(b)
215        } else {
216            u64::from_le_bytes(b)
217        })
218    }
219
220    /// A length-prefixed string, kept. Invalid UTF-8 is replaced, not refused.
221    fn string(&mut self, what: &str) -> Result<String> {
222        let len = self.u64(what)?;
223        if len > MAX_GGUF_STRING {
224            return Err(eyre!(
225                "{what} is {len} bytes, longer than datui reads in a GGUF header"
226            ));
227        }
228        self.need(len, what)?;
229        let mut buf = Vec::new();
230        (&mut self.inner)
231            .take(len)
232            .read_to_end(&mut buf)
233            .map_err(|e| eyre!("cannot read {what} in the GGUF header: {e}"))?;
234        if buf.len() as u64 != len {
235            return Err(eyre!("{what} in the GGUF header is cut short"));
236        }
237        self.pos += len;
238        Ok(String::from_utf8_lossy(&buf).into_owned())
239    }
240
241    /// Read past `n` bytes. Read rather than sought: the skips are a vocabulary's
242    /// strings, a few bytes each, and a seek would throw the read buffer away for each.
243    fn skip(&mut self, n: u64, what: &str) -> Result<()> {
244        self.need(n, what)?;
245        let skipped = std::io::copy(&mut (&mut self.inner).take(n), &mut std::io::sink())
246            .map_err(|e| eyre!("cannot read {what} in the GGUF header: {e}"))?;
247        if skipped != n {
248            return Err(eyre!("{what} in the GGUF header is cut short"));
249        }
250        self.pos += n;
251        Ok(())
252    }
253}
254
255/// Whether the first bytes of a file are a SafeTensors header: a little-endian length
256/// a header could have, then the `{` that opens its JSON.
257pub fn looks_like_safetensors(head: &[u8]) -> bool {
258    if head.len() < 9 {
259        return false;
260    }
261    let len = u64::from_le_bytes(head[..8].try_into().expect("eight bytes"));
262    (2..=MAX_SAFETENSORS_HEADER).contains(&len) && head[8] == b'{'
263}
264
265/// Whether the first bytes of a file are GGUF's magic.
266pub fn looks_like_gguf(head: &[u8]) -> bool {
267    head.starts_with(b"GGUF")
268}
269
270/// The header length a SafeTensors file's first eight bytes state, refused when it is
271/// more than the spec allows or than the `len`-byte file holds.
272fn safetensors_header_len(prefix: [u8; 8], len: u64) -> Result<u64> {
273    let header_len = u64::from_le_bytes(prefix);
274    if header_len > MAX_SAFETENSORS_HEADER {
275        return Err(eyre!(
276            "the SafeTensors header is {header_len} bytes, more than the {MAX_SAFETENSORS_HEADER} allowed"
277        ));
278    }
279    if header_len > len.saturating_sub(8) {
280        return Err(eyre!(
281            "the SafeTensors header claims {header_len} bytes and the file has {}",
282            len.saturating_sub(8)
283        ));
284    }
285    Ok(header_len)
286}
287
288/// Read one SafeTensors header from `reader`, which holds `len` bytes in all.
289pub fn read_safetensors<R: Read>(reader: R, len: u64) -> Result<Header> {
290    let mut reader = reader;
291    let mut prefix = [0u8; 8];
292    reader
293        .read_exact(&mut prefix)
294        .map_err(|_| eyre!("the file is shorter than its SafeTensors header length"))?;
295    let header_len = safetensors_header_len(prefix, len)?;
296    let mut json = Vec::new();
297    reader
298        .take(header_len)
299        .read_to_end(&mut json)
300        .map_err(|e| eyre!("cannot read the SafeTensors header: {e}"))?;
301    if json.len() as u64 != header_len {
302        return Err(eyre!("the SafeTensors header is cut short"));
303    }
304    parse_safetensors_json(&json, len.saturating_sub(8).saturating_sub(header_len))
305}
306
307/// The header's JSON, without its length prefix. `data_len` is how many bytes of
308/// tensor data follow it, which every tensor's `data_offsets` must stay inside.
309///
310/// Read straight into what is kept rather than through `serde_json::Value`: a hostile
311/// 100 MB header of tiny arrays would be gigabytes as a `Value` tree. Fields the spec
312/// does not name are passed over without being stored, and keys are seen in the order
313/// the file has them, which is the order the metadata is shown in.
314fn parse_safetensors_json(json: &[u8], data_len: u64) -> Result<Header> {
315    let mut de = serde_json::Deserializer::from_slice(json);
316    let parsed = serde::Deserializer::deserialize_map(&mut de, StHeaderVisitor)
317        .and_then(|header| de.end().map(|()| header))
318        .map_err(|e| eyre!("the SafeTensors header is not valid: {e}"))?;
319    let (mut tensors, metadata) = parsed;
320    for t in &tensors {
321        if t.offset_end.is_some_and(|end| end > data_len) {
322            return Err(eyre!(
323                "tensor \"{}\" runs past the end of the file ({data_len} bytes of data)",
324                t.name
325            ));
326        }
327    }
328    // The order the data is in.
329    tensors.sort_by(|a, b| a.offset.cmp(&b.offset).then_with(|| a.name.cmp(&b.name)));
330    Ok(Header {
331        kind: ModelKind::SafeTensors,
332        tensors,
333        metadata,
334    })
335}
336
337/// One tensor's entry. Any other field is skipped, not kept.
338#[derive(serde::Deserialize)]
339struct StEntry {
340    dtype: String,
341    shape: StShape,
342    data_offsets: (u64, u64),
343}
344
345/// A shape, refused past [`MAX_DIMS`] before a longer list is stored.
346struct StShape(Vec<u64>);
347
348impl<'de> serde::Deserialize<'de> for StShape {
349    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
350        struct V;
351        impl<'de> serde::de::Visitor<'de> for V {
352            type Value = StShape;
353            fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
354                write!(f, "a list of at most {MAX_DIMS} dimensions")
355            }
356            fn visit_seq<A: serde::de::SeqAccess<'de>>(
357                self,
358                mut seq: A,
359            ) -> std::result::Result<StShape, A::Error> {
360                let mut dims = Vec::new();
361                while let Some(d) = seq.next_element::<u64>()? {
362                    if dims.len() == MAX_DIMS {
363                        return Err(serde::de::Error::custom("more dimensions than datui reads"));
364                    }
365                    dims.push(d);
366                }
367                Ok(StShape(dims))
368            }
369        }
370        d.deserialize_seq(V)
371    }
372}
373
374/// A `__metadata__` value. The spec says text; a number or a bool is shown as written,
375/// and anything nested is passed over and named by what it is.
376struct StMetaValue(String);
377
378impl<'de> serde::Deserialize<'de> for StMetaValue {
379    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
380        struct V;
381        impl<'de> serde::de::Visitor<'de> for V {
382            type Value = StMetaValue;
383            fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
384                f.write_str("a metadata value")
385            }
386            fn visit_str<E>(self, v: &str) -> std::result::Result<StMetaValue, E> {
387                Ok(StMetaValue(v.to_string()))
388            }
389            fn visit_string<E>(self, v: String) -> std::result::Result<StMetaValue, E> {
390                Ok(StMetaValue(v))
391            }
392            fn visit_bool<E>(self, v: bool) -> std::result::Result<StMetaValue, E> {
393                Ok(StMetaValue(v.to_string()))
394            }
395            fn visit_i64<E>(self, v: i64) -> std::result::Result<StMetaValue, E> {
396                Ok(StMetaValue(v.to_string()))
397            }
398            fn visit_u64<E>(self, v: u64) -> std::result::Result<StMetaValue, E> {
399                Ok(StMetaValue(v.to_string()))
400            }
401            fn visit_f64<E>(self, v: f64) -> std::result::Result<StMetaValue, E> {
402                Ok(StMetaValue(v.to_string()))
403            }
404            fn visit_unit<E>(self) -> std::result::Result<StMetaValue, E> {
405                Ok(StMetaValue("null".to_string()))
406            }
407            fn visit_seq<A: serde::de::SeqAccess<'de>>(
408                self,
409                mut seq: A,
410            ) -> std::result::Result<StMetaValue, A::Error> {
411                while seq.next_element::<serde::de::IgnoredAny>()?.is_some() {}
412                Ok(StMetaValue("[array]".to_string()))
413            }
414            fn visit_map<A: serde::de::MapAccess<'de>>(
415                self,
416                mut map: A,
417            ) -> std::result::Result<StMetaValue, A::Error> {
418                while map
419                    .next_entry::<serde::de::IgnoredAny, serde::de::IgnoredAny>()?
420                    .is_some()
421                {}
422                Ok(StMetaValue("{object}".to_string()))
423            }
424        }
425        d.deserialize_any(V)
426    }
427}
428
429/// `__metadata__`, in the order the file has it.
430struct StMetadata(Metadata);
431
432impl<'de> serde::Deserialize<'de> for StMetadata {
433    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
434        struct V;
435        impl<'de> serde::de::Visitor<'de> for V {
436            type Value = StMetadata;
437            fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
438                f.write_str("an object of metadata")
439            }
440            fn visit_map<A: serde::de::MapAccess<'de>>(
441                self,
442                mut map: A,
443            ) -> std::result::Result<StMetadata, A::Error> {
444                let mut out: Metadata = Vec::new();
445                while let Some((key, StMetaValue(value))) =
446                    map.next_entry::<String, StMetaValue>()?
447                {
448                    out.push((key, MetaValue::Text(value)));
449                }
450                Ok(StMetadata(out))
451            }
452        }
453        d.deserialize_map(V)
454    }
455}
456
457/// The whole header: its tensors, and `__metadata__`.
458struct StHeaderVisitor;
459
460impl<'de> serde::de::Visitor<'de> for StHeaderVisitor {
461    type Value = (Vec<Tensor>, Metadata);
462    fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
463        f.write_str("a JSON object of tensors")
464    }
465    fn visit_map<A: serde::de::MapAccess<'de>>(
466        self,
467        mut map: A,
468    ) -> std::result::Result<Self::Value, A::Error> {
469        use serde::de::Error;
470        let mut tensors = Vec::new();
471        let mut seen = std::collections::HashSet::new();
472        let mut metadata = None;
473        while let Some(name) = map.next_key::<String>()? {
474            if name == "__metadata__" {
475                if metadata.is_some() {
476                    return Err(A::Error::custom("__metadata__ appears twice"));
477                }
478                let StMetadata(m) = map
479                    .next_value()
480                    .map_err(|e| A::Error::custom(format!("__metadata__: {e}")))?;
481                metadata = Some(m);
482                continue;
483            }
484            let entry: StEntry = map
485                .next_value()
486                .map_err(|e| A::Error::custom(format!("tensor \"{name}\": {e}")))?;
487            if !seen.insert(name.clone()) {
488                return Err(A::Error::custom(format!("tensor \"{name}\" appears twice")));
489            }
490            let (start, end) = entry.data_offsets;
491            if end < start {
492                return Err(A::Error::custom(format!(
493                    "tensor \"{name}\" ends before it starts"
494                )));
495            }
496            let shape = entry.shape.0;
497            tensors.push(Tensor {
498                name,
499                dtype: entry.dtype,
500                params: product(&shape),
501                shape,
502                bytes: Some(end - start),
503                offset: start,
504                offset_end: Some(end),
505            });
506        }
507        Ok((tensors, metadata.unwrap_or_default()))
508    }
509}
510
511/// The product of a shape; 1 for a scalar, `None` on overflow.
512fn product(shape: &[u64]) -> Option<u64> {
513    shape.iter().try_fold(1u64, |acc, d| acc.checked_mul(*d))
514}
515
516/// A GGML type: its name, and how many elements a block of how many bytes holds.
517fn ggml_type(id: u32) -> Option<(&'static str, u64, u64)> {
518    Some(match id {
519        0 => ("F32", 1, 4),
520        1 => ("F16", 1, 2),
521        2 => ("Q4_0", 32, 18),
522        3 => ("Q4_1", 32, 20),
523        6 => ("Q5_0", 32, 22),
524        7 => ("Q5_1", 32, 24),
525        8 => ("Q8_0", 32, 34),
526        9 => ("Q8_1", 32, 36),
527        10 => ("Q2_K", 256, 84),
528        11 => ("Q3_K", 256, 110),
529        12 => ("Q4_K", 256, 144),
530        13 => ("Q5_K", 256, 176),
531        14 => ("Q6_K", 256, 210),
532        15 => ("Q8_K", 256, 292),
533        16 => ("IQ2_XXS", 256, 66),
534        17 => ("IQ2_XS", 256, 74),
535        18 => ("IQ3_XXS", 256, 98),
536        19 => ("IQ1_S", 256, 50),
537        20 => ("IQ4_NL", 32, 18),
538        21 => ("IQ3_S", 256, 110),
539        22 => ("IQ2_S", 256, 82),
540        23 => ("IQ4_XS", 256, 136),
541        24 => ("I8", 1, 1),
542        25 => ("I16", 1, 2),
543        26 => ("I32", 1, 4),
544        27 => ("I64", 1, 8),
545        28 => ("F64", 1, 8),
546        29 => ("IQ1_M", 256, 56),
547        30 => ("BF16", 1, 2),
548        // Repacked Q4_0 and IQ4_NL, since removed from GGML; files written by
549        // llama.cpp in late 2024 still carry them, with the same block sizes.
550        31 => ("Q4_0_4_4", 32, 18),
551        32 => ("Q4_0_4_8", 32, 18),
552        33 => ("Q4_0_8_8", 32, 18),
553        34 => ("TQ1_0", 256, 54),
554        35 => ("TQ2_0", 256, 66),
555        36 => ("IQ4_NL_4_4", 32, 18),
556        37 => ("IQ4_NL_4_8", 32, 18),
557        38 => ("IQ4_NL_8_8", 32, 18),
558        39 => ("MXFP4", 32, 17),
559        _ => return None,
560    })
561}
562
563/// Where tensor data starts when `general.alignment` does not say.
564const GGUF_DEFAULT_ALIGNMENT: u64 = 32;
565
566/// GGUF metadata value types.
567const GGUF_STRING: u32 = 8;
568const GGUF_ARRAY: u32 = 9;
569
570/// The size of a fixed-width GGUF value type; `None` for a string or an array.
571fn gguf_fixed_size(ty: u32) -> Option<u64> {
572    match ty {
573        0 | 1 | 7 => Some(1),
574        2 | 3 => Some(2),
575        4..=6 => Some(4),
576        10..=12 => Some(8),
577        _ => None,
578    }
579}
580
581/// What an array of `ty` holds, as its summary says it.
582fn gguf_items_noun(ty: u32) -> &'static str {
583    match ty {
584        0..=5 | 10 | 11 => "integers",
585        6 | 12 => "floats",
586        7 => "bools",
587        GGUF_STRING => "strings",
588        _ => "arrays",
589    }
590}
591
592/// One fixed-width value, as text.
593fn gguf_scalar<R: Read>(r: &mut Bounded<R>, ty: u32) -> Result<String> {
594    let what = "a metadata value";
595    Ok(match ty {
596        0 => r.u8(what)?.to_string(),
597        1 => (r.u8(what)? as i8).to_string(),
598        2 => r.u16(what)?.to_string(),
599        3 => (r.u16(what)? as i16).to_string(),
600        4 => r.u32(what)?.to_string(),
601        5 => (r.u32(what)? as i32).to_string(),
602        6 => f32::from_bits(r.u32(what)?).to_string(),
603        7 => (r.u8(what)? != 0).to_string(),
604        10 => r.u64(what)?.to_string(),
605        11 => (r.u64(what)? as i64).to_string(),
606        12 => f64::from_bits(r.u64(what)?).to_string(),
607        other => return Err(eyre!("unknown GGUF metadata type {other}")),
608    })
609}
610
611/// A value of type `ty`, read whole when it is kept and skipped where it is long.
612fn gguf_value<R: Read>(r: &mut Bounded<R>, ty: u32, depth: u32) -> Result<MetaValue> {
613    match ty {
614        GGUF_STRING => Ok(MetaValue::Text(r.string("a metadata string")?)),
615        GGUF_ARRAY => {
616            if depth >= MAX_ARRAY_DEPTH {
617                return Err(eyre!("GGUF arrays nest deeper than datui reads"));
618            }
619            let item_ty = r.u32("an array's type")?;
620            let len = r.u64("an array's length")?;
621            // Every item takes at least this much, so the length is checked against
622            // what is left before a single item is read.
623            let least = match item_ty {
624                GGUF_STRING => 8,
625                GGUF_ARRAY => 12,
626                t => gguf_fixed_size(t).ok_or_else(|| eyre!("unknown GGUF array type {t}"))?,
627            };
628            r.need(
629                len.checked_mul(least)
630                    .ok_or_else(|| eyre!("a GGUF array's length overflows"))?,
631                "an array",
632            )?;
633            let of = gguf_items_noun(item_ty);
634            let listed = len <= LIST_ITEMS_SHOWN && item_ty != GGUF_ARRAY;
635            if !listed {
636                skip_items(r, item_ty, len, depth)?;
637                return Ok(MetaValue::List {
638                    of,
639                    len,
640                    items: Vec::new(),
641                });
642            }
643            let mut items = Vec::with_capacity(len as usize);
644            for _ in 0..len {
645                let item = if item_ty == GGUF_STRING {
646                    let s = r.string("an array's string")?;
647                    let cut: String = s.chars().take(LIST_ITEM_CHARS).collect();
648                    if cut.len() < s.len() {
649                        format!("{cut}...")
650                    } else {
651                        cut
652                    }
653                } else {
654                    gguf_scalar(r, item_ty)?
655                };
656                items.push(item);
657            }
658            Ok(MetaValue::List { of, len, items })
659        }
660        t => Ok(MetaValue::Text(gguf_scalar(r, t)?)),
661    }
662}
663
664/// Step over `len` items of `ty` without keeping them.
665fn skip_items<R: Read>(r: &mut Bounded<R>, ty: u32, len: u64, depth: u32) -> Result<()> {
666    match ty {
667        GGUF_STRING => {
668            for _ in 0..len {
669                let n = r.u64("an array's string")?;
670                r.skip(n, "an array's string")?;
671            }
672        }
673        GGUF_ARRAY => {
674            for _ in 0..len {
675                gguf_value(r, GGUF_ARRAY, depth + 1)?;
676            }
677        }
678        t => {
679            let size = gguf_fixed_size(t).ok_or_else(|| eyre!("unknown GGUF array type {t}"))?;
680            r.skip(len.saturating_mul(size), "an array")?;
681        }
682    }
683    Ok(())
684}
685
686/// Read one GGUF header (versions 2 and 3, either byte order) from `reader`, which
687/// holds `len` bytes in all.
688pub fn read_gguf<R: Read>(reader: R, len: u64) -> Result<Header> {
689    let mut r = Bounded {
690        inner: reader,
691        pos: 0,
692        end: len.min(MAX_GGUF_HEADER),
693        big_endian: false,
694    };
695    let magic = r
696        .fill::<4>("the magic number")
697        .map_err(|_| eyre!("the file is too short to be GGUF"))?;
698    if &magic != b"GGUF" {
699        return Err(eyre!("not a GGUF file: it does not start with GGUF"));
700    }
701    let raw = r.fill::<4>("the version")?;
702    let mut version = u32::from_le_bytes(raw);
703    // The magic reads the same either way; a big-endian file's version does not.
704    if version & 0xFFFF == 0 {
705        r.big_endian = true;
706        version = u32::from_be_bytes(raw);
707    }
708    match version {
709        2 | 3 => {}
710        1 => return Err(eyre!("GGUF version 1 files are not supported")),
711        v => return Err(eyre!("GGUF version {v} is not one datui reads (2 or 3)")),
712    }
713    let tensor_count = r.u64("the tensor count")?;
714    let kv_count = r.u64("the metadata count")?;
715    // A key/value pair takes at least 12 bytes and a tensor's entry at least 24, so
716    // either count is bounded by what is left before anything is allocated for it.
717    if tensor_count > MAX_GGUF_COUNT || tensor_count.saturating_mul(24) > r.left() {
718        return Err(eyre!(
719            "{tensor_count} tensors cannot fit in the GGUF header"
720        ));
721    }
722    if kv_count > MAX_GGUF_COUNT || kv_count.saturating_mul(12) > r.left() {
723        return Err(eyre!(
724            "{kv_count} metadata entries cannot fit in the GGUF header"
725        ));
726    }
727    let mut metadata = Vec::with_capacity(kv_count as usize);
728    for _ in 0..kv_count {
729        let key = r.string("a metadata key")?;
730        let ty = r.u32("a metadata type")?;
731        let value = gguf_value(&mut r, ty, 0).map_err(|e| eyre!("{e} (in \"{key}\")"))?;
732        metadata.push((key, value));
733    }
734    let mut tensors = Vec::with_capacity(tensor_count as usize);
735    for _ in 0..tensor_count {
736        let name = r.string("a tensor name")?;
737        let n_dims = r.u32("a tensor's dimension count")? as usize;
738        if n_dims > MAX_DIMS {
739            return Err(eyre!(
740                "tensor \"{name}\" has {n_dims} dimensions, more than datui reads"
741            ));
742        }
743        let mut shape = Vec::with_capacity(n_dims);
744        for _ in 0..n_dims {
745            shape.push(r.u64("a tensor dimension")?);
746        }
747        let ty = r.u32("a tensor's type")?;
748        let offset = r.u64("a tensor's offset")?;
749        let params = product(&shape);
750        let (dtype, bytes) = match ggml_type(ty) {
751            Some((name, block, size)) => (
752                name.to_string(),
753                params
754                    .filter(|p| p % block == 0)
755                    .and_then(|p| (p / block).checked_mul(size)),
756            ),
757            None => (format!("type {ty}"), None),
758        };
759        tensors.push(Tensor {
760            name,
761            dtype,
762            shape,
763            params,
764            bytes,
765            offset,
766            offset_end: None,
767        });
768    }
769    // The tensor data starts at the next multiple of the alignment after the header,
770    // and every tensor whose size is known must end inside the file: a download cut
771    // short is an error, not a table that looks whole.
772    let alignment = metadata
773        .iter()
774        .find(|(k, _)| k == "general.alignment")
775        .and_then(|(_, v)| match v {
776            MetaValue::Text(t) => t.parse::<u64>().ok(),
777            MetaValue::List { .. } => None,
778        })
779        .filter(|a| a.is_power_of_two())
780        .unwrap_or(GGUF_DEFAULT_ALIGNMENT);
781    let data_start = r.pos.next_multiple_of(alignment);
782    let data_len = len.saturating_sub(data_start);
783    for t in &tensors {
784        let end = t.bytes.and_then(|b| t.offset.checked_add(b));
785        if t.bytes.is_some() && end.is_none_or(|end| end > data_len) {
786            return Err(eyre!(
787                "tensor \"{}\" runs past the end of the file ({data_len} bytes of data)",
788                t.name
789            ));
790        }
791    }
792    Ok(Header {
793        kind: ModelKind::Gguf { version },
794        tensors,
795        metadata,
796    })
797}
798
799/// Parse `bytes` as whichever model header it starts with. For the fuzz target and the
800/// tests: both parsers over a slice, with no file.
801pub fn parse_header(bytes: &[u8]) -> Result<Header> {
802    if looks_like_gguf(bytes) {
803        read_gguf(bytes, bytes.len() as u64)
804    } else {
805        read_safetensors(bytes, bytes.len() as u64)
806    }
807}
808
809/// Read one file's header with the parser `format` names.
810fn read_file(path: &Path, format: FileFormat) -> Result<Header> {
811    let file = std::fs::File::open(path)?;
812    let len = file.metadata()?.len();
813    let reader = std::io::BufReader::new(file);
814    match format {
815        FileFormat::Gguf => read_gguf(reader, len),
816        _ => read_safetensors(reader, len),
817    }
818}
819
820/// A `model.safetensors.index.json`: the shard each tensor is in, and metadata. Any
821/// other field is skipped, not kept.
822#[derive(serde::Deserialize)]
823struct StIndex {
824    #[serde(default)]
825    metadata: Option<StMetadata>,
826    weight_map: std::collections::BTreeMap<String, String>,
827}
828
829/// Whether `path` is a SafeTensors index: `model.safetensors.index.json`.
830pub fn is_safetensors_index(path: &Path) -> bool {
831    path.file_name()
832        .and_then(|n| n.to_str())
833        .is_some_and(|n| n.to_ascii_lowercase().ends_with(".safetensors.index.json"))
834}
835
836/// The shards an index names, beside it and in name order, and the index's own
837/// metadata.
838fn read_index(path: &Path) -> Result<(Vec<PathBuf>, Metadata)> {
839    let named = |e: std::io::Error| crate::error_display::in_file(path, e.into());
840    let file = std::fs::File::open(path).map_err(named)?;
841    let len = file.metadata().map_err(named)?.len();
842    if len > MAX_INDEX_JSON {
843        return Err(FileError::new(
844            path,
845            format!("the index is {len} bytes, more than datui reads"),
846        )
847        .into());
848    }
849    let mut text = Vec::new();
850    file.take(MAX_INDEX_JSON)
851        .read_to_end(&mut text)
852        .map_err(named)?;
853    let (names, metadata) = parse_index(&text, &path.display().to_string())?;
854    let dir = path.parent().unwrap_or(Path::new(""));
855    Ok((names.iter().map(|name| dir.join(name)).collect(), metadata))
856}
857
858/// An index's shard names, each once and in name order, and its metadata. `named` is
859/// what errors call the index.
860fn parse_index(text: &[u8], named: &str) -> Result<(Vec<String>, Metadata)> {
861    let refused = |what: String| FileError::new(Path::new(named), what);
862    let index: StIndex = serde_json::from_slice(text)
863        .map_err(|e| refused(format!("not a SafeTensors index: {e}")))?;
864    let names: std::collections::BTreeSet<String> = index.weight_map.into_values().collect();
865    if names.len() > MAX_SHARDS {
866        return Err(refused("the index names too many shards".into()).into());
867    }
868    for name in &names {
869        // A shard is a file beside its index, never a path out of the directory.
870        let path = Path::new(name);
871        if path.components().count() != 1 || path.file_name().is_none() || name.contains('\\') {
872            return Err(refused(format!(
873                "the index names \"{name}\", which is not a file beside it"
874            ))
875            .into());
876        }
877    }
878    let metadata = index.metadata.map(|StMetadata(m)| m).unwrap_or_default();
879    Ok((names.into_iter().collect(), metadata))
880}
881
882/// Bytes of one remote object, fetched a range at a time: an HTTP server, or an object
883/// in a store.
884pub trait RangeSource {
885    /// Bytes `start..end` of the object, fewer only where the object ends first, and
886    /// the object's whole length. `start` is inside the object.
887    fn get(&mut self, start: u64, end: u64) -> std::result::Result<(Vec<u8>, u64), RangeError>;
888}
889
890/// Why a ranged read stopped.
891#[derive(Debug, Clone, PartialEq, Eq)]
892pub enum RangeError {
893    /// The server sent the whole file where a range was asked for: the header cannot be
894    /// read on its own, and the file is downloaded instead.
895    NoRanges,
896    /// Anything else, as the user is told it.
897    Failed(String),
898}
899
900impl From<color_eyre::Report> for RangeError {
901    fn from(e: color_eyre::Report) -> Self {
902        RangeError::Failed(e.to_string())
903    }
904}
905
906/// The first range a GGUF header is read in. Each read after it is twice the one
907/// before, up to [`MAX_RANGE`], so a header of a few KB costs one request and one with
908/// a vocabulary (5 to 10 MB) four or five. Larger, fewer requests fetch up to twice
909/// the header; see `a_vocabulary_sized_gguf_header_takes_a_few_ranges`.
910pub const FIRST_GGUF_RANGE: u64 = 256 * 1024;
911/// The first read of a SafeTensors file: its header's length and, for most files, the
912/// whole of its JSON in the same request. A checkpoint shard's header is a few KB to
913/// some tens of KB; one longer than this takes a second request, for the rest of it.
914pub const FIRST_SAFETENSORS_RANGE: u64 = 64 * 1024;
915/// The most one ranged request asks for.
916const MAX_RANGE: u64 = 16 * 1024 * 1024;
917/// The first read of a remote index; one larger than this takes a second.
918const FIRST_INDEX_RANGE: u64 = 1024 * 1024;
919
920/// `start..end` of `src`, checked: exactly the bytes asked for up to the object's end,
921/// and the same length the first answer gave, if there was one.
922fn fetch(
923    src: &mut dyn RangeSource,
924    start: u64,
925    end: u64,
926    known_len: Option<u64>,
927) -> std::result::Result<(Vec<u8>, u64), RangeError> {
928    let (bytes, len) = src.get(start, end)?;
929    if known_len.is_some_and(|known| known != len) {
930        return Err(RangeError::Failed(format!(
931            "the file changed size while its header was read ({} then {len} bytes)",
932            known_len.unwrap_or_default()
933        )));
934    }
935    let want = end.min(len).saturating_sub(start);
936    if bytes.len() as u64 != want {
937        return Err(RangeError::Failed(format!(
938            "asked for bytes {start}..{} and got {} bytes",
939            end.min(len),
940            bytes.len()
941        )));
942    }
943    Ok((bytes, len))
944}
945
946/// A remote object read forward through ranged requests that grow as the read goes on,
947/// and never past `limit`. One range is held at a time.
948struct Ranged<'a> {
949    src: &'a mut dyn RangeSource,
950    len: u64,
951    limit: u64,
952    buf: Vec<u8>,
953    buf_start: u64,
954    pos: u64,
955    next: u64,
956    stop: &'a dyn Fn() -> bool,
957}
958
959impl Read for Ranged<'_> {
960    fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
961        if self.pos >= self.limit || out.is_empty() {
962            return Ok(0);
963        }
964        let buf_end = self.buf_start + self.buf.len() as u64;
965        if self.pos < self.buf_start || self.pos >= buf_end {
966            if (self.stop)() {
967                return Err(std::io::Error::other("cancelled"));
968            }
969            let end = self.pos.saturating_add(self.next).min(self.limit);
970            let (bytes, _) = fetch(self.src, self.pos, end, Some(self.len)).map_err(|e| {
971                std::io::Error::other(match e {
972                    RangeError::NoRanges => "the server stopped serving byte ranges".to_string(),
973                    RangeError::Failed(message) => message,
974                })
975            })?;
976            self.buf = bytes;
977            self.buf_start = self.pos;
978            self.next = (self.next * 2).min(MAX_RANGE);
979        }
980        let at = (self.pos - self.buf_start) as usize;
981        let n = out.len().min(self.buf.len() - at);
982        out[..n].copy_from_slice(&self.buf[at..at + n]);
983        self.pos += n as u64;
984        Ok(n)
985    }
986}
987
988/// Read one model header from `src` with ranged requests: SafeTensors as its first
989/// [`FIRST_SAFETENSORS_RANGE`] and, when its JSON runs past that, the rest of the JSON
990/// and no further; GGUF forward in growing ranges until its tensor infos end. Every
991/// bound the file readers keep is kept. `stop` is asked before each request.
992pub fn read_header_ranged(
993    src: &mut dyn RangeSource,
994    format: FileFormat,
995    stop: &dyn Fn() -> bool,
996) -> std::result::Result<Header, RangeError> {
997    let first = match format {
998        FileFormat::Gguf => FIRST_GGUF_RANGE,
999        _ => FIRST_SAFETENSORS_RANGE,
1000    };
1001    read_header_ranged_from(src, format, first, stop)
1002}
1003
1004/// As [`read_header_ranged`], with the first range `first` (at least the 8 bytes of a
1005/// SafeTensors length): small, for the fuzz target and the tests, so a header crosses
1006/// many ranges.
1007pub fn read_header_ranged_from(
1008    src: &mut dyn RangeSource,
1009    format: FileFormat,
1010    first: u64,
1011    stop: &dyn Fn() -> bool,
1012) -> std::result::Result<Header, RangeError> {
1013    if format == FileFormat::Gguf {
1014        let (head, len) = fetch(src, 0, first.max(1), None)?;
1015        let reader = Ranged {
1016            src,
1017            len,
1018            limit: len.min(MAX_GGUF_HEADER),
1019            buf: head,
1020            buf_start: 0,
1021            pos: 0,
1022            next: first.max(1).saturating_mul(2).min(MAX_RANGE),
1023            stop,
1024        };
1025        return Ok(read_gguf(reader, len)?);
1026    }
1027    let (mut head, len) = fetch(src, 0, first.max(8), None)?;
1028    let prefix: [u8; 8] = head
1029        .get(..8)
1030        .and_then(|prefix| prefix.try_into().ok())
1031        .ok_or_else(|| eyre!("the file is shorter than its SafeTensors header length"))?;
1032    let header_len = safetensors_header_len(prefix, len)?;
1033    let end = 8 + header_len;
1034    // The first read holds the whole JSON, or the front of it: the rest is asked for
1035    // once, up to its end and no further.
1036    if (head.len() as u64) < end {
1037        if stop() {
1038            return Err(RangeError::Failed("cancelled".to_string()));
1039        }
1040        let rest = fetch(src, head.len() as u64, end, Some(len))?.0;
1041        head.extend(rest);
1042    }
1043    let json = &head[8..end as usize];
1044    Ok(parse_safetensors_json(json, len - end)?)
1045}
1046
1047/// A remote index: its shard names and metadata, read whole within
1048/// [`MAX_INDEX_JSON`].
1049fn read_index_ranged(
1050    src: &mut dyn RangeSource,
1051    named: &str,
1052) -> std::result::Result<(Vec<String>, Metadata), RangeError> {
1053    let (mut text, len) = fetch(src, 0, FIRST_INDEX_RANGE, None)?;
1054    if len > MAX_INDEX_JSON {
1055        return Err(RangeError::Failed(crate::error_display::file_message(
1056            Path::new(named),
1057            &format!("the index is {len} bytes, more than datui reads"),
1058        )));
1059    }
1060    if len > text.len() as u64 {
1061        text.extend(fetch(src, text.len() as u64, len, Some(len))?.0);
1062    }
1063    Ok(parse_index(&text, named)?)
1064}
1065
1066/// A ranged source for a URL.
1067pub type OpenRanges<'a> =
1068    dyn Fn(&str) -> std::result::Result<Box<dyn RangeSource>, RangeError> + Sync + 'a;
1069
1070/// How a remote model's files are reached: a source for a URL, and the URL of a file
1071/// named beside another. `open` and `stop` are called from the threads that read
1072/// shards at once ([`SHARD_READS`]).
1073pub struct Remote<'a> {
1074    pub open: &'a OpenRanges<'a>,
1075    pub sibling: &'a dyn Fn(&str, &str) -> String,
1076    pub stop: &'a (dyn Fn() -> bool + Sync),
1077}
1078
1079/// Shards whose headers are read at once. A model hub's checkpoint is up to some
1080/// hundreds of shards, each a request or two: one at a time, their round trips add up
1081/// to minutes. A few at once is most of the gain without a burst at the server.
1082pub const SHARD_READS: usize = 8;
1083
1084/// The last segment of a URL, without a query: what a file's row and its errors call it.
1085pub fn url_file_name(url: &str) -> &str {
1086    let path = url.split(['?', '#']).next().unwrap_or(url);
1087    path.rsplit('/').next().unwrap_or(path)
1088}
1089
1090/// Read `urls` — remote model files, or SafeTensors indexes that name them — as one
1091/// table of tensors, as [`read_model`] reads files on disk, fetching only their
1092/// headers. [`RangeError::NoRanges`] only for one file named on its own, which can be
1093/// downloaded instead; for shards it is an error that says so.
1094pub fn read_remote_model(
1095    urls: &[String],
1096    format: FileFormat,
1097    remote: &Remote,
1098) -> std::result::Result<(LazyFrame, ModelSummary), RangeError> {
1099    let no_ranges = |url: &str| {
1100        RangeError::Failed(crate::error_display::file_message(
1101            Path::new(url),
1102            "the server does not serve byte ranges, which reading a sharded model's headers needs",
1103        ))
1104    };
1105    // Each failure names the file it came from: a shard, or the index naming them.
1106    let named = |url: &str, e: RangeError| match e {
1107        RangeError::Failed(what) => {
1108            RangeError::Failed(crate::error_display::file_message(Path::new(url), &what))
1109        }
1110        e => e,
1111    };
1112    let mut files: Vec<String> = Vec::new();
1113    let mut seen = std::collections::HashSet::new();
1114    let mut metadata: Metadata = Vec::new();
1115    for url in urls {
1116        if format == FileFormat::Safetensors && is_safetensors_index(Path::new(url_file_name(url)))
1117        {
1118            let (names, index_meta) = (remote.open)(url)
1119                .and_then(|mut src| read_index_ranged(src.as_mut(), url))
1120                .map_err(|e| match e {
1121                    RangeError::NoRanges => no_ranges(url),
1122                    e => named(url, e),
1123                })?;
1124            merge_metadata(&mut metadata, index_meta);
1125            for name in names {
1126                let shard = (remote.sibling)(url, &name);
1127                if seen.insert(shard.clone()) {
1128                    files.push(shard);
1129                }
1130            }
1131        } else if seen.insert(url.clone()) {
1132            files.push(url.clone());
1133        }
1134    }
1135    if files.is_empty() {
1136        return Err(RangeError::Failed("no model files to read".to_string()));
1137    }
1138    // Named on its own, a file the server sends whole is downloaded instead.
1139    let alone = files.len() == 1 && urls.len() == 1 && files[0] == urls[0];
1140    let headers = read_headers(&files, format, remote).map_err(|(file, e)| match e {
1141        RangeError::NoRanges if alone => RangeError::NoRanges,
1142        RangeError::NoRanges => no_ranges(file),
1143        e => named(file, e),
1144    })?;
1145    let names: Vec<String> = files.iter().map(|f| url_file_name(f).to_string()).collect();
1146    Ok(build(&headers, &names, metadata)?)
1147}
1148
1149/// Each of `files`' headers, in their order, read [`SHARD_READS`] at a time. The first
1150/// read to fail is the error, with its file; once one has failed, or the open is
1151/// stopped, no more requests are made.
1152fn read_headers<'f>(
1153    files: &'f [String],
1154    format: FileFormat,
1155    remote: &Remote,
1156) -> std::result::Result<Vec<Header>, (&'f str, RangeError)> {
1157    use std::sync::Mutex;
1158    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1159    let next = AtomicUsize::new(0);
1160    let failed = AtomicBool::new(false);
1161    let first_error: Mutex<Option<(usize, RangeError)>> = Mutex::new(None);
1162    let read: Vec<Mutex<Option<Header>>> = files.iter().map(|_| Mutex::new(None)).collect();
1163    let stop = || failed.load(Ordering::Relaxed) || (remote.stop)();
1164    std::thread::scope(|scope| {
1165        for _ in 0..SHARD_READS.min(files.len()) {
1166            scope.spawn(|| {
1167                loop {
1168                    let at = next.fetch_add(1, Ordering::Relaxed);
1169                    if at >= files.len() || stop() {
1170                        return;
1171                    }
1172                    match (remote.open)(&files[at])
1173                        .and_then(|mut src| read_header_ranged(src.as_mut(), format, &stop))
1174                    {
1175                        Ok(header) => {
1176                            *read[at].lock().unwrap_or_else(|e| e.into_inner()) = Some(header);
1177                        }
1178                        Err(e) => {
1179                            // The others stop at their next request: only the first is
1180                            // the reason.
1181                            if !failed.swap(true, Ordering::Relaxed) {
1182                                *first_error.lock().unwrap_or_else(|e| e.into_inner()) =
1183                                    Some((at, e));
1184                            }
1185                            return;
1186                        }
1187                    }
1188                }
1189            });
1190        }
1191    });
1192    if let Some((at, e)) = first_error.into_inner().unwrap_or_else(|e| e.into_inner()) {
1193        return Err((&files[at], e));
1194    }
1195    read.into_iter()
1196        .map(|slot| slot.into_inner().unwrap_or_else(|e| e.into_inner()))
1197        .collect::<Option<Vec<Header>>>()
1198        // Stopped before every file was read.
1199        .ok_or((
1200            files.first().map_or("", String::as_str),
1201            RangeError::Failed("cancelled".to_string()),
1202        ))
1203}
1204
1205/// Read `paths` — model files, or SafeTensors indexes that name them — as one table of
1206/// tensors, with a `file` column when there is more than one file.
1207pub fn read_model(paths: &[PathBuf], format: FileFormat) -> Result<(LazyFrame, ModelSummary)> {
1208    let mut files: Vec<PathBuf> = Vec::new();
1209    // Each file once, however many indexes and names reach it.
1210    let mut seen = std::collections::HashSet::new();
1211    let mut metadata: Vec<(String, MetaValue)> = Vec::new();
1212    for path in paths {
1213        if format == FileFormat::Safetensors && is_safetensors_index(path) {
1214            let (shards, index_meta) = read_index(path)?;
1215            merge_metadata(&mut metadata, index_meta);
1216            for shard in shards {
1217                if seen.insert(shard.clone()) {
1218                    files.push(shard);
1219                }
1220            }
1221        } else if seen.insert(path.clone()) {
1222            files.push(path.clone());
1223        }
1224    }
1225    if files.is_empty() {
1226        return Err(eyre!("No model files to read"));
1227    }
1228    let mut headers = Vec::with_capacity(files.len());
1229    for file in &files {
1230        // The open names the path it was given; of several files, say which one.
1231        let header = read_file(file, format).map_err(|e| match files.len() {
1232            1 => e,
1233            _ => crate::error_display::in_file(file, e),
1234        })?;
1235        headers.push(header);
1236    }
1237    let names: Vec<String> = files
1238        .iter()
1239        .map(|f| {
1240            f.file_name()
1241                .map(|n| n.to_string_lossy().into_owned())
1242                .unwrap_or_else(|| f.display().to_string())
1243        })
1244        .collect();
1245    build(&headers, &names, metadata)
1246}
1247
1248/// Keep each key's first value.
1249fn merge_metadata(into: &mut Vec<(String, MetaValue)>, from: Vec<(String, MetaValue)>) {
1250    let mut seen: std::collections::HashSet<String> = into.iter().map(|(k, _)| k.clone()).collect();
1251    into.extend(from.into_iter().filter(|(key, _)| seen.insert(key.clone())));
1252}
1253
1254/// The table and the summary for headers already read; `names` are their files.
1255pub fn build(
1256    headers: &[Header],
1257    names: &[String],
1258    mut metadata: Vec<(String, MetaValue)>,
1259) -> Result<(LazyFrame, ModelSummary)> {
1260    let kind = headers
1261        .first()
1262        .map(|h| h.kind)
1263        .ok_or_else(|| eyre!("No model files to read"))?;
1264    let safetensors = kind == ModelKind::SafeTensors;
1265    let many = headers.len() > 1;
1266    let rows: usize = headers.iter().map(|h| h.tensors.len()).sum();
1267
1268    let mut file_col = Vec::with_capacity(if many { rows } else { 0 });
1269    let mut name = Vec::with_capacity(rows);
1270    let mut dtype = Vec::with_capacity(rows);
1271    let values: usize = headers
1272        .iter()
1273        .flat_map(|h| &h.tensors)
1274        .map(|t| t.shape.len())
1275        .sum();
1276    let mut shape = ListPrimitiveChunkedBuilder::<UInt64Type>::new(
1277        "shape".into(),
1278        rows,
1279        values,
1280        DataType::UInt64,
1281    );
1282    let mut params = Vec::with_capacity(rows);
1283    let mut bytes = Vec::with_capacity(rows);
1284    let mut start = Vec::with_capacity(rows);
1285    let mut end = Vec::with_capacity(rows);
1286    // By name, then into a list: a hostile header can name a type per tensor.
1287    let mut types: std::collections::HashMap<&str, TypeShare> = Default::default();
1288    let (mut total_params, mut total_bytes) = (0u64, 0u64);
1289    merge_metadata(
1290        &mut metadata,
1291        headers.iter().flat_map(|h| h.metadata.clone()).collect(),
1292    );
1293    for (header, file) in headers.iter().zip(names) {
1294        for t in &header.tensors {
1295            if many {
1296                file_col.push(file.as_str());
1297            }
1298            name.push(t.name.as_str());
1299            dtype.push(t.dtype.as_str());
1300            shape.append_slice(&t.shape);
1301            params.push(t.params);
1302            bytes.push(t.bytes);
1303            start.push(t.offset);
1304            end.push(t.offset_end);
1305            let p = t.params.unwrap_or(0);
1306            total_params = total_params.saturating_add(p);
1307            total_bytes = total_bytes.saturating_add(t.bytes.unwrap_or(0));
1308            let share = types.entry(t.dtype.as_str()).or_insert_with(|| TypeShare {
1309                name: t.dtype.clone(),
1310                tensors: 0,
1311                params: 0,
1312            });
1313            share.tensors += 1;
1314            share.params = share.params.saturating_add(p);
1315        }
1316    }
1317    let mut types: Vec<TypeShare> = types.into_values().collect();
1318    types.sort_by(|a, b| b.params.cmp(&a.params).then_with(|| a.name.cmp(&b.name)));
1319
1320    let shape = shape.finish().into_series();
1321    let mut columns: Vec<Column> = Vec::new();
1322    if many {
1323        columns.push(Series::new("file".into(), file_col).into());
1324    }
1325    columns.push(Series::new("name".into(), name).into());
1326    columns.push(Series::new(if safetensors { "dtype" } else { "type" }.into(), dtype).into());
1327    columns.push(shape.into());
1328    columns.push(Series::new("params".into(), params).into());
1329    columns.push(Series::new("bytes".into(), bytes).into());
1330    if safetensors {
1331        columns.push(Series::new("offset_start".into(), start).into());
1332        columns.push(Series::new("offset_end".into(), end).into());
1333    } else {
1334        columns.push(Series::new("offset".into(), start).into());
1335    }
1336    let df = DataFrame::new(rows, columns)?;
1337    let summary = ModelSummary {
1338        kind,
1339        files: headers.len(),
1340        tensors: rows,
1341        params: total_params,
1342        bytes: total_bytes,
1343        types,
1344        metadata,
1345    };
1346    Ok((df.lazy(), summary))
1347}
1348
1349/// Each type's share of the parameters, most first: `Q4_K 87% · Q6_K 12% · F32 <1%`.
1350/// By tensors when no tensor has a parameter count.
1351fn type_mix(types: &[TypeShare], sep: &str) -> String {
1352    let by_params = types.iter().any(|t| t.params > 0);
1353    let total: u64 = if by_params {
1354        types.iter().map(|t| t.params).fold(0, u64::saturating_add)
1355    } else {
1356        types.iter().map(|t| t.tensors as u64).sum()
1357    };
1358    types
1359        .iter()
1360        .map(|t| {
1361            let part = if by_params {
1362                t.params
1363            } else {
1364                t.tensors as u64
1365            };
1366            let pct = if total == 0 {
1367                0.0
1368            } else {
1369                part as f64 * 100.0 / total as f64
1370            };
1371            if pct > 0.0 && pct < 1.0 {
1372                format!("{} <1%", t.name)
1373            } else {
1374                format!("{} {:.0}%", t.name, pct)
1375            }
1376        })
1377        .collect::<Vec<_>>()
1378        .join(sep)
1379}
1380
1381/// The Model tab: the model's totals, then its metadata as key and value, every value
1382/// whole.
1383pub fn detail(model: &ModelSummary) -> crate::text_formats::Detail {
1384    use crate::widgets::info::{count_of, format_bytes, group_u64, short_count};
1385    let sep = format!(" {} ", crate::glyphs::get().middot);
1386    let mut head = model.kind.label();
1387    head.push_str(&sep);
1388    head.push_str(&count_of(model.tensors as u64, "tensor", "tensors"));
1389    if model.files > 1 {
1390        head.push_str(&sep);
1391        head.push_str(&count_of(model.files as u64, "file", "files"));
1392    }
1393    let mut lines = vec![
1394        head,
1395        format!(
1396            "Parameters: {}{}{sep}Size: {}",
1397            group_u64(model.params),
1398            // The short form only where it is shorter.
1399            if model.params >= 1000 {
1400                format!(" ({})", short_count(model.params))
1401            } else {
1402                String::new()
1403            },
1404            format_bytes(model.bytes)
1405        ),
1406    ];
1407    if !model.types.is_empty() {
1408        lines.push(format!("Types: {}", type_mix(&model.types, &sep)));
1409    }
1410    crate::text_formats::Detail {
1411        tab: crate::text_formats::tab(crate::FileFormat::Safetensors),
1412        lines,
1413        list_title: "Metadata",
1414        list: model.metadata.clone(),
1415        // A model's schema is the same seven columns every time; what is particular
1416        // to it is here.
1417        first: true,
1418        own_columns: true,
1419        ..Default::default()
1420    }
1421}
1422
1423/// What a model's header says besides its tensors, as the dataset takes it.
1424pub(crate) fn opened(summary: &ModelSummary) -> crate::members::Opened {
1425    crate::members::Opened {
1426        detail: Some(std::sync::Arc::new(detail(summary))),
1427        ..Default::default()
1428    }
1429}
1430
1431/// The scan of model files: their tensors, with the header's totals and metadata.
1432fn scan(input: crate::readers::ScanIn<'_>) -> Result<crate::scan::Scan> {
1433    let (lf, summary) = read_model(input.paths, input.format)?;
1434    input.report.opened = Some(std::sync::Arc::new(opened(&summary)));
1435    Ok(lf.into())
1436}
1437
1438#[cfg(test)]
1439pub(crate) mod tests {
1440    use super::*;
1441
1442    /// A SafeTensors file: the length, the JSON, then `data` bytes of tensor data.
1443    pub(crate) fn safetensors_bytes(json: &str, data: usize) -> Vec<u8> {
1444        let mut out = (json.len() as u64).to_le_bytes().to_vec();
1445        out.extend_from_slice(json.as_bytes());
1446        out.extend(std::iter::repeat_n(0u8, data));
1447        out
1448    }
1449
1450    /// Every way a model file is refused names the file, in the one shape.
1451    #[test]
1452    fn errors_name_the_file() {
1453        let past = r#"{"t":{"dtype":"F32","shape":[4],"data_offsets":[0,16]}}"#;
1454        crate::readers::bad_input::each_names_its_file(
1455            FileFormat::Safetensors,
1456            &[
1457                (
1458                    "claims.safetensors",
1459                    &[0xff, 0, 0, 0, 0, 0, 0, 0, b'{'],
1460                    "claims",
1461                ),
1462                (
1463                    "json.safetensors",
1464                    &safetensors_bytes("{nope", 0),
1465                    "not valid",
1466                ),
1467                (
1468                    "past.safetensors",
1469                    &safetensors_bytes(past, 8),
1470                    "Tensor \"t\"",
1471                ),
1472                (
1473                    "model.safetensors.index.json",
1474                    br#"{"weight_map":{"a":"../b.safetensors"}}"#,
1475                    "not a file beside it",
1476                ),
1477            ],
1478        );
1479        crate::readers::bad_input::each_names_its_file(
1480            FileFormat::Gguf,
1481            &[
1482                ("short.gguf", b"GG", "too short"),
1483                ("magic.gguf", b"GGML\x03\0\0\0", "does not start with GGUF"),
1484                ("v1.gguf", b"GGUF\x01\0\0\0", "version 1"),
1485            ],
1486        );
1487    }
1488
1489    /// Writes GGUF version 3, little-endian, as llama.cpp does.
1490    pub(crate) struct GgufWriter {
1491        pub out: Vec<u8>,
1492    }
1493
1494    impl GgufWriter {
1495        pub(crate) fn new(tensors: u64, kvs: u64) -> Self {
1496            let mut out = b"GGUF".to_vec();
1497            out.extend_from_slice(&3u32.to_le_bytes());
1498            out.extend_from_slice(&tensors.to_le_bytes());
1499            out.extend_from_slice(&kvs.to_le_bytes());
1500            Self { out }
1501        }
1502        pub(crate) fn str(&mut self, s: &str) -> &mut Self {
1503            self.out.extend_from_slice(&(s.len() as u64).to_le_bytes());
1504            self.out.extend_from_slice(s.as_bytes());
1505            self
1506        }
1507        pub(crate) fn u32(&mut self, v: u32) -> &mut Self {
1508            self.out.extend_from_slice(&v.to_le_bytes());
1509            self
1510        }
1511        pub(crate) fn u64(&mut self, v: u64) -> &mut Self {
1512            self.out.extend_from_slice(&v.to_le_bytes());
1513            self
1514        }
1515        /// Pad to the default alignment, then `n` bytes of tensor data.
1516        pub(crate) fn data(&mut self, n: usize) -> &mut Self {
1517            let padded = self.out.len().next_multiple_of(32);
1518            self.out.resize(padded + n, 0);
1519            self
1520        }
1521        pub(crate) fn kv_str(&mut self, key: &str, value: &str) -> &mut Self {
1522            self.str(key).u32(GGUF_STRING).str(value)
1523        }
1524        pub(crate) fn kv_u32(&mut self, key: &str, value: u32) -> &mut Self {
1525            self.str(key).u32(4).u32(value)
1526        }
1527        pub(crate) fn kv_strings(&mut self, key: &str, items: &[&str]) -> &mut Self {
1528            self.str(key).u32(GGUF_ARRAY).u32(GGUF_STRING);
1529            self.u64(items.len() as u64);
1530            for item in items {
1531                self.str(item);
1532            }
1533            self
1534        }
1535        pub(crate) fn tensor(
1536            &mut self,
1537            name: &str,
1538            shape: &[u64],
1539            ty: u32,
1540            offset: u64,
1541        ) -> &mut Self {
1542            self.str(name).u32(shape.len() as u32);
1543            for d in shape {
1544                self.u64(*d);
1545            }
1546            self.u32(ty).u64(offset)
1547        }
1548    }
1549
1550    #[test]
1551    fn safetensors_tensors_are_rows_in_data_order() {
1552        let json = r#"{"__metadata__":{"format":"pt"},
1553            "b.weight":{"dtype":"F32","shape":[2,3],"data_offsets":[8,32]},
1554            "a.bias":{"dtype":"BF16","shape":[4],"data_offsets":[0,8]}}"#;
1555        let bytes = safetensors_bytes(json, 32);
1556        let header = parse_header(&bytes).unwrap();
1557        assert_eq!(header.kind, ModelKind::SafeTensors);
1558        assert_eq!(
1559            header
1560                .tensors
1561                .iter()
1562                .map(|t| t.name.as_str())
1563                .collect::<Vec<_>>(),
1564            ["a.bias", "b.weight"],
1565            "the order the data is in"
1566        );
1567        let w = &header.tensors[1];
1568        assert_eq!(
1569            (w.params, w.bytes, w.offset, w.offset_end),
1570            (Some(6), Some(24), 8, Some(32))
1571        );
1572        assert_eq!(
1573            header.metadata,
1574            vec![("format".to_string(), MetaValue::Text("pt".to_string()))]
1575        );
1576    }
1577
1578    #[test]
1579    fn a_safetensors_header_longer_than_the_file_is_refused() {
1580        let mut bytes = safetensors_bytes("{}", 0);
1581        bytes[..8].copy_from_slice(&50_000_000u64.to_le_bytes());
1582        let err = parse_header(&bytes).unwrap_err().to_string();
1583        assert!(err.contains("claims"), "{err}");
1584        bytes[..8].copy_from_slice(&u64::MAX.to_le_bytes());
1585        assert!(parse_header(&bytes).is_err());
1586        for bad in [
1587            r#"{"t":{"dtype":"F32","shape":[2],"data_offsets":[8,0]}}"#,
1588            r#"{"t":{"dtype":"F32","shape":[-1],"data_offsets":[0,8]}}"#,
1589            r#"{"t":{"shape":[2],"data_offsets":[0,8]}}"#,
1590            r#"[1,2]"#,
1591            r#"{"__metadata__":3}"#,
1592        ] {
1593            assert!(parse_header(&safetensors_bytes(bad, 8)).is_err(), "{bad}");
1594        }
1595    }
1596
1597    /// The spec's rules a viewer can check from the header: one entry per name, two
1598    /// offsets inside the data, a shape of counts. A field it does not name is passed
1599    /// over, and `__metadata__` keeps the file's order.
1600    #[test]
1601    fn safetensors_entries_follow_the_spec() {
1602        let ok = r#"{"__metadata__":{"z":"1","a":"2","n":3},
1603            "t":{"dtype":"F32","shape":[2],"data_offsets":[0,8],"extra":[[1,2],{"x":1}]}}"#;
1604        let header = parse_header(&safetensors_bytes(ok, 8)).unwrap();
1605        let keys: Vec<&str> = header.metadata.iter().map(|(k, _)| k.as_str()).collect();
1606        assert_eq!(keys, ["z", "a", "n"], "the file's order");
1607        assert_eq!(header.metadata[2].1, MetaValue::Text("3".into()));
1608        for (bad, why) in [
1609            (
1610                r#"{"t":{"dtype":"F32","shape":[2],"data_offsets":[0,8]},
1611                    "t":{"dtype":"F32","shape":[2],"data_offsets":[0,8]}}"#,
1612                "appears twice",
1613            ),
1614            (
1615                r#"{"t":{"dtype":"F32","shape":[4],"data_offsets":[0,16]}}"#,
1616                "past the end",
1617            ),
1618            (
1619                r#"{"t":{"dtype":"F32","shape":[2],"data_offsets":[0,4,8]}}"#,
1620                "tensor \"t\"",
1621            ),
1622            (
1623                r#"{"t":{"dtype":"F32","shape":[1,1,1,1,1,1,1,1,1],"data_offsets":[0,4]}}"#,
1624                "dimensions",
1625            ),
1626        ] {
1627            let err = parse_header(&safetensors_bytes(bad, 8))
1628                .unwrap_err()
1629                .to_string();
1630            assert!(err.contains(why), "{bad}: {err}");
1631        }
1632    }
1633
1634    #[test]
1635    fn gguf_tensors_and_metadata_are_read() {
1636        let mut w = GgufWriter::new(2, 3);
1637        w.kv_str("general.architecture", "llama")
1638            .kv_u32("llama.context_length", 4096)
1639            .kv_strings("tokenizer.ggml.tokens", &["a"; 40]);
1640        w.tensor("token_embd.weight", &[256, 4], 12, 0).tensor(
1641            "output_norm.weight",
1642            &[256],
1643            0,
1644            576,
1645        );
1646        w.data(576 + 1024);
1647        let header = parse_header(&w.out).unwrap();
1648        assert_eq!(header.kind, ModelKind::Gguf { version: 3 });
1649        assert_eq!(header.metadata[0].1, MetaValue::Text("llama".into()));
1650        assert_eq!(header.metadata[1].1, MetaValue::Text("4096".into()));
1651        assert_eq!(
1652            header.metadata[2].1,
1653            MetaValue::List {
1654                of: "strings",
1655                len: 40,
1656                items: vec![]
1657            },
1658            "a long array is its length"
1659        );
1660        let embd = &header.tensors[0];
1661        assert_eq!(embd.dtype, "Q4_K");
1662        assert_eq!((embd.params, embd.bytes), (Some(1024), Some(4 * 144)));
1663        assert_eq!(header.tensors[1].bytes, Some(1024));
1664    }
1665
1666    #[test]
1667    fn a_big_endian_gguf_is_read() {
1668        let mut out = b"GGUF".to_vec();
1669        out.extend_from_slice(&3u32.to_be_bytes());
1670        out.extend_from_slice(&1u64.to_be_bytes());
1671        out.extend_from_slice(&0u64.to_be_bytes());
1672        out.extend_from_slice(&1u64.to_be_bytes());
1673        out.push(b'x');
1674        out.extend_from_slice(&1u32.to_be_bytes());
1675        out.extend_from_slice(&8u64.to_be_bytes());
1676        out.extend_from_slice(&1u32.to_be_bytes());
1677        out.extend_from_slice(&0u64.to_be_bytes());
1678        out.resize(out.len().next_multiple_of(32) + 16, 0);
1679        let header = parse_header(&out).unwrap();
1680        assert_eq!(header.tensors[0].shape, vec![8]);
1681        assert_eq!(header.tensors[0].dtype, "F16");
1682    }
1683
1684    #[test]
1685    fn corrupt_gguf_lengths_are_errors_not_allocations() {
1686        // Counts no header could hold.
1687        let w = GgufWriter::new(u64::MAX, 0);
1688        assert!(parse_header(&w.out).is_err());
1689        let w = GgufWriter::new(0, 1 << 40);
1690        assert!(parse_header(&w.out).is_err());
1691        // A string longer than the file.
1692        let mut w = GgufWriter::new(0, 1);
1693        w.u64(u64::MAX - 3);
1694        assert!(parse_header(&w.out).is_err());
1695        // An array longer than the file.
1696        let mut w = GgufWriter::new(0, 1);
1697        w.str("k").u32(GGUF_ARRAY).u32(4).u64(1 << 60);
1698        assert!(parse_header(&w.out).is_err());
1699        // Too many dimensions.
1700        let mut w = GgufWriter::new(1, 0);
1701        w.str("t").u32(1_000_000);
1702        assert!(parse_header(&w.out).is_err());
1703        // Version 1, and a version from the future.
1704        let mut w = GgufWriter::new(0, 0);
1705        w.out[4..8].copy_from_slice(&1u32.to_le_bytes());
1706        assert!(parse_header(&w.out).is_err());
1707        w.out[4..8].copy_from_slice(&9u32.to_le_bytes());
1708        assert!(parse_header(&w.out).is_err());
1709        // Cut short anywhere.
1710        let mut w = GgufWriter::new(1, 1);
1711        w.kv_str("general.name", "tiny")
1712            .tensor("t", &[4, 4], 0, 0)
1713            .data(64);
1714        for cut in 0..w.out.len() {
1715            assert!(parse_header(&w.out[..cut]).is_err(), "cut at {cut}");
1716        }
1717        assert!(parse_header(&w.out).is_ok());
1718    }
1719
1720    /// A GGUF whose tensor data is cut short is refused, measured from where the data
1721    /// starts: after the header, at `general.alignment` or 32.
1722    #[test]
1723    fn a_gguf_tensor_must_end_inside_the_file() {
1724        let tensor = |w: &mut GgufWriter| {
1725            w.tensor("t", &[4], 0, 0);
1726        };
1727        let mut w = GgufWriter::new(1, 0);
1728        tensor(&mut w);
1729        w.data(15);
1730        let err = parse_header(&w.out).unwrap_err().to_string();
1731        assert!(err.contains("past the end"), "{err}");
1732        w.data(16);
1733        assert!(parse_header(&w.out).is_ok());
1734
1735        // A wider alignment moves the start of the data further on.
1736        let mut w = GgufWriter::new(1, 1);
1737        w.kv_u32("general.alignment", 256);
1738        tensor(&mut w);
1739        let header_end = w.out.len();
1740        w.data(16);
1741        w.out.truncate(header_end.next_multiple_of(256) + 15);
1742        assert!(parse_header(&w.out).is_err(), "measured from 256");
1743        w.out.resize(header_end.next_multiple_of(256) + 16, 0);
1744        assert!(parse_header(&w.out).is_ok());
1745    }
1746
1747    /// A header can name a type per tensor and repeat a key many times over; the
1748    /// totals and the merge stay linear rather than comparing each against all.
1749    #[test]
1750    fn many_types_and_keys_build_in_one_pass() {
1751        let n = 50_000;
1752        let header = Header {
1753            kind: ModelKind::SafeTensors,
1754            tensors: (0..n)
1755                .map(|i| Tensor {
1756                    name: format!("t{i}"),
1757                    dtype: format!("X{i}"),
1758                    shape: vec![2],
1759                    params: Some(2),
1760                    bytes: Some(0),
1761                    offset: 0,
1762                    offset_end: Some(0),
1763                })
1764                .collect(),
1765            metadata: (0..n)
1766                .map(|i| (format!("k{}", i % 7), MetaValue::Text(i.to_string())))
1767                .collect(),
1768        };
1769        let (_, summary) = build(&[header], &["a".into()], vec![]).unwrap();
1770        assert_eq!(summary.types.len(), n);
1771        assert_eq!(summary.metadata.len(), 7, "each key once");
1772        assert_eq!(
1773            summary.metadata[0].1,
1774            MetaValue::Text("0".into()),
1775            "the first"
1776        );
1777    }
1778
1779    /// Bytes served by range, counting what is asked of them.
1780    #[derive(Clone, Default)]
1781    pub(crate) struct Served {
1782        pub files: std::collections::BTreeMap<String, Vec<u8>>,
1783        /// Each request: the file, and the range.
1784        pub asked: std::sync::Arc<std::sync::Mutex<Vec<(String, u64, u64)>>>,
1785        pub no_ranges: bool,
1786        /// How long each request takes, and how many are under way: now and at most.
1787        pub wait: std::time::Duration,
1788        pub busy: std::sync::Arc<(
1789            std::sync::atomic::AtomicUsize,
1790            std::sync::atomic::AtomicUsize,
1791        )>,
1792    }
1793
1794    struct ServedFile {
1795        served: Served,
1796        url: String,
1797    }
1798
1799    impl RangeSource for ServedFile {
1800        fn get(&mut self, start: u64, end: u64) -> std::result::Result<(Vec<u8>, u64), RangeError> {
1801            let bytes = self
1802                .served
1803                .files
1804                .get(&self.url)
1805                .ok_or_else(|| RangeError::Failed(format!("{}: 404", self.url)))?;
1806            if self.served.no_ranges {
1807                return Err(RangeError::NoRanges);
1808            }
1809            self.served
1810                .asked
1811                .lock()
1812                .unwrap()
1813                .push((self.url.clone(), start, end));
1814            if !self.served.wait.is_zero() {
1815                use std::sync::atomic::Ordering::SeqCst;
1816                let (now, most) = &*self.served.busy;
1817                most.fetch_max(now.fetch_add(1, SeqCst) + 1, SeqCst);
1818                std::thread::sleep(self.served.wait);
1819                now.fetch_sub(1, SeqCst);
1820            }
1821            let len = bytes.len() as u64;
1822            let (from, to) = (start.min(len) as usize, end.min(len) as usize);
1823            Ok((bytes[from..to].to_vec(), len))
1824        }
1825    }
1826
1827    impl Served {
1828        fn bytes(&self) -> u64 {
1829            self.asked
1830                .lock()
1831                .unwrap()
1832                .iter()
1833                .map(|(_, a, b)| b - a)
1834                .sum()
1835        }
1836
1837        /// The one GGUF file `url`, its header read from a first range of `first`.
1838        fn read_gguf_from(&self, url: &str, first: u64) -> Header {
1839            let mut src = ServedFile {
1840                served: self.clone(),
1841                url: url.to_string(),
1842            };
1843            read_header_ranged_from(&mut src, FileFormat::Gguf, first, &|| false).unwrap()
1844        }
1845
1846        fn read(
1847            &self,
1848            urls: &[&str],
1849            format: FileFormat,
1850        ) -> std::result::Result<(LazyFrame, ModelSummary), RangeError> {
1851            let open = |url: &str| -> std::result::Result<Box<dyn RangeSource>, RangeError> {
1852                Ok(Box::new(ServedFile {
1853                    served: self.clone(),
1854                    url: url.to_string(),
1855                }))
1856            };
1857            let sibling = |url: &str, name: &str| {
1858                format!("{}/{name}", url.rsplit_once('/').map_or(url, |(d, _)| d))
1859            };
1860            let urls: Vec<String> = urls.iter().map(|u| u.to_string()).collect();
1861            read_remote_model(
1862                &urls,
1863                format,
1864                &Remote {
1865                    open: &open,
1866                    sibling: &sibling,
1867                    stop: &|| false,
1868                },
1869            )
1870        }
1871    }
1872
1873    fn ranged(
1874        bytes: &[u8],
1875        format: FileFormat,
1876        first: u64,
1877    ) -> std::result::Result<Header, RangeError> {
1878        let served = Served {
1879            files: [("f".to_string(), bytes.to_vec())].into(),
1880            ..Default::default()
1881        };
1882        let mut src = ServedFile {
1883            served,
1884            url: "f".to_string(),
1885        };
1886        read_header_ranged_from(&mut src, format, first, &|| false)
1887    }
1888
1889    fn sample_gguf() -> Vec<u8> {
1890        let mut w = GgufWriter::new(2, 3);
1891        w.kv_str("general.architecture", "llama")
1892            .kv_u32("llama.context_length", 4096)
1893            .kv_strings("tokenizer.ggml.tokens", &["token"; 300]);
1894        w.tensor("token_embd.weight", &[256, 4], 12, 0).tensor(
1895            "output_norm.weight",
1896            &[256],
1897            0,
1898            576,
1899        );
1900        w.data(576 + 1024);
1901        w.out
1902    }
1903
1904    /// Read by range, a header is the same header the file reader finds, however small
1905    /// the ranges; and so is the error, for a header cut short anywhere.
1906    #[test]
1907    fn a_ranged_read_finds_what_the_file_reader_finds() {
1908        let st = safetensors_bytes(
1909            r#"{"__metadata__":{"format":"pt"},"x":{"dtype":"F16","shape":[2,2],"data_offsets":[0,8]}}"#,
1910            8,
1911        );
1912        let gguf = sample_gguf();
1913        for first in [1, 3, 64, FIRST_GGUF_RANGE] {
1914            assert_eq!(
1915                ranged(&gguf, FileFormat::Gguf, first).unwrap(),
1916                parse_header(&gguf).unwrap(),
1917                "first range {first}"
1918            );
1919        }
1920        assert_eq!(
1921            ranged(&st, FileFormat::Safetensors, 1).unwrap(),
1922            parse_header(&st).unwrap()
1923        );
1924        for cut in 1..gguf.len() / 4 {
1925            assert!(
1926                ranged(&gguf[..cut], FileFormat::Gguf, 7).is_err(),
1927                "cut at {cut}"
1928            );
1929        }
1930        for cut in 1..st.len() {
1931            assert!(
1932                ranged(&st[..cut], FileFormat::Safetensors, 7).is_err(),
1933                "cut at {cut}"
1934            );
1935        }
1936    }
1937
1938    /// SafeTensors asks for its first 64 KiB, which holds most headers whole, and
1939    /// for a longer header the rest of its JSON: never past it. A GGUF header stops
1940    /// being read where its tensor infos end, give or take the last range.
1941    #[test]
1942    fn only_the_header_is_fetched() {
1943        let url = "s3://b/m.safetensors";
1944        let json = r#"{"x":{"dtype":"F32","shape":[1024],"data_offsets":[0,4096]}}"#;
1945        let served = Served {
1946            files: [(url.to_string(), safetensors_bytes(json, 4096))].into(),
1947            ..Default::default()
1948        };
1949        assert!(served.read(&[url], FileFormat::Safetensors).is_ok());
1950        assert_eq!(
1951            *served.asked.lock().unwrap(),
1952            [(url.to_string(), 0, FIRST_SAFETENSORS_RANGE)],
1953            "one request"
1954        );
1955
1956        // A header of 3,000 tensors, past the first read, before a gigabyte of data.
1957        let tensors: Vec<String> = (0..3000)
1958            .map(|i| {
1959                format!(
1960                    r#""layer.{i}.weight":{{"dtype":"F32","shape":[1],"data_offsets":[{},{}]}}"#,
1961                    i * 4,
1962                    i * 4 + 4
1963                )
1964            })
1965            .collect();
1966        let json = format!("{{{}}}", tensors.join(","));
1967        let end = 8 + json.len() as u64;
1968        assert!(end > FIRST_SAFETENSORS_RANGE);
1969        let mut st = safetensors_bytes(&json, 3000 * 4);
1970        st.resize(st.len() + (1 << 20), 0);
1971        let served = Served {
1972            files: [(url.to_string(), st)].into(),
1973            ..Default::default()
1974        };
1975        let (_, summary) = served.read(&[url], FileFormat::Safetensors).unwrap();
1976        assert_eq!(summary.tensors, 3000);
1977        assert_eq!(
1978            *served.asked.lock().unwrap(),
1979            [
1980                (url.to_string(), 0, FIRST_SAFETENSORS_RANGE),
1981                (url.to_string(), FIRST_SAFETENSORS_RANGE, end)
1982            ],
1983            "the rest of the JSON, and no data"
1984        );
1985
1986        // Megabytes of data after a header of a few KB.
1987        let mut gguf = sample_gguf();
1988        gguf.resize(gguf.len() + 8 * 1024 * 1024, 0);
1989        let served = Served {
1990            files: [("https://h/m.gguf".to_string(), gguf)].into(),
1991            ..Default::default()
1992        };
1993        assert!(served.read(&["https://h/m.gguf"], FileFormat::Gguf).is_ok());
1994        assert_eq!(
1995            served.asked.lock().unwrap().len(),
1996            1,
1997            "one range for a small header"
1998        );
1999        assert!(served.bytes() <= FIRST_GGUF_RANGE, "{}", served.bytes());
2000    }
2001
2002    /// The lengths a hostile header states are refused before they are fetched.
2003    #[test]
2004    fn a_hostile_remote_header_is_refused_before_it_is_fetched() {
2005        // A SafeTensors length past the file, or past the spec.
2006        for claim in [50_000_000u64, MAX_SAFETENSORS_HEADER + 1, u64::MAX] {
2007            let mut st = safetensors_bytes("{}", 64);
2008            st[..8].copy_from_slice(&claim.to_le_bytes());
2009            let served = Served {
2010                files: [("u".to_string(), st)].into(),
2011                ..Default::default()
2012            };
2013            assert!(served.read(&["u"], FileFormat::Safetensors).is_err());
2014            assert_eq!(
2015                served.asked.lock().unwrap().len(),
2016                1,
2017                "only the first read, for {claim}"
2018            );
2019        }
2020        // A GGUF string or count longer than the file: one range, then the error.
2021        let mut w = GgufWriter::new(0, 1);
2022        w.u64(u64::MAX - 3);
2023        w.out.resize(1 << 20, 0);
2024        let served = Served {
2025            files: [("g".to_string(), w.out)].into(),
2026            ..Default::default()
2027        };
2028        assert!(served.read(&["g"], FileFormat::Gguf).is_err());
2029        assert_eq!(served.asked.lock().unwrap().len(), 1);
2030        let w = GgufWriter::new(u64::MAX, 0);
2031        assert!(ranged(&w.out, FileFormat::Gguf, 4).is_err());
2032    }
2033
2034    /// A GGUF header the size of a Llama 3 vocabulary (128k tokens, 280k merges, about
2035    /// 7 MB), at the front of a much larger file.
2036    fn llama3_sized_gguf() -> Vec<u8> {
2037        let tokens = vec!["tok_ab"; 128_256];
2038        let merges = vec!["Ġab Ġcdefg"; 280_147];
2039        let mut w = GgufWriter::new(291, 3);
2040        w.kv_str("general.architecture", "llama")
2041            .kv_strings("tokenizer.ggml.tokens", &tokens)
2042            .kv_strings("tokenizer.ggml.merges", &merges);
2043        for i in 0..291 {
2044            w.tensor(&format!("blk.{i}.attn_q.weight"), &[1], 0, i * 32);
2045        }
2046        w.data(291 * 32 + (32 << 20));
2047        w.out
2048    }
2049
2050    /// A stopped open asks for nothing more: not the next range of a header, nor the
2051    /// shards no read has started on.
2052    #[test]
2053    fn a_stopped_read_asks_for_nothing_more() {
2054        let served = Served {
2055            files: [("g".to_string(), llama3_sized_gguf())].into(),
2056            ..Default::default()
2057        };
2058        let asked = served.asked.clone();
2059        let stop = || !asked.lock().unwrap().is_empty();
2060        let mut src = ServedFile {
2061            served: served.clone(),
2062            url: "g".to_string(),
2063        };
2064        let err = read_header_ranged_from(&mut src, FileFormat::Gguf, 1024, &stop).unwrap_err();
2065        assert!(
2066            matches!(err, RangeError::Failed(ref m) if m.contains("cancelled")),
2067            "{err:?}"
2068        );
2069        assert_eq!(served.asked.lock().unwrap().len(), 1);
2070
2071        let st = safetensors_bytes(
2072            r#"{"x":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#,
2073            4,
2074        );
2075        let shards: Vec<String> = (0..SHARD_READS * 4).map(|i| format!("s{i:03}")).collect();
2076        let served = Served {
2077            files: shards.iter().map(|s| (s.clone(), st.clone())).collect(),
2078            ..Default::default()
2079        };
2080        let asked = served.asked.clone();
2081        let open = |url: &str| -> std::result::Result<Box<dyn RangeSource>, RangeError> {
2082            Ok(Box::new(ServedFile {
2083                served: served.clone(),
2084                url: url.to_string(),
2085            }))
2086        };
2087        let stop = || !asked.lock().unwrap().is_empty();
2088        let read = read_remote_model(
2089            &shards,
2090            FileFormat::Safetensors,
2091            &Remote {
2092                open: &open,
2093                sibling: &|_, name| name.to_string(),
2094                stop: &stop,
2095            },
2096        );
2097        assert!(
2098            matches!(read, Err(RangeError::Failed(ref m)) if m.to_lowercase().contains("cancelled")),
2099            "{:?}",
2100            read.err()
2101        );
2102        // At most each reader's request already under way when the stop came.
2103        let n = served.asked.lock().unwrap().len();
2104        assert!((1..=SHARD_READS).contains(&n), "{n}");
2105    }
2106
2107    /// A vocabulary-sized header costs a handful of requests and not much more than
2108    /// itself on the wire. Measured at 7.2 MB: a first range of 64 KiB took 7 requests
2109    /// (7.9 MiB), 256 KiB 5 (7.8 MiB), 1 MiB 4 (15 MiB).
2110    #[test]
2111    fn a_vocabulary_sized_gguf_header_takes_a_few_ranges() {
2112        let gguf = llama3_sized_gguf();
2113        let served = Served {
2114            files: [("g".to_string(), gguf.clone())].into(),
2115            ..Default::default()
2116        };
2117        let header = served.read_gguf_from("g", FIRST_GGUF_RANGE);
2118        assert_eq!(header.tensors.len(), 291);
2119        // The header and its few bytes of tensor data, before the padding.
2120        let end = (gguf.len() - (32 << 20)) as u64;
2121        assert!(
2122            served.asked.lock().unwrap().len() <= 5,
2123            "{:?}",
2124            served.asked.lock().unwrap()
2125        );
2126        assert!(served.bytes() < end * 2, "{} for {end}", served.bytes());
2127    }
2128
2129    /// A source that answers with more or fewer bytes than were asked for, or whose
2130    /// length changes between requests, is an error rather than a header.
2131    #[test]
2132    fn a_lying_source_is_an_error() {
2133        struct Liar(u32);
2134        impl RangeSource for Liar {
2135            fn get(
2136                &mut self,
2137                start: u64,
2138                end: u64,
2139            ) -> std::result::Result<(Vec<u8>, u64), RangeError> {
2140                self.0 += 1;
2141                let st = safetensors_bytes(
2142                    r#"{"x":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#,
2143                    4,
2144                );
2145                let bytes = st[start as usize..end.min(st.len() as u64) as usize].to_vec();
2146                Ok(match self.0 {
2147                    // More than asked.
2148                    1 => ([bytes, vec![0; 4]].concat(), st.len() as u64),
2149                    2 => (bytes, st.len() as u64),
2150                    // A length that changed.
2151                    _ => (bytes, 1 << 30),
2152                })
2153            }
2154        }
2155        let err = read_header_ranged_from(&mut Liar(0), FileFormat::Safetensors, 8, &|| false)
2156            .unwrap_err();
2157        assert!(
2158            matches!(err, RangeError::Failed(ref m) if m.contains("got")),
2159            "{err:?}"
2160        );
2161        // From 8 bytes, so the JSON takes a second request.
2162        let err = read_header_ranged_from(&mut Liar(1), FileFormat::Safetensors, 8, &|| false)
2163            .unwrap_err();
2164        assert!(
2165            matches!(err, RangeError::Failed(ref m) if m.contains("changed size")),
2166            "{err:?}"
2167        );
2168    }
2169
2170    /// An index names its shards beside its own URL, each read once; a name that leaves
2171    /// the directory is refused.
2172    #[test]
2173    fn a_remote_index_resolves_its_shards_beside_it() {
2174        let shard = |n: u64| {
2175            safetensors_bytes(
2176                &format!(r#"{{"t{n}":{{"dtype":"F32","shape":[2],"data_offsets":[0,8]}}}}"#),
2177                8,
2178            )
2179        };
2180        let index = r#"{"metadata":{"total_size":16},"weight_map":{
2181            "t1":"model-00001-of-00002.safetensors","t2":"model-00002-of-00002.safetensors"}}"#;
2182        let served = Served {
2183            files: [
2184                (
2185                    "gs://b/m/model.safetensors.index.json".to_string(),
2186                    index.as_bytes().to_vec(),
2187                ),
2188                (
2189                    "gs://b/m/model-00001-of-00002.safetensors".to_string(),
2190                    shard(1),
2191                ),
2192                (
2193                    "gs://b/m/model-00002-of-00002.safetensors".to_string(),
2194                    shard(2),
2195                ),
2196            ]
2197            .into(),
2198            ..Default::default()
2199        };
2200        let (lf, summary) = served
2201            .read(
2202                &[
2203                    "gs://b/m/model.safetensors.index.json",
2204                    "gs://b/m/model-00001-of-00002.safetensors",
2205                ],
2206                FileFormat::Safetensors,
2207            )
2208            .unwrap();
2209        assert_eq!((summary.files, summary.tensors), (2, 2));
2210        assert_eq!(summary.metadata[0].0, "total_size");
2211        let df = lf.collect().unwrap();
2212        let files: Vec<&str> = df
2213            .column("file")
2214            .unwrap()
2215            .str()
2216            .unwrap()
2217            .iter()
2218            .map(|v| v.unwrap())
2219            .collect();
2220        assert_eq!(
2221            files,
2222            [
2223                "model-00001-of-00002.safetensors",
2224                "model-00002-of-00002.safetensors"
2225            ]
2226        );
2227
2228        for bad in [
2229            "../x.safetensors",
2230            "a/b.safetensors",
2231            "..",
2232            "a\\\\b.safetensors",
2233        ] {
2234            let index = format!(r#"{{"weight_map":{{"t":"{bad}"}}}}"#);
2235            let served = Served {
2236                files: [(
2237                    "i/model.safetensors.index.json".to_string(),
2238                    index.into_bytes(),
2239                )]
2240                .into(),
2241                ..Default::default()
2242            };
2243            let err = served
2244                .read(&["i/model.safetensors.index.json"], FileFormat::Safetensors)
2245                .err()
2246                .expect("an error");
2247            assert!(
2248                matches!(err, RangeError::Failed(ref m) if m.contains("not a file beside it")),
2249                "{bad}: {err:?}"
2250            );
2251        }
2252    }
2253
2254    /// A checkpoint's shards are read a few at a time, one request each, and come out
2255    /// in the index's order however their reads finish; a shard that fails is named.
2256    #[test]
2257    fn shards_are_read_a_few_at_a_time() {
2258        let n = SHARD_READS * 3;
2259        let names: Vec<String> = (1..=n)
2260            .map(|i| format!("model-{i:05}-of-{n:05}.safetensors"))
2261            .collect();
2262        let map: Vec<String> = names
2263            .iter()
2264            .enumerate()
2265            .map(|(i, name)| format!(r#""t{i}":"{name}""#))
2266            .collect();
2267        let index = format!(r#"{{"weight_map":{{{}}}}}"#, map.join(","));
2268        let mut files: std::collections::BTreeMap<String, Vec<u8>> = names
2269            .iter()
2270            .enumerate()
2271            .map(|(i, name)| {
2272                let json =
2273                    format!(r#"{{"t{i}":{{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}}}"#);
2274                (format!("h/{name}"), safetensors_bytes(&json, 4))
2275            })
2276            .collect();
2277        files.insert(
2278            "h/model.safetensors.index.json".to_string(),
2279            index.into_bytes(),
2280        );
2281        let served = Served {
2282            files,
2283            wait: std::time::Duration::from_millis(20),
2284            ..Default::default()
2285        };
2286        let (lf, summary) = served
2287            .read(&["h/model.safetensors.index.json"], FileFormat::Safetensors)
2288            .unwrap();
2289        assert_eq!((summary.files, summary.tensors), (n, n));
2290        let df = lf.collect().unwrap();
2291        let read: Vec<&str> = df
2292            .column("file")
2293            .unwrap()
2294            .str()
2295            .unwrap()
2296            .iter()
2297            .flatten()
2298            .collect();
2299        assert_eq!(read, names, "in the index's order");
2300        assert_eq!(
2301            served.asked.lock().unwrap().len(),
2302            1 + n,
2303            "one request a shard"
2304        );
2305        let most = served.busy.1.load(std::sync::atomic::Ordering::SeqCst);
2306        assert!((2..=SHARD_READS).contains(&most), "{most} at once");
2307
2308        let mut broken = served.clone();
2309        broken.wait = std::time::Duration::ZERO;
2310        broken
2311            .files
2312            .insert(format!("h/{}", names[5]), b"not a header".to_vec());
2313        let err = broken
2314            .read(&["h/model.safetensors.index.json"], FileFormat::Safetensors)
2315            .err()
2316            .expect("an error");
2317        assert!(
2318            matches!(err, RangeError::Failed(ref m) if m.starts_with(&format!("\"h/{}\": ", names[5]))),
2319            "{err:?}"
2320        );
2321    }
2322
2323    /// A server that sends whole files: one file named on its own is downloaded
2324    /// instead; shards cannot be, and say why.
2325    #[test]
2326    fn no_ranges_is_a_download_for_one_file_only() {
2327        let served = Served {
2328            files: [
2329                ("h/m.safetensors".to_string(), safetensors_bytes("{}", 0)),
2330                (
2331                    "h/model.safetensors.index.json".to_string(),
2332                    br#"{"weight_map":{}}"#.to_vec(),
2333                ),
2334            ]
2335            .into(),
2336            no_ranges: true,
2337            ..Default::default()
2338        };
2339        assert_eq!(
2340            served
2341                .read(&["h/m.safetensors"], FileFormat::Safetensors)
2342                .err()
2343                .expect("an error"),
2344            RangeError::NoRanges
2345        );
2346        let err = served
2347            .read(&["h/model.safetensors.index.json"], FileFormat::Safetensors)
2348            .err()
2349            .expect("an error");
2350        assert!(
2351            matches!(err, RangeError::Failed(ref m) if m.contains("byte ranges")),
2352            "{err:?}"
2353        );
2354    }
2355
2356    #[test]
2357    fn the_frame_and_the_summary_agree() {
2358        let st = parse_header(&safetensors_bytes(
2359            r#"{"x":{"dtype":"F16","shape":[2,2],"data_offsets":[0,8]},
2360                "y":{"dtype":"F32","shape":[3],"data_offsets":[8,20]}}"#,
2361            20,
2362        ))
2363        .unwrap();
2364        let (lf, summary) = build(
2365            &[st.clone(), st],
2366            &["a.safetensors".into(), "b.safetensors".into()],
2367            vec![],
2368        )
2369        .unwrap();
2370        let df = lf.collect().unwrap();
2371        assert_eq!(
2372            df.get_column_names()
2373                .iter()
2374                .map(|n| n.as_str())
2375                .collect::<Vec<_>>(),
2376            [
2377                "file",
2378                "name",
2379                "dtype",
2380                "shape",
2381                "params",
2382                "bytes",
2383                "offset_start",
2384                "offset_end"
2385            ]
2386        );
2387        assert_eq!(df.height(), 4);
2388        assert_eq!((summary.files, summary.tensors), (2, 4));
2389        assert_eq!((summary.params, summary.bytes), (14, 40));
2390        assert_eq!(summary.types[0].name, "F16", "most parameters first");
2391    }
2392}