1use 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
23pub(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
34pub(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
45pub const MAX_SAFETENSORS_HEADER: u64 = 100_000_000;
47const MAX_INDEX_JSON: u64 = 64 * 1024 * 1024;
49const MAX_GGUF_HEADER: u64 = 1024 * 1024 * 1024;
52const MAX_GGUF_STRING: u64 = 16 * 1024 * 1024;
54const MAX_DIMS: usize = 8;
56const MAX_GGUF_COUNT: u64 = 1 << 20;
59const MAX_ARRAY_DEPTH: u32 = 4;
61const LIST_ITEMS_SHOWN: u64 = 16;
63const LIST_ITEM_CHARS: usize = 120;
65const MAX_SHARDS: usize = 100_000;
67
68#[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#[derive(Debug, Clone, PartialEq)]
86pub enum MetaValue {
87 Text(String),
89 List {
91 of: &'static str,
93 len: u64,
94 items: Vec<String>,
96 },
97}
98
99#[derive(Debug, Clone, PartialEq)]
101pub struct ModelSummary {
102 pub kind: ModelKind,
103 pub files: usize,
105 pub tensors: usize,
106 pub params: u64,
108 pub bytes: u64,
110 pub types: Vec<TypeShare>,
113 pub metadata: Vec<(String, MetaValue)>,
116}
117
118#[derive(Debug, Clone, PartialEq)]
120pub struct TypeShare {
121 pub name: String,
122 pub tensors: usize,
123 pub params: u64,
124}
125
126#[derive(Debug, Clone, PartialEq)]
128pub struct Tensor {
129 pub name: String,
130 pub dtype: String,
132 pub shape: Vec<u64>,
133 pub params: Option<u64>,
135 pub bytes: Option<u64>,
137 pub offset: u64,
140 pub offset_end: Option<u64>,
142}
143
144pub type Metadata = Vec<(String, MetaValue)>;
146
147#[derive(Debug, Clone, PartialEq)]
149pub struct Header {
150 pub kind: ModelKind,
151 pub tensors: Vec<Tensor>,
152 pub metadata: Vec<(String, MetaValue)>,
153}
154
155struct 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 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 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 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
255pub 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
265pub fn looks_like_gguf(head: &[u8]) -> bool {
267 head.starts_with(b"GGUF")
268}
269
270fn 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
288pub 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
307fn 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 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#[derive(serde::Deserialize)]
339struct StEntry {
340 dtype: String,
341 shape: StShape,
342 data_offsets: (u64, u64),
343}
344
345struct 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
374struct 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
429struct 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
457struct 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
511fn product(shape: &[u64]) -> Option<u64> {
513 shape.iter().try_fold(1u64, |acc, d| acc.checked_mul(*d))
514}
515
516fn 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 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
563const GGUF_DEFAULT_ALIGNMENT: u64 = 32;
565
566const GGUF_STRING: u32 = 8;
568const GGUF_ARRAY: u32 = 9;
569
570fn 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
581fn 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
592fn 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
611fn 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 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
664fn 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
686pub 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 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 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 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
799pub 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
809fn 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#[derive(serde::Deserialize)]
823struct StIndex {
824 #[serde(default)]
825 metadata: Option<StMetadata>,
826 weight_map: std::collections::BTreeMap<String, String>,
827}
828
829pub 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
836fn 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
858fn 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 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
882pub trait RangeSource {
885 fn get(&mut self, start: u64, end: u64) -> std::result::Result<(Vec<u8>, u64), RangeError>;
888}
889
890#[derive(Debug, Clone, PartialEq, Eq)]
892pub enum RangeError {
893 NoRanges,
896 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
906pub const FIRST_GGUF_RANGE: u64 = 256 * 1024;
911pub const FIRST_SAFETENSORS_RANGE: u64 = 64 * 1024;
915const MAX_RANGE: u64 = 16 * 1024 * 1024;
917const FIRST_INDEX_RANGE: u64 = 1024 * 1024;
919
920fn 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
946struct 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
988pub 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
1004pub 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 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
1047fn 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
1066pub type OpenRanges<'a> =
1068 dyn Fn(&str) -> std::result::Result<Box<dyn RangeSource>, RangeError> + Sync + 'a;
1069
1070pub 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
1079pub const SHARD_READS: usize = 8;
1083
1084pub 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
1090pub 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 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 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
1149fn 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 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 .ok_or((
1200 files.first().map_or("", String::as_str),
1201 RangeError::Failed("cancelled".to_string()),
1202 ))
1203}
1204
1205pub fn read_model(paths: &[PathBuf], format: FileFormat) -> Result<(LazyFrame, ModelSummary)> {
1208 let mut files: Vec<PathBuf> = Vec::new();
1209 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 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
1248fn 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
1254pub 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 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
1349fn 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
1381pub 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 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 first: true,
1418 own_columns: true,
1419 ..Default::default()
1420 }
1421}
1422
1423pub(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
1431fn 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 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 #[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 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 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 #[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 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 let mut w = GgufWriter::new(0, 1);
1693 w.u64(u64::MAX - 3);
1694 assert!(parse_header(&w.out).is_err());
1695 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 let mut w = GgufWriter::new(1, 0);
1701 w.str("t").u32(1_000_000);
1702 assert!(parse_header(&w.out).is_err());
1703 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 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 #[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 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 #[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 #[derive(Clone, Default)]
1781 pub(crate) struct Served {
1782 pub files: std::collections::BTreeMap<String, Vec<u8>>,
1783 pub asked: std::sync::Arc<std::sync::Mutex<Vec<(String, u64, u64)>>>,
1785 pub no_ranges: bool,
1786 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 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 #[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 #[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 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 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 #[test]
2004 fn a_hostile_remote_header_is_refused_before_it_is_fetched() {
2005 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 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 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 #[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 let n = served.asked.lock().unwrap().len();
2104 assert!((1..=SHARD_READS).contains(&n), "{n}");
2105 }
2106
2107 #[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 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 #[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 1 => ([bytes, vec![0; 4]].concat(), st.len() as u64),
2149 2 => (bytes, st.len() as u64),
2150 _ => (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 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 #[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 #[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 #[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}