use crate::checkpoint::{TensorCatalog, TensorDescriptor, TensorDtype, TensorStorage};
pub use eredu_checkpoint::artifact::ArtifactFile;
use eredu_checkpoint::{
artifact::{fingerprint_artifact_files, ArtifactFingerprintError, ArtifactMemberFingerprint},
safetensors::SafetensorsShards,
store::{SharedCheckpointSource, TensorMetadata},
StoredDtype,
};
use eredu_gguf::{Checkpoint as GgufCheckpoint, GgmlType, MetadataValue};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::{
collections::{BTreeMap, BTreeSet},
fs::File,
io::Read,
path::{Path, PathBuf},
sync::Arc,
};
#[derive(Clone, Copy, Eq, Hash, PartialEq)]
pub struct ArtifactIdentity([u8; 32]);
impl ArtifactIdentity {
pub const fn digest(self) -> [u8; 32] {
self.0
}
}
#[derive(Clone)]
pub struct ArtifactAdmissionToken(Arc<()>);
impl ArtifactAdmissionToken {
fn new() -> Self {
Self(Arc::new(()))
}
pub fn same_admission(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.0, &other.0)
}
}
impl std::fmt::Debug for ArtifactAdmissionToken {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("ArtifactAdmissionToken(..)")
}
}
impl std::fmt::Display for ArtifactIdentity {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(formatter, "sha256:{}", hex_digest(&self.0))
}
}
impl std::fmt::Debug for ArtifactIdentity {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(self, formatter)
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ArtifactMemberIdentity {
logical_role: String,
length: u64,
digest: [u8; 32],
}
impl ArtifactMemberIdentity {
pub fn new(logical_role: impl Into<String>, length: u64, digest: [u8; 32]) -> Self {
Self {
logical_role: logical_role.into(),
length,
digest,
}
}
pub fn logical_role(&self) -> &str {
&self.logical_role
}
pub const fn length(&self) -> u64 {
self.length
}
pub const fn digest(&self) -> [u8; 32] {
self.digest
}
}
impl From<ArtifactMemberFingerprint> for ArtifactMemberIdentity {
fn from(member: ArtifactMemberFingerprint) -> Self {
Self::new(member.logical_role(), member.length(), member.digest())
}
}
pub fn fingerprint_artifact(
domain: &str,
members: impl IntoIterator<Item = ArtifactMemberIdentity>,
) -> Result<ArtifactIdentity, ArtifactError> {
if domain.is_empty() {
return Err(ArtifactError::InvalidArtifactIdentity(
"artifact identity domain must not be empty".into(),
));
}
let mut members = members.into_iter().collect::<Vec<_>>();
if members.is_empty() {
return Err(ArtifactError::InvalidArtifactIdentity(
"artifact identity requires at least one file".into(),
));
}
if members.iter().any(|member| member.logical_role.is_empty()) {
return Err(ArtifactError::InvalidArtifactIdentity(
"artifact member has an empty logical role".into(),
));
}
members.sort_unstable_by(|left, right| left.logical_role.cmp(&right.logical_role));
if let Some(pair) = members
.windows(2)
.find(|pair| pair[0].logical_role == pair[1].logical_role)
{
return Err(ArtifactError::InvalidArtifactIdentity(format!(
"duplicate artifact logical role {:?}",
pair[0].logical_role
)));
}
let mut hasher = sha2::Sha256::new();
hash_identity_component(&mut hasher, b"eredu-checkpoint-artifact-v2");
hash_identity_component(&mut hasher, domain.as_bytes());
use sha2::Digest as _;
hasher.update((members.len() as u64).to_le_bytes());
for member in members {
hash_identity_component(&mut hasher, member.logical_role.as_bytes());
hasher.update(member.length.to_le_bytes());
hasher.update(member.digest);
}
Ok(ArtifactIdentity(hasher.finalize().into()))
}
pub fn fingerprint_filesystem_artifact(
domain: &str,
files: impl IntoIterator<Item = ArtifactFile>,
) -> Result<ArtifactIdentity, ArtifactError> {
fingerprint_artifact(
domain,
fingerprint_artifact_files(files)?
.into_iter()
.map(ArtifactMemberIdentity::from),
)
}
pub fn fingerprint_safetensors_artifact(
domain: &str,
shards: &SafetensorsShards,
) -> Result<ArtifactIdentity, ArtifactError> {
fingerprint_filesystem_artifact(
domain,
shards
.logical_payload_paths()
.iter()
.map(|(role, path)| ArtifactFile::new(role, path)),
)
}
pub fn fingerprint_gguf_artifact(
domain: &str,
checkpoint: &GgufCheckpoint,
) -> Result<ArtifactIdentity, ArtifactError> {
fingerprint_filesystem_artifact(
domain,
checkpoint
.shards()
.iter()
.map(|shard| ArtifactFile::new(format!("split/{:05}", shard.split_no()), shard.path())),
)
}
fn hash_identity_component(hasher: &mut sha2::Sha256, value: &[u8]) {
use sha2::Digest as _;
hasher.update((value.len() as u64).to_le_bytes());
hasher.update(value);
}
fn hex_digest(bytes: &[u8]) -> String {
const DIGITS: &[u8; 16] = b"0123456789abcdef";
let mut output = String::with_capacity(bytes.len() * 2);
for &byte in bytes {
output.push(DIGITS[usize::from(byte >> 4)] as char);
output.push(DIGITS[usize::from(byte & 0x0f)] as char);
}
output
}
#[derive(Debug, Clone, Copy, Eq, Hash, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum LoadingProtocol {
Model,
Realtime,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ArtifactFormat {
SafeTensors,
Gguf,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelConfiguration {
declared_model_type: String,
effective_model_type: String,
family: String,
loading_protocol: LoadingProtocol,
#[serde(skip_serializing_if = "Option::is_none")]
json: Option<Value>,
}
impl ModelConfiguration {
pub fn new(
declared_model_type: impl Into<String>,
effective_model_type: impl Into<String>,
family: impl Into<String>,
loading_protocol: LoadingProtocol,
json: Option<Value>,
) -> Result<Self, ArtifactError> {
let configuration = Self {
declared_model_type: declared_model_type.into(),
effective_model_type: effective_model_type.into(),
family: family.into(),
loading_protocol,
json,
};
if [
&configuration.declared_model_type,
&configuration.effective_model_type,
&configuration.family,
]
.into_iter()
.any(|value| value.trim().is_empty())
{
return Err(ArtifactError::InvalidArtifact(
"model configuration identities must be non-empty".into(),
));
}
Ok(configuration)
}
pub fn declared_model_type(&self) -> &str {
&self.declared_model_type
}
pub fn effective_model_type(&self) -> &str {
&self.effective_model_type
}
pub fn family(&self) -> &str {
&self.family
}
pub const fn loading_protocol(&self) -> LoadingProtocol {
self.loading_protocol
}
pub const fn json(&self) -> Option<&Value> {
self.json.as_ref()
}
}
#[derive(Debug, Clone)]
pub struct ResolvedModelConfiguration<P> {
configuration: ModelConfiguration,
architecture_plan: P,
}
impl<P> ResolvedModelConfiguration<P> {
pub fn new(configuration: ModelConfiguration, architecture_plan: P) -> Self {
Self {
configuration,
architecture_plan,
}
}
pub const fn configuration(&self) -> &ModelConfiguration {
&self.configuration
}
pub const fn architecture_plan(&self) -> &P {
&self.architecture_plan
}
pub fn into_parts(self) -> (ModelConfiguration, P) {
(self.configuration, self.architecture_plan)
}
}
pub trait ModelConfigurationResolver {
type ArtifactPlan: Clone + std::fmt::Debug;
fn resolve_safetensors(
&self,
json: &Value,
) -> Result<ResolvedModelConfiguration<Self::ArtifactPlan>, ArtifactError>;
fn resolve_gguf(
&self,
architecture: &str,
checkpoint: &GgufCheckpoint,
) -> Result<ResolvedModelConfiguration<Self::ArtifactPlan>, ArtifactError>;
fn gguf_companion_requirements(
&self,
architecture: &str,
checkpoint: &GgufCheckpoint,
) -> Result<Vec<GgufCompanionRequirement>, ArtifactError>;
fn artifact_plan(
&self,
_path: &Path,
_format: ArtifactFormat,
_configuration: &ModelConfiguration,
_tensors: &TensorCatalog,
_validated_gguf: Option<&ValidatedGguf>,
resolved_plan: Self::ArtifactPlan,
) -> Result<Self::ArtifactPlan, ArtifactError> {
Ok(resolved_plan)
}
}
#[derive(Debug, Clone, Eq, Hash, Ord, PartialEq, PartialOrd)]
#[non_exhaustive]
pub enum GgufCompanionRole {
MediaProjector,
Named(String),
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum GgufCompanionEncoding {
DenseRequired,
DensePreferred,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct GgufCompanionRequirement {
role: GgufCompanionRole,
required: bool,
filename_prefix: String,
parent_search_depth: usize,
encoding: GgufCompanionEncoding,
}
impl GgufCompanionRequirement {
pub fn new(
role: GgufCompanionRole,
required: bool,
filename_prefix: impl Into<String>,
parent_search_depth: usize,
encoding: GgufCompanionEncoding,
) -> Result<Self, ArtifactError> {
let filename_prefix = filename_prefix.into();
if filename_prefix.trim().is_empty()
|| matches!(&role, GgufCompanionRole::Named(name) if name.trim().is_empty())
{
return Err(ArtifactError::InvalidArtifact(
"GGUF companion roles and filename prefixes must be non-empty".into(),
));
}
Ok(Self {
role,
required,
filename_prefix,
parent_search_depth,
encoding,
})
}
pub fn role(&self) -> &GgufCompanionRole {
&self.role
}
}
pub fn gguf_u32_metadata_values(
key: &str,
value: Option<&MetadataValue>,
) -> Result<Vec<u32>, ArtifactError> {
let Some(value) = value else {
return Ok(Vec::new());
};
value.to_u32_vec().ok_or_else(|| {
ArtifactError::InvalidArtifact(format!(
"GGUF metadata key {key:?} must contain an integer or integer array whose values fit in u32"
))
})
}
#[derive(Debug, Clone)]
pub struct ArtifactInspection<P = ()> {
admission_token: ArtifactAdmissionToken,
path: PathBuf,
format: ArtifactFormat,
configuration: ModelConfiguration,
tensors: TensorCatalog,
safetensors_shards: Option<SafetensorsShards>,
validated_gguf: Option<ValidatedGguf>,
architecture_plan: P,
}
#[derive(Debug, Clone)]
pub struct ValidatedGguf {
checkpoint: GgufCheckpoint,
companions: BTreeMap<GgufCompanionRole, ValidatedGgufCompanion>,
}
#[derive(Debug, Clone)]
pub struct ValidatedGgufCompanion {
path: PathBuf,
checkpoint: GgufCheckpoint,
}
impl ValidatedGgufCompanion {
pub fn path(&self) -> &Path {
&self.path
}
pub fn checkpoint(&self) -> &GgufCheckpoint {
&self.checkpoint
}
}
impl ValidatedGguf {
pub fn checkpoint(&self) -> &GgufCheckpoint {
&self.checkpoint
}
pub fn companion(&self, role: &GgufCompanionRole) -> Option<&ValidatedGgufCompanion> {
self.companions.get(role)
}
pub fn companions(
&self,
) -> impl Iterator<Item = (&GgufCompanionRole, &ValidatedGgufCompanion)> {
self.companions.iter()
}
pub fn into_parts(
self,
) -> (
GgufCheckpoint,
BTreeMap<GgufCompanionRole, ValidatedGgufCompanion>,
) {
(self.checkpoint, self.companions)
}
}
impl<P> ArtifactInspection<P> {
pub fn admission_token(&self) -> ArtifactAdmissionToken {
self.admission_token.clone()
}
pub fn path(&self) -> &Path {
&self.path
}
pub const fn format(&self) -> ArtifactFormat {
self.format
}
pub fn configuration(&self) -> &ModelConfiguration {
&self.configuration
}
pub fn tensors(&self) -> &TensorCatalog {
&self.tensors
}
pub fn safetensors_shards(&self) -> Option<&SafetensorsShards> {
self.safetensors_shards.as_ref()
}
pub fn validated_gguf(&self) -> Option<&ValidatedGguf> {
self.validated_gguf.as_ref()
}
pub fn gguf_checkpoint(&self) -> Option<&GgufCheckpoint> {
self.validated_gguf().map(ValidatedGguf::checkpoint)
}
pub fn architecture_plan(&self) -> &P {
&self.architecture_plan
}
pub fn architecture_plan_mut(&mut self) -> &mut P {
&mut self.architecture_plan
}
pub fn map_architecture_plan<Q>(self, map: impl FnOnce(P) -> Q) -> ArtifactInspection<Q> {
ArtifactInspection {
admission_token: self.admission_token,
path: self.path,
format: self.format,
configuration: self.configuration,
tensors: self.tensors,
safetensors_shards: self.safetensors_shards,
validated_gguf: self.validated_gguf,
architecture_plan: map(self.architecture_plan),
}
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
#[non_exhaustive]
pub enum QuantizationRequest {
Affine {
group_size: u32,
bits: u8,
},
MxFp4,
}
#[derive(Debug, Clone, Copy, Default, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ResidencyRequest {
#[default]
FullyResident,
LayerwiseHost,
DenseDiskStream,
AddressableParameterBanks,
}
#[derive(Debug, Clone, Copy, Default, Eq, PartialEq, Serialize, Deserialize)]
pub struct PreparationPolicy {
quantization: Option<QuantizationRequest>,
residency: ResidencyRequest,
topology: Option<crate::topology::ParallelTopology>,
required_session_capabilities: crate::backend::SessionCapabilities,
}
impl PreparationPolicy {
pub const fn new(
quantization: Option<QuantizationRequest>,
residency: ResidencyRequest,
) -> Self {
Self {
quantization,
residency,
topology: None,
required_session_capabilities: crate::backend::SessionCapabilities::new(
false, false, false,
),
}
}
pub const fn quantization(self) -> Option<QuantizationRequest> {
self.quantization
}
pub const fn residency(self) -> ResidencyRequest {
self.residency
}
pub const fn topology(self) -> Option<crate::topology::ParallelTopology> {
self.topology
}
pub const fn required_session_capabilities(self) -> crate::backend::SessionCapabilities {
self.required_session_capabilities
}
pub const fn with_topology(mut self, topology: crate::topology::ParallelTopology) -> Self {
self.topology = Some(topology);
self
}
pub const fn with_required_session_capabilities(
mut self,
capabilities: crate::backend::SessionCapabilities,
) -> Self {
self.required_session_capabilities = capabilities;
self
}
pub fn validate_session_capabilities(
&self,
available: &crate::backend::SessionCapabilities,
) -> Result<(), crate::backend::SessionCapabilityError> {
self.required_session_capabilities.validate(available)
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum MaterializationRoute {
Resident,
Layerwise,
AddressableParameterBanks,
}
#[derive(Debug, Clone)]
pub struct ModelPreparationPlan<P = ()> {
inspection: ArtifactInspection<P>,
policy: PreparationPolicy,
route: MaterializationRoute,
admitted_session_capabilities: crate::backend::SessionCapabilities,
}
impl<P> ModelPreparationPlan<P> {
pub fn from_retained_admission(
inspection: ArtifactInspection<P>,
admission: crate::PreparationAdmission,
) -> Result<Self, ArtifactError> {
if admission.request().format() != inspection.format() {
return Err(ArtifactError::InvalidArtifact(
"retained preparation admission has a different artifact format".into(),
));
}
Ok(Self {
inspection,
policy: admission.request().policy(),
route: admission.route(),
admitted_session_capabilities: admission.session_capabilities(),
})
}
pub fn inspection(&self) -> &ArtifactInspection<P> {
&self.inspection
}
pub const fn policy(&self) -> PreparationPolicy {
self.policy
}
pub const fn route(&self) -> MaterializationRoute {
self.route
}
pub const fn admitted_session_capabilities(&self) -> crate::backend::SessionCapabilities {
self.admitted_session_capabilities
}
pub fn into_artifact(self) -> ModelArtifact {
match self.inspection.validated_gguf {
Some(validated) => ModelArtifact::Gguf {
path: self.inspection.path,
configuration: self.inspection.configuration,
tensors: self.inspection.tensors,
validated,
},
None => ModelArtifact::SafeTensors {
path: self.inspection.path,
configuration: self.inspection.configuration,
tensors: self.inspection.tensors,
shards: self
.inspection
.safetensors_shards
.expect("SafeTensors inspection retains its admitted shard set"),
},
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub enum ModelArtifact {
SafeTensors {
path: PathBuf,
configuration: ModelConfiguration,
tensors: TensorCatalog,
shards: SafetensorsShards,
},
Gguf {
path: PathBuf,
configuration: ModelConfiguration,
tensors: TensorCatalog,
validated: ValidatedGguf,
},
}
pub fn open_prepared_safetensors_artifact(
tensors: &TensorCatalog,
shards: SafetensorsShards,
resolution: eredu_checkpoint::validation::ResolvedCheckpointPlan,
max_cached_shards: usize,
) -> Result<SharedCheckpointSource, ArtifactError> {
let catalog = tensors
.descriptors()
.map(|tensor| {
let storage = tensor.storage.as_ref().ok_or_else(|| {
ArtifactError::InvalidArtifact(format!(
"prepared SafeTensors tensor {:?} has no storage provenance",
tensor.name
))
})?;
let stored_dtype = tensor_dtype_to_stored(&tensor.dtype);
Ok((
tensor.name.clone(),
TensorMetadata {
name: tensor.name.clone(),
logical_shape: tensor.shape.clone(),
physical_shape: tensor.shape.clone(),
stored_dtype,
encoded_byte_len: storage.length,
backing_shard: Some(PathBuf::from(&storage.member)),
},
))
})
.collect::<Result<BTreeMap<_, _>, ArtifactError>>()?;
eredu_checkpoint::store::open_prepared_safetensors_source(
shards,
catalog,
resolution,
max_cached_shards,
)
.map_err(Into::into)
}
fn tensor_dtype_to_stored(dtype: &TensorDtype) -> StoredDtype {
match dtype {
TensorDtype::Bool => StoredDtype::Bool,
TensorDtype::U8 => StoredDtype::U8,
TensorDtype::I8 => StoredDtype::I8,
TensorDtype::I16 => StoredDtype::I16,
TensorDtype::U16 => StoredDtype::U16,
TensorDtype::F16 => StoredDtype::F16,
TensorDtype::Bf16 => StoredDtype::BF16,
TensorDtype::I32 => StoredDtype::I32,
TensorDtype::U32 => StoredDtype::U32,
TensorDtype::F32 => StoredDtype::F32,
TensorDtype::F64 => StoredDtype::F64,
TensorDtype::I64 => StoredDtype::I64,
TensorDtype::U64 => StoredDtype::U64,
TensorDtype::Complex64 => StoredDtype::C64,
TensorDtype::Encoded(name) => match name.as_str() {
"F8_E4M3" => StoredDtype::F8E4M3,
"F4" => StoredDtype::F4,
"F8_E8M0" => StoredDtype::F8E8M0,
"F8_E5M2" => StoredDtype::F8E5M2,
_ => StoredDtype::Other(name.clone()),
},
}
}
pub fn inspect_artifact<R: ModelConfigurationResolver>(
path: impl AsRef<Path>,
resolver: &R,
) -> Result<ArtifactInspection<R::ArtifactPlan>, ArtifactError> {
let path = path.as_ref();
if is_gguf(path) {
inspect_gguf(path, resolver)
} else if path.is_dir() {
inspect_safetensors(path, resolver)
} else if !path.exists() {
Err(ArtifactError::MissingArtifact(path.to_path_buf()))
} else {
Err(ArtifactError::UnsupportedContainer(path.to_path_buf()))
}
}
pub fn plan_model_preparation<P>(
inspection: ArtifactInspection<P>,
policy: PreparationPolicy,
admitted_session_capabilities: crate::backend::SessionCapabilities,
) -> Result<ModelPreparationPlan<P>, ArtifactError> {
let route = validate_preparation_policy(inspection.configuration.loading_protocol, policy)?;
Ok(ModelPreparationPlan {
inspection,
policy,
route,
admitted_session_capabilities,
})
}
pub fn validate_preparation_policy(
protocol: LoadingProtocol,
policy: PreparationPolicy,
) -> Result<MaterializationRoute, ArtifactError> {
if protocol != LoadingProtocol::Model {
return Err(ArtifactError::UnsupportedLoadingProtocol(protocol));
}
let route = match policy.residency {
ResidencyRequest::FullyResident => MaterializationRoute::Resident,
ResidencyRequest::LayerwiseHost | ResidencyRequest::DenseDiskStream => {
MaterializationRoute::Layerwise
}
ResidencyRequest::AddressableParameterBanks => {
MaterializationRoute::AddressableParameterBanks
}
};
Ok(route)
}
fn inspect_gguf<R: ModelConfigurationResolver>(
path: &Path,
resolver: &R,
) -> Result<ArtifactInspection<R::ArtifactPlan>, ArtifactError> {
let checkpoint = GgufCheckpoint::open(path)?;
let architecture_name = checkpoint
.metadata()
.get("general.architecture")
.and_then(MetadataValue::as_str)
.ok_or(ArtifactError::MissingGgufArchitecture)?;
let (configuration, resolved_plan) = resolver
.resolve_gguf(architecture_name, &checkpoint)?
.into_parts();
let requirements = resolver.gguf_companion_requirements(architecture_name, &checkpoint)?;
let companions = resolve_gguf_companions(path, &requirements)?;
validate_gguf_container(&checkpoint)?;
let tensors = checkpoint
.tensors()
.map(|tensor| {
let descriptor = tensor.descriptor();
let shape = descriptor
.dimensions
.iter()
.map(|&dimension| {
usize::try_from(dimension).map_err(|_| {
ArtifactError::InvalidArtifact(format!(
"GGUF tensor {:?} dimension {dimension} exceeds the host address space",
descriptor.name
))
})
})
.collect::<Result<Vec<_>, _>>()?;
Ok(TensorDescriptor {
name: descriptor.name.clone(),
shape,
dtype: gguf_dtype(descriptor.ggml_type),
storage: None,
})
})
.collect::<Result<Vec<_>, ArtifactError>>()?;
let tensors = TensorCatalog::new(tensors)?;
let validated_gguf = ValidatedGguf {
checkpoint,
companions,
};
let architecture_plan = resolver.artifact_plan(
path,
ArtifactFormat::Gguf,
&configuration,
&tensors,
Some(&validated_gguf),
resolved_plan,
)?;
Ok(ArtifactInspection {
admission_token: ArtifactAdmissionToken::new(),
path: path.to_path_buf(),
format: ArtifactFormat::Gguf,
configuration,
tensors,
safetensors_shards: None,
validated_gguf: Some(validated_gguf),
architecture_plan,
})
}
pub fn resolve_gguf_companions(
primary: &Path,
requirements: &[GgufCompanionRequirement],
) -> Result<BTreeMap<GgufCompanionRole, ValidatedGgufCompanion>, ArtifactError> {
let mut resolved = BTreeMap::new();
let mut declared_roles = BTreeSet::new();
for requirement in requirements {
if !declared_roles.insert(requirement.role.clone()) {
return Err(ArtifactError::InvalidArtifact(format!(
"GGUF companion role {:?} was declared more than once",
requirement.role
)));
}
let mut directories = Vec::new();
let mut directory = primary.parent().unwrap_or_else(|| Path::new("."));
directories.push(directory.to_path_buf());
for _ in 0..requirement.parent_search_depth {
let Some(parent) = directory.parent() else {
break;
};
if parent == directory {
break;
}
directories.push(parent.to_path_buf());
directory = parent;
}
let mut candidates = Vec::new();
for directory in &directories {
let candidate_start = candidates.len();
for entry in std::fs::read_dir(directory)? {
let path = entry?.path();
let name = path
.file_name()
.and_then(|name| name.to_str())
.unwrap_or_default();
if path != primary
&& path.is_file()
&& name
.get(..requirement.filename_prefix.len())
.is_some_and(|prefix| {
prefix.eq_ignore_ascii_case(&requirement.filename_prefix)
})
&& is_gguf(&path)
{
let checkpoint = GgufCheckpoint::open(&path)?;
if checkpoint.physical_tensor_count() == 0 {
return Err(ArtifactError::InvalidArtifact(format!(
"GGUF companion {} contains no tensors",
path.display()
)));
}
let dense = checkpoint.tensors().all(|tensor| {
matches!(
tensor.descriptor().ggml_type,
eredu_gguf::GgmlType::F32
| eredu_gguf::GgmlType::F16
| eredu_gguf::GgmlType::Bf16
)
});
candidates.push((path, checkpoint, dense));
}
}
if candidates.len() != candidate_start {
break;
}
}
candidates.sort_by(|left, right| left.0.cmp(&right.0));
candidates.dedup_by(|left, right| left.0 == right.0);
let dense = candidates
.iter()
.filter(|candidate| candidate.2)
.collect::<Vec<_>>();
let selected = match requirement.encoding {
GgufCompanionEncoding::DenseRequired => match dense.as_slice() {
[candidate] => Some(*candidate),
[] if candidates.is_empty() => None,
[] => {
return Err(ArtifactError::InvalidArtifact(format!(
"GGUF companion {:?} requires dense F32, F16, or BF16 tensors, but all {} matching candidates are quantized",
requirement.role,
candidates.len()
)))
}
_ => return Err(ambiguous_companion(requirement, &directories, dense.len())),
},
GgufCompanionEncoding::DensePreferred => match dense.as_slice() {
[candidate] => Some(*candidate),
[] => match candidates.as_slice() {
[candidate] => Some(candidate),
[] => None,
_ => {
return Err(ambiguous_companion(
requirement,
&directories,
candidates.len(),
))
}
},
_ => return Err(ambiguous_companion(requirement, &directories, dense.len())),
},
};
match selected {
Some((path, checkpoint, _)) => {
resolved.insert(
requirement.role.clone(),
ValidatedGgufCompanion {
path: path.clone(),
checkpoint: checkpoint.clone(),
},
);
}
None if requirement.required => {
return Err(ArtifactError::MissingRequiredGgufCompanion {
role: requirement.role.clone(),
filename_prefix: requirement.filename_prefix.clone(),
searched_directories: directories,
})
}
None => {}
}
}
Ok(resolved)
}
fn ambiguous_companion(
requirement: &GgufCompanionRequirement,
directories: &[PathBuf],
candidates: usize,
) -> ArtifactError {
ArtifactError::InvalidArtifact(format!(
"GGUF companion {:?} is ambiguous: found {candidates} preferred candidates in {}",
requirement.role,
display_directories(directories)
))
}
fn display_directories(directories: &[PathBuf]) -> String {
directories
.iter()
.map(|directory| directory.display().to_string())
.collect::<Vec<_>>()
.join(", ")
}
fn validate_gguf_container(checkpoint: &GgufCheckpoint) -> Result<(), ArtifactError> {
if checkpoint.physical_tensor_count() == 0 {
return Err(ArtifactError::InvalidArtifact(
"GGUF model checkpoint contains no tensors".into(),
));
}
Ok(())
}
fn inspect_safetensors<R: ModelConfigurationResolver>(
path: &Path,
resolver: &R,
) -> Result<ArtifactInspection<R::ArtifactPlan>, ArtifactError> {
let config_path = path.join("config.json");
let json: Value = serde_json::from_reader(File::open(&config_path)?)?;
let (configuration, resolved_plan) = resolver.resolve_safetensors(&json)?.into_parts();
let shards = SafetensorsShards::discover(path)?;
let mut descriptors = Vec::new();
let mut names = BTreeSet::new();
for shard in shards.payload_paths() {
for descriptor in inspect_safetensors_header(shard)? {
if !names.insert(descriptor.name.clone()) {
return Err(ArtifactError::DuplicateTensor(descriptor.name));
}
descriptors.push(descriptor);
}
}
let tensors = TensorCatalog::new(descriptors)?;
if tensors.is_empty() {
return Err(ArtifactError::InvalidArtifact(
"SafeTensors checkpoint contains no tensors".into(),
));
}
let architecture_plan = resolver.artifact_plan(
path,
ArtifactFormat::SafeTensors,
&configuration,
&tensors,
None,
resolved_plan,
)?;
Ok(ArtifactInspection {
admission_token: ArtifactAdmissionToken::new(),
path: path.to_path_buf(),
format: ArtifactFormat::SafeTensors,
configuration,
tensors,
safetensors_shards: Some(shards),
validated_gguf: None,
architecture_plan,
})
}
#[derive(Deserialize)]
struct RawSafetensorInfo {
dtype: String,
shape: Vec<usize>,
data_offsets: [u64; 2],
}
fn inspect_safetensors_header(path: &Path) -> Result<Vec<TensorDescriptor>, ArtifactError> {
const MAX_HEADER_BYTES: u64 = 100_000_000;
let mut file = File::open(path)?;
let file_len = file.metadata()?.len();
let mut length = [0_u8; 8];
file.read_exact(&mut length)?;
let header_len = u64::from_le_bytes(length);
if header_len > MAX_HEADER_BYTES {
return Err(ArtifactError::InvalidArtifact(format!(
"SafeTensors header in {} exceeds {MAX_HEADER_BYTES} bytes",
path.display()
)));
}
let mut header = vec![
0_u8;
usize::try_from(header_len).map_err(|_| {
ArtifactError::InvalidArtifact("SafeTensors header length overflows usize".into())
})?
];
file.read_exact(&mut header)?;
let raw: BTreeMap<String, Value> = serde_json::from_slice(&header)?;
let payload_start = 8_u64
.checked_add(header_len)
.ok_or_else(|| ArtifactError::InvalidArtifact("SafeTensors offset overflow".into()))?;
let mut entries = raw
.into_iter()
.filter(|(name, _)| name != "__metadata__")
.map(|(name, value)| {
serde_json::from_value::<RawSafetensorInfo>(value).map(|info| (name, info))
})
.collect::<Result<Vec<_>, _>>()?;
entries.sort_by_key(|(_, info)| info.data_offsets[0]);
let mut output = Vec::with_capacity(entries.len());
let mut expected_offset = 0_u64;
for (name, info) in entries {
if info.shape.contains(&0) {
return Err(ArtifactError::InvalidArtifact(format!(
"SafeTensors tensor {name:?} has an invalid shape"
)));
}
let [start, end] = info.data_offsets;
if start != expected_offset || end < start {
return Err(ArtifactError::InvalidArtifact(format!(
"SafeTensors tensor {name:?} has non-contiguous data offsets"
)));
}
expected_offset = end;
let absolute = payload_start
.checked_add(start)
.ok_or_else(|| ArtifactError::InvalidArtifact("SafeTensors offset overflow".into()))?;
output.push(TensorDescriptor {
name,
shape: info.shape,
dtype: safetensors_dtype(&info.dtype),
storage: Some(TensorStorage {
member: path.display().to_string(),
offset: absolute,
length: end - start,
}),
});
}
if payload_start
.checked_add(expected_offset)
.ok_or_else(|| ArtifactError::InvalidArtifact("SafeTensors length overflow".into()))?
!= file_len
{
return Err(ArtifactError::InvalidArtifact(format!(
"SafeTensors payload length does not match header in {}",
path.display()
)));
}
Ok(output)
}
fn safetensors_dtype(dtype: &str) -> TensorDtype {
match dtype {
"BOOL" => TensorDtype::Bool,
"U8" => TensorDtype::U8,
"I8" => TensorDtype::I8,
"I16" => TensorDtype::I16,
"U16" => TensorDtype::U16,
"F16" => TensorDtype::F16,
"BF16" => TensorDtype::Bf16,
"I32" => TensorDtype::I32,
"U32" => TensorDtype::U32,
"F32" => TensorDtype::F32,
"F64" => TensorDtype::F64,
"I64" => TensorDtype::I64,
"U64" => TensorDtype::U64,
"C64" => TensorDtype::Complex64,
other => TensorDtype::Encoded(other.into()),
}
}
fn gguf_dtype(dtype: GgmlType) -> TensorDtype {
match dtype {
GgmlType::F32 => TensorDtype::F32,
GgmlType::F16 => TensorDtype::F16,
GgmlType::Bf16 => TensorDtype::Bf16,
GgmlType::I8 => TensorDtype::I8,
GgmlType::I16 => TensorDtype::I16,
GgmlType::I32 => TensorDtype::I32,
GgmlType::I64 => TensorDtype::I64,
GgmlType::F64 => TensorDtype::F64,
encoded => TensorDtype::Encoded(format!("{encoded:?}")),
}
}
fn is_gguf(path: &Path) -> bool {
path.extension()
.and_then(|extension| extension.to_str())
.is_some_and(|extension| extension.eq_ignore_ascii_case("gguf"))
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ArtifactError {
#[error("invalid artifact identity: {0}")]
InvalidArtifactIdentity(String),
#[error("model artifact does not exist: {0}")]
MissingArtifact(PathBuf),
#[error("model artifact must be a SafeTensors directory or .gguf file: {0}")]
UnsupportedContainer(PathBuf),
#[error("unsupported model type: {0}")]
UnsupportedModelType(String),
#[error("unsupported GGUF architecture: {0}")]
UnsupportedGgufArchitecture(String),
#[error("GGUF metadata is missing string key \"general.architecture\"")]
MissingGgufArchitecture,
#[error(
"required GGUF companion {role:?} matching {filename_prefix:?} was not found in {searched}",
searched = display_directories(.searched_directories)
)]
MissingRequiredGgufCompanion {
role: GgufCompanionRole,
filename_prefix: String,
searched_directories: Vec<PathBuf>,
},
#[error("invalid model artifact: {0}")]
InvalidArtifact(String),
#[error("invalid architecture artifact plan: {0}")]
InvalidArchitecturePlan(String),
#[error("duplicate checkpoint tensor {0:?}")]
DuplicateTensor(String),
#[error(transparent)]
SafetensorsShards(#[from] eredu_checkpoint::safetensors::SafetensorsShardError),
#[error(transparent)]
ArtifactFingerprint(#[from] ArtifactFingerprintError),
#[error(transparent)]
CheckpointStore(#[from] eredu_checkpoint::store::StoreError),
#[error("unsupported model quantization policy: {0}")]
UnsupportedQuantizationPolicy(String),
#[error("unsupported model residency policy: {0}")]
UnsupportedResidencyPolicy(String),
#[error("model artifact requires the {0:?} loading protocol")]
UnsupportedLoadingProtocol(LoadingProtocol),
#[error(transparent)]
Io(#[from] std::io::Error),
#[error(transparent)]
Json(#[from] serde_json::Error),
#[error(transparent)]
Gguf(#[from] eredu_gguf::Error),
#[error(transparent)]
Catalog(#[from] crate::checkpoint::CatalogError),
}
#[cfg(test)]
mod tests {
use super::*;
use eredu_gguf::{GgmlType, MetadataArray, TensorInput, Writer};
use std::io::Write;
#[test]
fn artifact_identity_rejects_empty_and_duplicate_layouts() {
assert!(matches!(
fingerprint_artifact("", [ArtifactMemberIdentity::new("weights", 7, [1; 32])]),
Err(ArtifactError::InvalidArtifactIdentity(_))
));
assert!(matches!(
fingerprint_artifact("model", std::iter::empty::<ArtifactMemberIdentity>()),
Err(ArtifactError::InvalidArtifactIdentity(_))
));
assert!(matches!(
fingerprint_artifact("model", [ArtifactMemberIdentity::new("", 7, [1; 32])]),
Err(ArtifactError::InvalidArtifactIdentity(_))
));
assert!(matches!(
fingerprint_artifact(
"model",
[
ArtifactMemberIdentity::new("weights", 7, [1; 32]),
ArtifactMemberIdentity::new("weights", 7, [1; 32]),
],
),
Err(ArtifactError::InvalidArtifactIdentity(_))
));
}
#[test]
fn artifact_identity_is_order_and_location_independent_but_domain_and_layout_exact() {
let first = tempfile::tempdir().unwrap();
let relocated = tempfile::tempdir().unwrap();
for directory in [first.path(), relocated.path()] {
std::fs::write(directory.join("a"), b"first").unwrap();
std::fs::write(directory.join("b"), b"second").unwrap();
}
let fingerprint = |root: &Path, domain: &str, reversed: bool| {
let mut files = vec![
ArtifactFile::new("decoder", root.join("a")),
ArtifactFile::new("projector", root.join("b")),
];
if reversed {
files.reverse();
}
fingerprint_filesystem_artifact(domain, files).unwrap()
};
let identity = fingerprint(first.path(), "model", false);
assert_eq!(identity, fingerprint(first.path(), "model", true));
assert_eq!(identity, fingerprint(relocated.path(), "model", false));
assert_ne!(identity, fingerprint(first.path(), "assistant", false));
assert_ne!(
identity,
fingerprint_filesystem_artifact(
"model",
[
ArtifactFile::new("decoder-renamed", first.path().join("a")),
ArtifactFile::new("projector", first.path().join("b")),
],
)
.unwrap()
);
assert_eq!(format!("{identity:?}"), identity.to_string());
assert_eq!(identity.to_string().len(), 71);
}
#[test]
fn artifact_identity_includes_exact_content_and_length() {
let directory = tempfile::tempdir().unwrap();
let path = directory.path().join("weights");
let identify = || {
fingerprint_filesystem_artifact("model", [ArtifactFile::new("weights", &path)]).unwrap()
};
std::fs::write(&path, b"a").unwrap();
let short = identify();
std::fs::write(&path, b"b").unwrap();
assert_ne!(short, identify());
std::fs::write(&path, b"a\0").unwrap();
assert_ne!(short, identify());
}
#[test]
fn admitted_safetensors_identity_is_relocation_independent_and_content_exact() {
let first = tempfile::tempdir().unwrap();
let relocated = tempfile::tempdir().unwrap();
write_safetensors_fixture(first.path(), "llama");
write_safetensors_fixture(relocated.path(), "llama");
let first_path = first.path().join("original-name.safetensors");
let relocated_path = relocated.path().join("renamed.safetensors");
std::fs::rename(first.path().join("model.safetensors"), &first_path).unwrap();
std::fs::rename(relocated.path().join("model.safetensors"), &relocated_path).unwrap();
let first_shards = SafetensorsShards::discover(&first_path).unwrap();
let relocated_shards = SafetensorsShards::discover(&relocated_path).unwrap();
let identify = |shards: &SafetensorsShards| {
fingerprint_safetensors_artifact("speculative-test", shards).unwrap()
};
let original = identify(&first_shards);
assert_eq!(original, identify(&relocated_shards));
let payload = relocated_path;
let mut bytes = std::fs::read(&payload).unwrap();
let last = bytes.last_mut().unwrap();
*last ^= 0x01;
std::fs::write(&payload, bytes).unwrap();
assert_ne!(original, identify(&relocated_shards));
}
#[test]
fn admitted_gguf_identity_uses_split_roles_and_exact_content() {
let first = tempfile::tempdir().unwrap();
let relocated = tempfile::tempdir().unwrap();
let write = |path: &Path| {
let data = 1.0_f32.to_le_bytes();
Writer::default()
.write(
File::create(path).unwrap(),
&BTreeMap::from([(
"general.architecture".into(),
MetadataValue::String("llama".into()),
)]),
&[TensorInput {
name: "weight",
dimensions: &[1],
ggml_type: GgmlType::F32,
data: &data,
}],
)
.unwrap();
};
let first_path = first.path().join("first-name.gguf");
let relocated_path = relocated.path().join("different-name.gguf");
write(&first_path);
write(&relocated_path);
let first_checkpoint = GgufCheckpoint::open(&first_path).unwrap();
let relocated_checkpoint = GgufCheckpoint::open(&relocated_path).unwrap();
let original = fingerprint_gguf_artifact("speculative-test", &first_checkpoint).unwrap();
assert_eq!(
original,
fingerprint_gguf_artifact("speculative-test", &relocated_checkpoint).unwrap()
);
let mut bytes = std::fs::read(&relocated_path).unwrap();
*bytes.last_mut().unwrap() ^= 0x01;
std::fs::write(&relocated_path, bytes).unwrap();
assert_ne!(
original,
fingerprint_gguf_artifact("speculative-test", &relocated_checkpoint).unwrap()
);
}
#[test]
fn canonical_prepared_safetensors_open_rejects_substitution_without_payload_reads() {
use eredu_checkpoint::{
schema::{
CatalogPolicy, SafetensorsCheckpointPlan, SafetensorsTensorConstraint,
StoredDtypeConstraint,
},
store::SafetensorsWeightStore,
};
let directory = tempfile::tempdir().unwrap();
write_safetensors_fixture(directory.path(), "llama");
let inspection = inspect_artifact(directory.path(), &FixtureResolver).unwrap();
let shards = inspection.safetensors_shards().unwrap().clone();
let plan = SafetensorsCheckpointPlan::new(
"fixture",
vec![SafetensorsTensorConstraint::required(
"token_embd.weight",
vec![2, 2],
StoredDtypeConstraint::Exact(StoredDtype::F32),
)],
Vec::new(),
CatalogPolicy::strict(),
)
.unwrap();
let store = SafetensorsWeightStore::open_admitted(shards.clone(), 1).unwrap();
let resolution =
eredu_checkpoint::validation::resolve_safetensors_plan(&store, &plan).unwrap();
let prepared = open_prepared_safetensors_artifact(
inspection.tensors(),
shards.clone(),
resolution.clone(),
1,
)
.unwrap();
assert_eq!(prepared.source_diagnostics().unwrap().physical_reads, 0);
let header = br#"{"token_embd.weight":{"dtype":"F32","shape":[4],"data_offsets":[0,16]}}"#;
let mut file = File::create(directory.path().join("model.safetensors")).unwrap();
file.write_all(&(header.len() as u64).to_le_bytes())
.unwrap();
file.write_all(header).unwrap();
file.write_all(&[0_u8; 16]).unwrap();
drop(file);
assert!(matches!(
open_prepared_safetensors_artifact(inspection.tensors(), shards, resolution, 1,),
Err(ArtifactError::CheckpointStore(
eredu_checkpoint::store::StoreError::PreparedCatalogMismatch { .. }
))
));
}
struct FixtureResolver;
#[derive(Debug, Clone, Default, Eq, PartialEq)]
struct FixtureArtifactPlan {
format: Option<ArtifactFormat>,
}
impl ModelConfigurationResolver for FixtureResolver {
type ArtifactPlan = FixtureArtifactPlan;
fn resolve_safetensors(
&self,
json: &Value,
) -> Result<ResolvedModelConfiguration<Self::ArtifactPlan>, ArtifactError> {
let model_type = json
.get("model_type")
.and_then(Value::as_str)
.ok_or_else(|| ArtifactError::InvalidArtifact("missing model_type".into()))?;
let family = match model_type {
"llama" => "llama",
"gemma4" => "gemma4",
"future" => "future_family",
other => return Err(ArtifactError::UnsupportedModelType(other.into())),
};
Ok(ResolvedModelConfiguration::new(
ModelConfiguration {
declared_model_type: model_type.into(),
effective_model_type: model_type.into(),
family: family.into(),
loading_protocol: LoadingProtocol::Model,
json: Some(json.clone()),
},
FixtureArtifactPlan::default(),
))
}
fn resolve_gguf(
&self,
architecture: &str,
_checkpoint: &GgufCheckpoint,
) -> Result<ResolvedModelConfiguration<Self::ArtifactPlan>, ArtifactError> {
let family = match architecture {
"llama" => "llama",
"future" => "future_family",
other => return Err(ArtifactError::UnsupportedGgufArchitecture(other.into())),
};
Ok(ResolvedModelConfiguration::new(
ModelConfiguration {
declared_model_type: architecture.into(),
effective_model_type: architecture.into(),
family: family.into(),
loading_protocol: LoadingProtocol::Model,
json: None,
},
FixtureArtifactPlan::default(),
))
}
fn gguf_companion_requirements(
&self,
architecture: &str,
_checkpoint: &GgufCheckpoint,
) -> Result<Vec<GgufCompanionRequirement>, ArtifactError> {
if architecture == "future" {
return Ok(vec![GgufCompanionRequirement::new(
GgufCompanionRole::MediaProjector,
false,
"mmproj",
0,
GgufCompanionEncoding::DensePreferred,
)?]);
}
Ok(Vec::new())
}
fn artifact_plan(
&self,
_path: &Path,
format: ArtifactFormat,
_configuration: &ModelConfiguration,
_tensors: &TensorCatalog,
_validated_gguf: Option<&ValidatedGguf>,
_resolved_plan: Self::ArtifactPlan,
) -> Result<Self::ArtifactPlan, ArtifactError> {
Ok(FixtureArtifactPlan {
format: Some(format),
})
}
}
fn write_safetensors_fixture(root: &Path, model_type: &str) {
std::fs::write(
root.join("config.json"),
format!(r#"{{"model_type":"{model_type}"}}"#),
)
.unwrap();
let header =
br#"{"token_embd.weight":{"dtype":"F32","shape":[2,2],"data_offsets":[0,16]}}"#;
let mut file = File::create(root.join("model.safetensors")).unwrap();
file.write_all(&(header.len() as u64).to_le_bytes())
.unwrap();
file.write_all(header).unwrap();
file.write_all(&[0_u8; 16]).unwrap();
}
fn write_gguf_fixture(path: &Path, ggml_type: GgmlType) {
let metadata = BTreeMap::from([(
"general.architecture".into(),
MetadataValue::String("clip".into()),
)]);
let (dimensions, data) = match ggml_type {
GgmlType::F32 => (vec![1], 1.0_f32.to_le_bytes().to_vec()),
GgmlType::Q8_0 => (vec![32], vec![0_u8; 34]),
other => panic!("unsupported fixture encoding {other:?}"),
};
Writer::default()
.write(
File::create(path).unwrap(),
&metadata,
&[TensorInput {
name: "projector.weight",
dimensions: &dimensions,
ggml_type,
data: &data,
}],
)
.unwrap();
}
#[test]
fn companion_planning_selects_by_catalog_encoding_not_filename() {
let root = tempfile::tempdir().unwrap();
let primary = root.path().join("model.gguf");
write_gguf_fixture(&primary, GgmlType::F32);
let quantized_name = root.path().join("mmproj-f16.gguf");
let dense_name = root.path().join("mmproj-q4_k.gguf");
write_gguf_fixture(&quantized_name, GgmlType::Q8_0);
write_gguf_fixture(&dense_name, GgmlType::F32);
let requirement = GgufCompanionRequirement::new(
GgufCompanionRole::MediaProjector,
true,
"mmproj",
1,
GgufCompanionEncoding::DensePreferred,
)
.unwrap();
let companions = resolve_gguf_companions(&primary, &[requirement]).unwrap();
assert_eq!(
companions
.get(&GgufCompanionRole::MediaProjector)
.unwrap()
.path(),
dense_name
);
}
#[test]
fn dense_only_and_required_companion_policies_fail_closed() {
let root = tempfile::tempdir().unwrap();
let primary = root.path().join("model.gguf");
write_gguf_fixture(&primary, GgmlType::F32);
write_gguf_fixture(&root.path().join("mmproj.gguf"), GgmlType::Q8_0);
let optional = GgufCompanionRequirement::new(
GgufCompanionRole::MediaProjector,
false,
"mmproj",
0,
GgufCompanionEncoding::DenseRequired,
)
.unwrap();
assert!(resolve_gguf_companions(&primary, &[optional]).is_err());
let required = GgufCompanionRequirement::new(
GgufCompanionRole::MediaProjector,
true,
"mmproj",
0,
GgufCompanionEncoding::DenseRequired,
)
.unwrap();
assert!(resolve_gguf_companions(&primary, &[required]).is_err());
}
#[test]
fn missing_required_companion_preserves_its_semantic_role() {
let root = tempfile::tempdir().unwrap();
let primary = root.path().join("model.gguf");
write_gguf_fixture(&primary, GgmlType::F32);
let requirement = GgufCompanionRequirement::new(
GgufCompanionRole::MediaProjector,
true,
"mmproj",
0,
GgufCompanionEncoding::DensePreferred,
)
.unwrap();
let error = resolve_gguf_companions(&primary, &[requirement]).unwrap_err();
assert!(matches!(
error,
ArtifactError::MissingRequiredGgufCompanion {
role: GgufCompanionRole::MediaProjector,
filename_prefix,
searched_directories,
} if filename_prefix == "mmproj" && searched_directories == [root.path()]
));
}
#[test]
fn loading_protocol_is_family_agnostic() {
assert!(matches!(
validate_preparation_policy(LoadingProtocol::Realtime, PreparationPolicy::default()),
Err(ArtifactError::UnsupportedLoadingProtocol(
LoadingProtocol::Realtime
))
));
}
#[test]
fn gguf_u32_metadata_is_lossless_and_fail_closed() {
let values = MetadataValue::Array(MetadataArray::Uint64(vec![0, u32::MAX.into()]));
assert_eq!(
gguf_u32_metadata_values("tokenizer.ids", Some(&values)).unwrap(),
vec![0, u32::MAX]
);
assert!(gguf_u32_metadata_values(
"tokenizer.ids",
Some(&MetadataValue::Uint64(u64::from(u32::MAX) + 1))
)
.is_err());
assert!(
gguf_u32_metadata_values("tokenizer.ids", Some(&MetadataValue::Int32(-1))).is_err()
);
assert!(gguf_u32_metadata_values(
"tokenizer.ids",
Some(&MetadataValue::String("1".into()))
)
.is_err());
assert!(gguf_u32_metadata_values("tokenizer.ids", None)
.unwrap()
.is_empty());
}
#[test]
fn safetensors_inspection_and_planning_are_backend_neutral() {
let root = tempfile::tempdir().unwrap();
write_safetensors_fixture(root.path(), "llama");
let inspection = inspect_artifact(root.path(), &FixtureResolver).unwrap();
assert_eq!(inspection.configuration().family, "llama");
assert_eq!(inspection.tensors().len(), 1);
assert_eq!(
inspection
.safetensors_shards()
.unwrap()
.payload_paths()
.len(),
1
);
let plan = plan_model_preparation(
inspection,
PreparationPolicy::default(),
crate::backend::SessionCapabilities::default(),
)
.unwrap();
assert_eq!(plan.route(), MaterializationRoute::Resident);
let architecture_plan = plan.inspection().architecture_plan().clone();
let artifact = plan.into_artifact();
let ModelArtifact::SafeTensors { shards, .. } = artifact else {
panic!("expected SafeTensors artifact");
};
assert_eq!(shards.payload_paths().len(), 1);
assert_eq!(architecture_plan.format, Some(ArtifactFormat::SafeTensors));
}
#[test]
fn safetensors_inspection_rejects_index_entries_missing_from_their_shard() {
let root = tempfile::tempdir().unwrap();
write_safetensors_fixture(root.path(), "llama");
std::fs::rename(
root.path().join("model.safetensors"),
root.path().join("model-00001.safetensors"),
)
.unwrap();
std::fs::write(
root.path().join("model.safetensors.index.json"),
r#"{"weight_map":{"missing.weight":"model-00001.safetensors"}}"#,
)
.unwrap();
assert!(matches!(
inspect_artifact(root.path(), &FixtureResolver),
Err(ArtifactError::SafetensorsShards(
eredu_checkpoint::safetensors::SafetensorsShardError::MalformedIndex { .. }
))
));
}
#[cfg(unix)]
#[test]
fn safetensors_inspection_rejects_indexed_symlinks_outside_the_access_root() {
use std::os::unix::fs::symlink;
let parent = tempfile::tempdir().unwrap();
let outside = parent.path().join("outside");
std::fs::create_dir(&outside).unwrap();
write_safetensors_fixture(&outside, "llama");
let checkpoint = parent.path().join("checkpoint");
std::fs::create_dir(&checkpoint).unwrap();
std::fs::write(checkpoint.join("config.json"), r#"{"model_type":"llama"}"#).unwrap();
symlink(
outside.join("model.safetensors"),
checkpoint.join("model-00001.safetensors"),
)
.unwrap();
std::fs::write(
checkpoint.join("model.safetensors.index.json"),
r#"{"weight_map":{"token_embd.weight":"model-00001.safetensors"}}"#,
)
.unwrap();
assert!(matches!(
inspect_artifact(&checkpoint, &FixtureResolver),
Err(ArtifactError::SafetensorsShards(
eredu_checkpoint::safetensors::SafetensorsShardError::UnsafeShardPath { .. }
))
));
}
#[test]
fn safetensors_native_dtypes_remain_typed_in_the_portable_catalog() {
assert_eq!(safetensors_dtype("BOOL"), TensorDtype::Bool);
assert_eq!(safetensors_dtype("I64"), TensorDtype::I64);
assert_eq!(safetensors_dtype("U32"), TensorDtype::U32);
assert_eq!(safetensors_dtype("F64"), TensorDtype::F64);
assert_eq!(safetensors_dtype("C64"), TensorDtype::Complex64);
assert_eq!(
safetensors_dtype("F8_E4M3"),
TensorDtype::Encoded("F8_E4M3".into())
);
}
#[test]
fn gguf_dense_dtypes_remain_typed_in_the_portable_catalog() {
assert_eq!(gguf_dtype(GgmlType::F16), TensorDtype::F16);
assert_eq!(gguf_dtype(GgmlType::Bf16), TensorDtype::Bf16);
assert_eq!(gguf_dtype(GgmlType::F32), TensorDtype::F32);
assert_eq!(
gguf_dtype(GgmlType::Q4K),
TensorDtype::Encoded("Q4K".into())
);
}
#[test]
fn core_accepts_families_defined_only_by_the_resolver() {
let root = tempfile::tempdir().unwrap();
write_safetensors_fixture(root.path(), "future");
let inspection = inspect_artifact(root.path(), &FixtureResolver).unwrap();
assert_eq!(inspection.configuration().family, "future_family");
assert_eq!(
inspection.configuration().loading_protocol,
LoadingProtocol::Model
);
assert!(plan_model_preparation(
inspection,
PreparationPolicy::default(),
crate::backend::SessionCapabilities::default(),
)
.is_ok());
}
#[test]
fn safetensors_inspection_accepts_rank_zero_scalar_parameters() {
let root = tempfile::tempdir().unwrap();
std::fs::write(
root.path().join("config.json"),
r#"{"model_type":"gemma4"}"#,
)
.unwrap();
let header = br#"{"clip.output_max":{"dtype":"F32","shape":[],"data_offsets":[0,4]}}"#;
let mut file = File::create(root.path().join("model.safetensors")).unwrap();
file.write_all(&(header.len() as u64).to_le_bytes())
.unwrap();
file.write_all(header).unwrap();
file.write_all(&0.0_f32.to_le_bytes()).unwrap();
let inspection = inspect_artifact(root.path(), &FixtureResolver).unwrap();
assert_eq!(
inspection.tensors().get("clip.output_max").unwrap().shape,
Vec::<usize>::new()
);
}
#[test]
fn parallel_policy_binds_the_exact_neutral_topology() {
let root = tempfile::tempdir().unwrap();
write_safetensors_fixture(root.path(), "llama");
let topology = crate::topology::ParallelTopology::new(2, 3, 4, 1).unwrap();
let policy = PreparationPolicy {
topology: Some(topology),
..PreparationPolicy::default()
};
let plan = plan_model_preparation(
inspect_artifact(root.path(), &FixtureResolver).unwrap(),
policy,
crate::backend::SessionCapabilities::default(),
)
.unwrap();
assert_eq!(plan.policy(), policy);
assert_eq!(plan.policy().topology, Some(topology));
assert_eq!(plan.route(), MaterializationRoute::Resident);
}
#[test]
fn policy_leaves_expert_cache_capability_to_architecture_and_backend() {
let root = tempfile::tempdir().unwrap();
write_safetensors_fixture(root.path(), "llama");
let plan = plan_model_preparation(
inspect_artifact(root.path(), &FixtureResolver).unwrap(),
PreparationPolicy {
residency: ResidencyRequest::AddressableParameterBanks,
..PreparationPolicy::default()
},
crate::backend::SessionCapabilities::default(),
)
.unwrap();
assert_eq!(
plan.route(),
MaterializationRoute::AddressableParameterBanks
);
}
#[test]
fn policy_leaves_nonresident_quantization_capability_to_architecture_and_backend() {
let root = tempfile::tempdir().unwrap();
write_safetensors_fixture(root.path(), "llama");
let policy = PreparationPolicy {
quantization: Some(QuantizationRequest::MxFp4),
residency: ResidencyRequest::LayerwiseHost,
..PreparationPolicy::default()
};
let plan = plan_model_preparation(
inspect_artifact(root.path(), &FixtureResolver).unwrap(),
policy,
crate::backend::SessionCapabilities::default(),
)
.unwrap();
assert_eq!(plan.policy(), policy);
assert_eq!(plan.route(), MaterializationRoute::Layerwise);
}
#[test]
fn gguf_plan_owns_the_portable_checkpoint_for_later_materialization() {
let root = tempfile::tempdir().unwrap();
let path = root.path().join("model.gguf");
let data = [0_u8; 2];
let metadata = BTreeMap::from([
(
"general.architecture".into(),
MetadataValue::String("llama".into()),
),
("llama.block_count".into(), MetadataValue::Uint32(1)),
("llama.embedding_length".into(), MetadataValue::Uint32(1)),
]);
Writer::default()
.write(
File::create(&path).unwrap(),
&metadata,
&[TensorInput {
name: "token_embd.weight",
dimensions: &[1],
ggml_type: GgmlType::F16,
data: &data,
}],
)
.unwrap();
let inspection = inspect_artifact(&path, &FixtureResolver).unwrap();
let validated = inspection.validated_gguf().unwrap();
assert_eq!(validated.checkpoint().physical_tensor_count(), 1);
assert_eq!(inspection.configuration().declared_model_type, "llama");
assert_eq!(
inspection.tensors().get("token_embd.weight").unwrap().dtype,
TensorDtype::F16
);
let plan = plan_model_preparation(
inspection,
PreparationPolicy::default(),
crate::backend::SessionCapabilities::default(),
)
.unwrap();
let architecture_plan = plan.inspection().architecture_plan().clone();
let artifact = plan.into_artifact();
let ModelArtifact::Gguf {
configuration,
validated,
..
} = artifact
else {
panic!("expected GGUF artifact");
};
assert_eq!(architecture_plan.format, Some(ArtifactFormat::Gguf));
assert_eq!(configuration.family, "llama");
assert_eq!(validated.checkpoint().physical_tensor_count(), 1);
}
#[test]
fn core_accepts_architecture_owned_gguf_schema() {
let root = tempfile::tempdir().unwrap();
let path = root.path().join("model.gguf");
let data = 1.0_f32.to_le_bytes();
let metadata = BTreeMap::from([
(
"general.architecture".into(),
MetadataValue::String("future".into()),
),
("future.state_width".into(), MetadataValue::Uint32(1)),
]);
Writer::default()
.write(
File::create(&path).unwrap(),
&metadata,
&[TensorInput {
name: "state.in_proj",
dimensions: &[1],
ggml_type: GgmlType::F32,
data: &data,
}],
)
.unwrap();
let inspection = inspect_artifact(&path, &FixtureResolver).unwrap();
assert_eq!(inspection.configuration().family, "future_family");
assert!(inspection.tensors().get("state.in_proj").is_some());
}
#[test]
fn preparation_plan_carries_the_exact_inspected_companion() {
let root = tempfile::tempdir().unwrap();
let primary = root.path().join("model.gguf");
let scalar = 1.0_f32.to_le_bytes();
Writer::default()
.write(
File::create(&primary).unwrap(),
&BTreeMap::from([(
"general.architecture".into(),
MetadataValue::String("future".into()),
)]),
&[TensorInput {
name: "state.in_proj",
dimensions: &[1],
ggml_type: GgmlType::F32,
data: &scalar,
}],
)
.unwrap();
let projector = root.path().join("mmproj.gguf");
write_gguf_fixture(&projector, GgmlType::F32);
let inspection = inspect_artifact(&primary, &FixtureResolver).unwrap();
assert_eq!(
inspection
.validated_gguf()
.unwrap()
.companion(&GgufCompanionRole::MediaProjector)
.unwrap()
.path(),
projector
);
let ModelArtifact::Gguf { validated, .. } = plan_model_preparation(
inspection,
PreparationPolicy::default(),
crate::backend::SessionCapabilities::default(),
)
.unwrap()
.into_artifact() else {
panic!("expected GGUF preparation artifact");
};
assert_eq!(
validated
.companion(&GgufCompanionRole::MediaProjector)
.unwrap()
.path(),
projector
);
}
}