use std::collections::HashMap;
use std::fs::File;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use memmap2::Mmap;
use serde_json::Value;
use thiserror::Error;
use tracing::{debug, info};
use crate::ir::lazy::{LazyMeta, LazyTensor, LazyTensorMap, MaterializeError};
use crate::ir::{DType, TensorMap};
use crate::progress::ProgressReporter;
#[derive(Error, Debug)]
pub enum SafetensorsError {
#[error("No safetensors files found in {path}")]
NoFiles { path: String },
#[error("Failed to read shard '{shard}': {source}")]
ShardReadError {
shard: String,
source: std::io::Error,
},
#[error("Failed to memory-map shard '{shard}': {source}")]
MmapError {
shard: String,
source: std::io::Error,
},
#[error("Failed to parse safetensors header in '{shard}': {reason}")]
HeaderParseError { shard: String, reason: String },
#[error("Failed to parse index.json: {0}")]
IndexParseError(String),
#[error("Unsupported dtype '{dtype}' in tensor '{tensor}'")]
UnsupportedDtype { dtype: String, tensor: String },
#[error("Tensor '{tensor}' data range [{start}..{end}] exceeds file size {file_size} in shard '{shard}'")]
DataOutOfBounds {
tensor: String,
shard: String,
start: usize,
end: usize,
file_size: usize,
},
#[error("Tensor '{tensor}' header size {header_bytes} disagrees with shape×dtype {derived_bytes} in shard '{shard}'")]
HeaderSizeMismatch {
tensor: String,
shard: String,
header_bytes: usize,
derived_bytes: usize,
},
#[error("Materialise error: {0}")]
Materialize(#[from] MaterializeError),
#[error("I/O error: {0}")]
Io(#[from] std::io::Error),
}
pub fn read_tensors_lazy(
model_dir: &Path,
progress: &ProgressReporter,
) -> Result<LazyTensorMap, SafetensorsError> {
let shard_paths = discover_shards(model_dir)?;
if shard_paths.is_empty() {
return Err(SafetensorsError::NoFiles {
path: model_dir.display().to_string(),
});
}
info!(
shard_count = shard_paths.len(),
"Discovered safetensors shards (lazy)"
);
let pb = progress.bar(shard_paths.len() as u64, "Indexing shards");
let mut lazy_map = LazyTensorMap::new();
for shard_path in &shard_paths {
let shard_name = shard_path
.file_name()
.map(|n| n.to_string_lossy().to_string())
.unwrap_or_else(|| shard_path.display().to_string());
debug!(shard = %shard_name, "Indexing shard");
index_shard_lazy(shard_path, &shard_name, &mut lazy_map)?;
pb.inc(1);
}
pb.finish_with_message(format!(
"Indexed {} tensors from {} shards",
lazy_map.len(),
shard_paths.len()
));
Ok(lazy_map)
}
pub fn read_tensors(
model_dir: &Path,
progress: &ProgressReporter,
) -> Result<TensorMap, SafetensorsError> {
let lazy_map = read_tensors_lazy(model_dir, progress)?;
let len = lazy_map.len();
let tensor_map = lazy_map.materialize_all()?;
debug!(tensors = len, "Materialised lazy map → eager TensorMap");
Ok(tensor_map)
}
fn discover_shards(model_dir: &Path) -> Result<Vec<PathBuf>, SafetensorsError> {
let index_path = model_dir.join("model.safetensors.index.json");
if index_path.exists() {
return discover_shards_from_index(&index_path, model_dir);
}
let single_path = model_dir.join("model.safetensors");
if single_path.exists() {
return Ok(vec![single_path]);
}
let mut paths: Vec<PathBuf> = std::fs::read_dir(model_dir)?
.filter_map(|e| e.ok())
.map(|e| e.path())
.filter(|p| {
p.extension()
.map(|ext| ext == "safetensors")
.unwrap_or(false)
})
.collect();
paths.sort();
Ok(paths)
}
fn discover_shards_from_index(
index_path: &Path,
model_dir: &Path,
) -> Result<Vec<PathBuf>, SafetensorsError> {
let content = std::fs::read_to_string(index_path).map_err(|e| {
SafetensorsError::IndexParseError(format!("Failed to read {}: {}", index_path.display(), e))
})?;
let index: Value = serde_json::from_str(&content).map_err(|e| {
SafetensorsError::IndexParseError(format!(
"Failed to parse {}: {}",
index_path.display(),
e
))
})?;
let weight_map = index
.get("weight_map")
.and_then(|v| v.as_object())
.ok_or_else(|| {
SafetensorsError::IndexParseError("index.json missing weight_map".to_string())
})?;
let mut shard_names: Vec<String> = weight_map
.values()
.filter_map(|v| v.as_str().map(|s| s.to_string()))
.collect();
shard_names.sort();
shard_names.dedup();
let paths: Vec<PathBuf> = shard_names
.into_iter()
.map(|name| model_dir.join(&name))
.collect();
for path in &paths {
if !path.exists() {
return Err(SafetensorsError::ShardReadError {
shard: path.display().to_string(),
source: std::io::Error::new(
std::io::ErrorKind::NotFound,
format!("Shard file not found: {}", path.display()),
),
});
}
}
Ok(paths)
}
fn index_shard_lazy(
shard_path: &Path,
shard_name: &str,
lazy_map: &mut LazyTensorMap,
) -> Result<(), SafetensorsError> {
let file = File::open(shard_path).map_err(|e| SafetensorsError::ShardReadError {
shard: shard_name.to_string(),
source: e,
})?;
let mmap = unsafe {
Mmap::map(&file).map_err(|e| SafetensorsError::MmapError {
shard: shard_name.to_string(),
source: e,
})?
};
let file_size = mmap.len();
if file_size < 8 {
return Err(SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!("File too small ({} bytes)", file_size),
});
}
let header_size = u64::from_le_bytes(mmap[..8].try_into().map_err(|_| {
SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: "Failed to read header size".to_string(),
}
})?) as usize;
if 8 + header_size > file_size {
return Err(SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!(
"Header size ({}) exceeds file size ({})",
header_size,
file_size - 8
),
});
}
let header_bytes = &mmap[8..8 + header_size];
let header: HashMap<String, Value> =
serde_json::from_slice(header_bytes).map_err(|e| SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!("JSON parse error: {}", e),
})?;
let data_start = 8 + header_size;
let mmap = Arc::new(mmap);
for (name, info) in &header {
if name == "__metadata__" {
continue;
}
let dtype_str = info.get("dtype").and_then(|v| v.as_str()).ok_or_else(|| {
SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!("Tensor '{}' missing dtype", name),
}
})?;
let dtype = DType::from_safetensors_str(dtype_str).ok_or_else(|| {
SafetensorsError::UnsupportedDtype {
dtype: dtype_str.to_string(),
tensor: name.clone(),
}
})?;
let shape: Vec<usize> = info
.get("shape")
.and_then(|v| v.as_array())
.ok_or_else(|| SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!("Tensor '{}' missing shape", name),
})?
.iter()
.filter_map(|v| v.as_u64().map(|u| u as usize))
.collect();
let offsets = info
.get("data_offsets")
.and_then(|v| v.as_array())
.ok_or_else(|| SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!("Tensor '{}' missing data_offsets", name),
})?;
if offsets.len() != 2 {
return Err(SafetensorsError::HeaderParseError {
shard: shard_name.to_string(),
reason: format!(
"Tensor '{}' has {} data_offsets, expected 2",
name,
offsets.len()
),
});
}
let offset_start = offsets[0].as_u64().unwrap_or(0) as usize;
let offset_end = offsets[1].as_u64().unwrap_or(0) as usize;
let abs_start = data_start + offset_start;
let abs_end = data_start + offset_end;
if abs_end > file_size {
return Err(SafetensorsError::DataOutOfBounds {
tensor: name.clone(),
shard: shard_name.to_string(),
start: abs_start,
end: abs_end,
file_size,
});
}
let header_byte_len = abs_end - abs_start;
let meta = LazyMeta::new(name.clone(), shape, dtype);
if meta.byte_len != header_byte_len {
return Err(SafetensorsError::HeaderSizeMismatch {
tensor: name.clone(),
shard: shard_name.to_string(),
header_bytes: header_byte_len,
derived_bytes: meta.byte_len,
});
}
let mmap_clone = Arc::clone(&mmap);
let load = move || -> Result<Vec<u8>, MaterializeError> {
Ok(mmap_clone[abs_start..abs_end].to_vec())
};
let lazy = LazyTensor::from_closure(meta, load);
lazy_map.insert(lazy);
}
debug!(
shard = %shard_name,
tensor_count = lazy_map.len(),
"Indexed shard"
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
fn create_test_safetensors(tensors: &[(&str, &[usize], &str, &[u8])]) -> Vec<u8> {
let mut header_map = serde_json::Map::new();
let mut current_offset = 0usize;
for (name, shape, dtype, data) in tensors {
let mut tensor_info = serde_json::Map::new();
tensor_info.insert(
"dtype".to_string(),
serde_json::Value::String(dtype.to_string()),
);
tensor_info.insert(
"shape".to_string(),
serde_json::Value::Array(
shape
.iter()
.map(|&s| serde_json::Value::Number(s.into()))
.collect(),
),
);
let end_offset = current_offset + data.len();
tensor_info.insert(
"data_offsets".to_string(),
serde_json::Value::Array(vec![
serde_json::Value::Number(current_offset.into()),
serde_json::Value::Number(end_offset.into()),
]),
);
header_map.insert(name.to_string(), serde_json::Value::Object(tensor_info));
current_offset = end_offset;
}
let header_json = serde_json::to_string(&header_map).unwrap();
let header_bytes = header_json.as_bytes();
let header_size = header_bytes.len() as u64;
let mut file_data = Vec::new();
file_data.extend_from_slice(&header_size.to_le_bytes());
file_data.extend_from_slice(header_bytes);
for (_, _, _, data) in tensors {
file_data.extend_from_slice(data);
}
file_data
}
#[test]
fn test_read_single_shard() {
let tmp = tempfile::tempdir().unwrap();
let model_dir = tmp.path();
let tensor_data: Vec<u8> = (0..6u32).flat_map(|v| (v as f32).to_le_bytes()).collect();
let safetensors_data =
create_test_safetensors(&[("test_weight", &[2, 3], "F32", &tensor_data)]);
std::fs::write(model_dir.join("model.safetensors"), &safetensors_data).unwrap();
let progress = ProgressReporter::new();
let tensor_map = read_tensors(model_dir, &progress).unwrap();
assert_eq!(tensor_map.len(), 1);
let tensor = tensor_map.get("test_weight").unwrap();
assert_eq!(tensor.shape, vec![2, 3]);
assert_eq!(tensor.dtype, DType::F32);
assert_eq!(tensor.data.len(), 24); assert_eq!(*tensor.data, tensor_data);
}
#[test]
fn test_read_multiple_tensors() {
let tmp = tempfile::tempdir().unwrap();
let model_dir = tmp.path();
let weight_data: Vec<u8> = vec![0u8; 4 * 2]; let bias_data: Vec<u8> = vec![0u8; 2 * 2];
let safetensors_data = create_test_safetensors(&[
("layer.weight", &[2, 2], "F16", &weight_data),
("layer.bias", &[2], "F16", &bias_data),
]);
std::fs::write(model_dir.join("model.safetensors"), &safetensors_data).unwrap();
let progress = ProgressReporter::new();
let tensor_map = read_tensors(model_dir, &progress).unwrap();
assert_eq!(tensor_map.len(), 2);
assert!(tensor_map.get("layer.weight").is_some());
assert!(tensor_map.get("layer.bias").is_some());
}
#[test]
fn test_no_safetensors_error() {
let tmp = tempfile::tempdir().unwrap();
let progress = ProgressReporter::new();
let result = read_tensors(tmp.path(), &progress);
assert!(result.is_err());
}
#[test]
fn test_lazy_reader_metadata_only() {
let tmp = tempfile::tempdir().unwrap();
let model_dir = tmp.path();
let tensor_data: Vec<u8> = (0..6u32).flat_map(|v| (v as f32).to_le_bytes()).collect();
let safetensors_data = create_test_safetensors(&[("w", &[2, 3], "F32", &tensor_data)]);
std::fs::write(model_dir.join("model.safetensors"), &safetensors_data).unwrap();
let progress = ProgressReporter::new();
let lazy_map = read_tensors_lazy(model_dir, &progress).unwrap();
assert_eq!(lazy_map.len(), 1);
let lazy = lazy_map.get("w").unwrap();
assert_eq!(lazy.shape(), &[2, 3]);
assert_eq!(lazy.dtype(), DType::F32);
assert_eq!(lazy.byte_len(), 24);
let realised = lazy_map
.into_iter()
.next()
.unwrap()
.1
.materialize()
.unwrap();
assert_eq!(*realised.data, tensor_data);
}
#[test]
fn test_lazy_byte_identical_to_eager_bridge() {
let tmp = tempfile::tempdir().unwrap();
let model_dir = tmp.path();
let f32_data: Vec<u8> = (0..6u32).flat_map(|v| (v as f32).to_le_bytes()).collect();
let f16_data: Vec<u8> = (0..4u32)
.flat_map(|v| half::f16::from_f32(v as f32).to_le_bytes())
.collect();
let bf16_data: Vec<u8> = (0..3u32)
.flat_map(|v| half::bf16::from_f32(v as f32).to_le_bytes())
.collect();
let safetensors_data = create_test_safetensors(&[
("zebra.weight", &[2, 3], "F32", &f32_data),
("alpha.weight", &[2, 2], "F16", &f16_data),
("mango.weight", &[3], "BF16", &bf16_data),
]);
std::fs::write(model_dir.join("model.safetensors"), &safetensors_data).unwrap();
let progress = ProgressReporter::new();
let lazy_map = read_tensors_lazy(model_dir, &progress).unwrap();
let materialised = lazy_map.materialize_all().unwrap();
let eager = read_tensors(model_dir, &progress).unwrap();
assert_eq!(materialised.len(), eager.len());
for name in ["zebra.weight", "alpha.weight", "mango.weight"] {
let m = materialised.get(name).unwrap();
let e = eager.get(name).unwrap();
assert_eq!(m.shape, e.shape, "{name} shape");
assert_eq!(m.dtype, e.dtype, "{name} dtype");
assert_eq!(*m.data, *e.data, "{name} bytes");
}
}
#[test]
fn test_lazy_reader_does_not_materialise_at_index_time() {
let tmp = tempfile::tempdir().unwrap();
let model_dir = tmp.path();
let big_data: Vec<u8> = (0..16 * 1024)
.flat_map(|v: u32| (v as f32).to_le_bytes())
.collect();
let safetensors_data =
create_test_safetensors(&[("big.weight", &[16 * 1024], "F32", &big_data)]);
std::fs::write(model_dir.join("model.safetensors"), &safetensors_data).unwrap();
let progress = ProgressReporter::new();
let probe = AtomicUsize::new(0);
let lazy_map = read_tensors_lazy(model_dir, &progress).unwrap();
for _ in 0..100 {
let lazy = lazy_map.get("big.weight").unwrap();
assert_eq!(lazy.shape(), &[16 * 1024]);
assert_eq!(lazy.dtype(), DType::F32);
probe.fetch_add(1, Ordering::SeqCst);
}
assert_eq!(probe.load(Ordering::SeqCst), 100);
let realised = lazy_map
.into_iter()
.find(|(k, _)| k == "big.weight")
.unwrap()
.1
.materialize()
.unwrap();
assert_eq!(*realised.data, big_data);
}
#[test]
fn test_lazy_reader_rejects_header_size_mismatch() {
let tmp = tempfile::tempdir().unwrap();
let model_dir = tmp.path();
let mut header_map = serde_json::Map::new();
let mut tensor_info = serde_json::Map::new();
tensor_info.insert(
"dtype".to_string(),
serde_json::Value::String("F32".to_string()),
);
tensor_info.insert(
"shape".to_string(),
serde_json::Value::Array(vec![
serde_json::Value::Number(2.into()),
serde_json::Value::Number(3.into()),
]),
);
tensor_info.insert(
"data_offsets".to_string(),
serde_json::Value::Array(vec![
serde_json::Value::Number(0u64.into()),
serde_json::Value::Number(12u64.into()),
]),
);
header_map.insert("liar".to_string(), serde_json::Value::Object(tensor_info));
let header_json = serde_json::to_string(&header_map).unwrap();
let header_size = header_json.len() as u64;
let mut file_data = Vec::new();
file_data.extend_from_slice(&header_size.to_le_bytes());
file_data.extend_from_slice(header_json.as_bytes());
file_data.extend_from_slice(&vec![0u8; 12]);
std::fs::write(model_dir.join("model.safetensors"), &file_data).unwrap();
let progress = ProgressReporter::new();
let err = read_tensors_lazy(model_dir, &progress).unwrap_err();
match err {
SafetensorsError::HeaderSizeMismatch {
tensor,
header_bytes,
derived_bytes,
..
} => {
assert_eq!(tensor, "liar");
assert_eq!(header_bytes, 12);
assert_eq!(derived_bytes, 24);
}
other => panic!("expected HeaderSizeMismatch, got {other:?}"),
}
}
}