Skip to main content

model_artifact/
lib.rs

1pub mod gguf;
2
3use std::path::Path;
4
5use anyhow::{Result, bail};
6use async_trait::async_trait;
7use model_ref::{
8    ModelRef, format_canonical_ref, gguf_matches_quant_selector, normalize_gguf_distribution_id,
9    parse_model_ref, split_gguf_shard_info,
10};
11use serde::{Deserialize, Serialize};
12
13#[async_trait]
14pub trait ModelRepository: Send + Sync {
15    async fn resolve_revision(&self, repo: &str, revision: Option<&str>) -> Result<String>;
16
17    async fn list_files(&self, repo: &str, revision: &str) -> Result<Vec<ModelArtifactFile>>;
18}
19
20#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
21pub struct ResolvedModelArtifact {
22    pub model_id: String,
23    pub source_repo: String,
24    pub source_revision: String,
25    pub selector: Option<String>,
26    pub format: ModelFormat,
27    pub files: Vec<ModelArtifactFile>,
28    pub primary_file: String,
29    pub canonical_ref: String,
30    pub distribution_id: String,
31}
32
33#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
34pub struct ModelIdentity {
35    pub model_id: String,
36    #[serde(skip_serializing_if = "Option::is_none")]
37    pub source_repo: Option<String>,
38    #[serde(skip_serializing_if = "Option::is_none")]
39    pub source_revision: Option<String>,
40    #[serde(skip_serializing_if = "Option::is_none")]
41    pub source_file: Option<String>,
42    #[serde(skip_serializing_if = "Option::is_none")]
43    pub canonical_ref: Option<String>,
44    #[serde(skip_serializing_if = "Option::is_none")]
45    pub distribution_id: Option<String>,
46    #[serde(skip_serializing_if = "Option::is_none")]
47    pub selector: Option<String>,
48}
49
50impl ModelIdentity {
51    pub fn from_model_id(model_id: impl Into<String>) -> Self {
52        Self {
53            model_id: model_id.into(),
54            source_repo: None,
55            source_revision: None,
56            source_file: None,
57            canonical_ref: None,
58            distribution_id: None,
59            selector: None,
60        }
61    }
62}
63
64impl From<&ResolvedModelArtifact> for ModelIdentity {
65    fn from(artifact: &ResolvedModelArtifact) -> Self {
66        Self {
67            model_id: artifact.model_id.clone(),
68            source_repo: Some(artifact.source_repo.clone()),
69            source_revision: Some(artifact.source_revision.clone()),
70            source_file: Some(artifact.primary_file.clone()),
71            canonical_ref: Some(artifact.canonical_ref.clone()),
72            distribution_id: Some(artifact.distribution_id.clone()),
73            selector: artifact.selector.clone(),
74        }
75    }
76}
77
78#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
79#[serde(rename_all = "snake_case")]
80pub enum ModelFormat {
81    Gguf,
82    Safetensors,
83}
84
85#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
86pub struct ModelArtifactFile {
87    pub path: String,
88    pub size_bytes: Option<u64>,
89    pub sha256: Option<String>,
90}
91
92impl ModelArtifactFile {
93    pub fn new(path: impl Into<String>) -> Self {
94        Self {
95            path: path.into(),
96            size_bytes: None,
97            sha256: None,
98        }
99    }
100}
101
102pub async fn resolve_model_artifact_ref(
103    model_ref: &str,
104    repository: &impl ModelRepository,
105) -> Result<ResolvedModelArtifact> {
106    let parsed = parse_model_ref(model_ref)?;
107    resolve_model_artifact(&parsed, repository).await
108}
109
110pub async fn resolve_model_artifact(
111    model_ref: &ModelRef,
112    repository: &impl ModelRepository,
113) -> Result<ResolvedModelArtifact> {
114    let source_revision = repository
115        .resolve_revision(&model_ref.repo, model_ref.revision.as_deref())
116        .await?;
117    let mut repo_files = repository
118        .list_files(&model_ref.repo, &source_revision)
119        .await?;
120    repo_files.sort_by(|left, right| left.path.cmp(&right.path));
121
122    let primary_file = select_primary_file(model_ref.selector.as_deref(), &repo_files)?;
123    let format = format_for_file(&primary_file.path)?;
124    let files = artifact_file_set(&primary_file.path, &repo_files);
125    let distribution_id = distribution_id_for_file(&primary_file.path)?;
126
127    Ok(ResolvedModelArtifact {
128        model_id: model_ref.display_id(),
129        source_repo: model_ref.repo.clone(),
130        source_revision: source_revision.clone(),
131        selector: model_ref.selector.clone(),
132        format,
133        files,
134        primary_file: primary_file.path.clone(),
135        canonical_ref: format_canonical_ref(&model_ref.repo, &source_revision, &primary_file.path),
136        distribution_id,
137    })
138}
139
140pub fn select_primary_artifact_file(
141    selector: Option<&str>,
142    files: &[ModelArtifactFile],
143) -> Result<ModelArtifactFile> {
144    select_primary_file(selector, files)
145}
146
147pub fn artifact_files_for_primary(
148    primary_file: &str,
149    files: &[ModelArtifactFile],
150) -> Vec<ModelArtifactFile> {
151    artifact_file_set(primary_file, files)
152}
153
154fn select_primary_file(
155    selector: Option<&str>,
156    files: &[ModelArtifactFile],
157) -> Result<ModelArtifactFile> {
158    let Some(selector) = selector else {
159        return select_default_file(files);
160    };
161
162    let selector_lower = selector.to_ascii_lowercase();
163    let gguf_exact = format!("{selector}.gguf").to_ascii_lowercase();
164    let gguf_split_prefix = format!("{selector}-00001-of-").to_ascii_lowercase();
165    let safetensors_exact = format!("{selector}.safetensors").to_ascii_lowercase();
166    let safetensors_split_prefix = format!("{selector}-00001-of-").to_ascii_lowercase();
167
168    select_ranked_file(files, |file, lower, basename| {
169        if lower == selector_lower || basename == selector_lower {
170            Some(0)
171        } else if gguf_matches_quant_selector(&file.path, selector) {
172            Some(1)
173        } else if basename == safetensors_exact {
174            Some(2)
175        } else if basename.starts_with(&safetensors_split_prefix)
176            && basename.ends_with(".safetensors")
177        {
178            Some(3)
179        } else if basename == gguf_exact {
180            Some(4)
181        } else if basename.starts_with(&gguf_split_prefix) && basename.ends_with(".gguf") {
182            Some(5)
183        } else {
184            None
185        }
186    })
187    .ok_or_else(|| {
188        anyhow::anyhow!("no model artifact matching selector '{selector}' in repository")
189    })
190}
191
192fn select_default_file(files: &[ModelArtifactFile]) -> Result<ModelArtifactFile> {
193    select_ranked_file(files, |_file, lower, basename| {
194        if basename == "model.safetensors" {
195            Some(0)
196        } else if is_split_safetensors_first_shard(basename) {
197            Some(1)
198        } else if lower.ends_with(".gguf") {
199            if is_known_gguf_sidecar(basename) {
200                return None;
201            }
202            if lower.contains("-000") && !lower.contains("-00001-of-") {
203                return None;
204            }
205            Some(if lower.contains("-00001-of-") { 2 } else { 3 })
206        } else {
207            None
208        }
209    })
210    .ok_or_else(|| anyhow::anyhow!("no supported model artifact files found in repository"))
211}
212
213fn select_ranked_file(
214    files: &[ModelArtifactFile],
215    mut rank: impl FnMut(&ModelArtifactFile, &str, &str) -> Option<u8>,
216) -> Option<ModelArtifactFile> {
217    files
218        .iter()
219        .filter_map(|file| {
220            let lower = file.path.to_ascii_lowercase();
221            let basename = basename_lower(&file.path);
222            rank(file, &lower, &basename).map(|rank| {
223                (
224                    rank,
225                    artifact_preference_score(&file.path),
226                    file.path.as_str(),
227                    file,
228                )
229            })
230        })
231        .min_by(|left, right| (left.0, left.1, left.2).cmp(&(right.0, right.1, right.2)))
232        .map(|(_, _, _, file)| file.clone())
233}
234
235fn artifact_file_set(primary_file: &str, files: &[ModelArtifactFile]) -> Vec<ModelArtifactFile> {
236    if let Some(primary) = split_gguf_shard_info(primary_file) {
237        let mut shards = files
238            .iter()
239            .filter(|file| {
240                split_gguf_shard_info(&file.path)
241                    .map(|candidate| {
242                        candidate.prefix == primary.prefix && candidate.total == primary.total
243                    })
244                    .unwrap_or(false)
245            })
246            .cloned()
247            .collect::<Vec<_>>();
248        shards.sort_by(|left, right| left.path.cmp(&right.path));
249        if !shards.is_empty() {
250            return shards;
251        }
252    }
253
254    vec![
255        files
256            .iter()
257            .find(|file| file.path == primary_file)
258            .cloned()
259            .unwrap_or_else(|| ModelArtifactFile::new(primary_file)),
260    ]
261}
262
263fn format_for_file(file: &str) -> Result<ModelFormat> {
264    if file.ends_with(".gguf") {
265        return Ok(ModelFormat::Gguf);
266    }
267    if file.ends_with(".safetensors") || file.ends_with(".safetensors.index.json") {
268        return Ok(ModelFormat::Safetensors);
269    }
270    bail!("unsupported model artifact file format: {file}")
271}
272
273fn distribution_id_for_file(file: &str) -> Result<String> {
274    if file.ends_with(".gguf") {
275        return normalize_gguf_distribution_id(file)
276            .ok_or_else(|| anyhow::anyhow!("invalid GGUF artifact file name: {file}"));
277    }
278    let basename = Path::new(file)
279        .file_name()
280        .and_then(|value| value.to_str())
281        .unwrap_or(file);
282    let stem = basename.strip_suffix(".safetensors").unwrap_or(basename);
283    Ok(split_safetensors_shard_stem_prefix(stem)
284        .unwrap_or(stem)
285        .to_string())
286}
287
288fn basename_lower(path: &str) -> String {
289    Path::new(path)
290        .file_name()
291        .and_then(|value| value.to_str())
292        .unwrap_or(path)
293        .to_ascii_lowercase()
294}
295
296fn artifact_preference_score(file: &str) -> usize {
297    if file.contains("-00001-of-") {
298        return 0;
299    }
300    const PREFERRED: &[&str] = &[
301        "Q4_K_M", "Q4_K_S", "Q4_1", "Q5_K_M", "Q5_K_S", "Q8_0", "BF16",
302    ];
303    PREFERRED
304        .iter()
305        .position(|needle| file.contains(needle))
306        .map(|pos| pos + 1)
307        .unwrap_or(PREFERRED.len() + 2)
308}
309
310fn is_known_gguf_sidecar(basename_lower: &str) -> bool {
311    basename_lower.starts_with("mmproj")
312}
313
314fn is_split_safetensors_first_shard(basename_lower: &str) -> bool {
315    let Some(stem) = basename_lower.strip_suffix(".safetensors") else {
316        return false;
317    };
318    split_safetensors_shard_info(stem)
319        .map(|(_, part, _)| part == "00001")
320        .unwrap_or(false)
321}
322
323fn split_safetensors_shard_stem_prefix(stem: &str) -> Option<&str> {
324    split_safetensors_shard_info(stem).map(|(prefix, _, _)| prefix)
325}
326
327fn split_safetensors_shard_info(stem: &str) -> Option<(&str, &str, &str)> {
328    let (prefix_and_part, total) = stem.rsplit_once("-of-")?;
329    if total.len() != 5 || !total.bytes().all(|byte| byte.is_ascii_digit()) {
330        return None;
331    }
332    let (prefix, part) = prefix_and_part.rsplit_once('-')?;
333    if part.len() != 5 || !part.bytes().all(|byte| byte.is_ascii_digit()) {
334        return None;
335    }
336    Some((prefix, part, total))
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342    use std::collections::HashMap;
343
344    struct MemoryRepository {
345        revision: String,
346        files: HashMap<String, Vec<ModelArtifactFile>>,
347    }
348
349    #[async_trait]
350    impl ModelRepository for MemoryRepository {
351        async fn resolve_revision(&self, _repo: &str, revision: Option<&str>) -> Result<String> {
352            Ok(revision.unwrap_or(&self.revision).to_string())
353        }
354
355        async fn list_files(&self, repo: &str, _revision: &str) -> Result<Vec<ModelArtifactFile>> {
356            Ok(self.files.get(repo).cloned().unwrap_or_default())
357        }
358    }
359
360    fn repo(files: Vec<&str>) -> MemoryRepository {
361        MemoryRepository {
362            revision: "abc123".to_string(),
363            files: HashMap::from([(
364                "org/repo".to_string(),
365                files.into_iter().map(ModelArtifactFile::new).collect(),
366            )]),
367        }
368    }
369
370    fn files(paths: &[&str]) -> Vec<ModelArtifactFile> {
371        paths.iter().copied().map(ModelArtifactFile::new).collect()
372    }
373
374    #[tokio::test]
375    async fn resolves_quant_selector_to_gguf_file() {
376        let repository = repo(vec!["Model-Q5_K_M.gguf", "Model-Q4_K_M.gguf", "README.md"]);
377
378        let resolved = resolve_model_artifact_ref("org/repo:Q4_K_M", &repository)
379            .await
380            .unwrap();
381
382        assert_eq!(resolved.model_id, "org/repo:Q4_K_M");
383        assert_eq!(resolved.source_revision, "abc123");
384        assert_eq!(resolved.primary_file, "Model-Q4_K_M.gguf");
385        assert_eq!(resolved.canonical_ref, "org/repo@abc123/Model-Q4_K_M.gguf");
386        assert_eq!(resolved.distribution_id, "Model-Q4_K_M");
387        assert_eq!(resolved.files.len(), 1);
388    }
389
390    #[tokio::test]
391    async fn resolves_split_gguf_selector_to_all_shards() {
392        let repository = repo(vec![
393            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00002-of-00003.gguf",
394            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00001-of-00003.gguf",
395            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00003-of-00003.gguf",
396            "UD-Q4_K_M/GLM-5.1-UD-Q4_K_M-00001-of-00003.gguf",
397        ]);
398
399        let resolved = resolve_model_artifact_ref("org/repo:UD-IQ2_M", &repository)
400            .await
401            .unwrap();
402
403        assert_eq!(
404            resolved.primary_file,
405            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00001-of-00003.gguf"
406        );
407        assert_eq!(resolved.distribution_id, "GLM-5.1-UD-IQ2_M");
408        assert_eq!(resolved.files.len(), 3);
409        assert_eq!(
410            resolved.files[2].path,
411            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00003-of-00003.gguf"
412        );
413    }
414
415    #[test]
416    fn public_selector_api_resolves_mesh_split_stem_to_first_part() {
417        let files = files(&[
418            "zai-org.GLM-5.1.Q2_K-00002-of-00018.gguf",
419            "zai-org.GLM-5.1.Q2_K-00001-of-00018.gguf",
420        ]);
421
422        let selected = select_primary_artifact_file(Some("zai-org.GLM-5.1.Q2_K"), &files).unwrap();
423
424        assert_eq!(selected.path, "zai-org.GLM-5.1.Q2_K-00001-of-00018.gguf");
425    }
426
427    #[test]
428    fn public_selector_api_resolves_mesh_quant_aliases() {
429        let files = files(&[
430            "qwen3.5-moe-0.87B-d0.8B.Q2_K.gguf",
431            "gemma-4-31B-it-Q4_0.gguf",
432            "Qwen3-8B-Q4_K_M.gguf",
433        ]);
434
435        assert_eq!(
436            select_primary_artifact_file(Some("Q2_K"), &files)
437                .unwrap()
438                .path,
439            "qwen3.5-moe-0.87B-d0.8B.Q2_K.gguf"
440        );
441        assert_eq!(
442            select_primary_artifact_file(Some("Q4_0"), &files)
443                .unwrap()
444                .path,
445            "gemma-4-31B-it-Q4_0.gguf"
446        );
447    }
448
449    #[test]
450    fn exact_filename_precedes_quant_selector_match() {
451        let files = files(&["Model-Q4_K_M.gguf", "Q4_K_M"]);
452
453        let selected = select_primary_artifact_file(Some("Q4_K_M"), &files).unwrap();
454
455        assert_eq!(selected.path, "Q4_K_M");
456    }
457
458    #[test]
459    fn selector_ranking_is_independent_of_input_order() {
460        let first_order = files(&["z/Model-Q4_K_M.gguf", "a/Model-Q4_K_M.gguf"]);
461        let second_order = files(&["a/Model-Q4_K_M.gguf", "z/Model-Q4_K_M.gguf"]);
462
463        let first = select_primary_artifact_file(Some("Q4_K_M"), &first_order).unwrap();
464        let second = select_primary_artifact_file(Some("Q4_K_M"), &second_order).unwrap();
465
466        assert_eq!(first.path, "a/Model-Q4_K_M.gguf");
467        assert_eq!(second.path, first.path);
468    }
469
470    #[test]
471    fn public_selector_api_resolves_mesh_mlx_shorthand() {
472        let files = files(&[
473            "model-00002-of-00048.safetensors",
474            "model-00001-of-00048.safetensors",
475            "model.safetensors.index.json",
476        ]);
477
478        let selected = select_primary_artifact_file(Some("model"), &files).unwrap();
479
480        assert_eq!(selected.path, "model-00001-of-00048.safetensors");
481    }
482
483    #[test]
484    fn public_default_api_preserves_mesh_default_ordering() {
485        let files = files(&[
486            "Qwen3-8B-Q8_0.gguf",
487            "mmproj-BF16.gguf",
488            "Qwen3-8B-Q4_K_M.gguf",
489        ]);
490
491        let selected = select_primary_artifact_file(None, &files).unwrap();
492
493        assert_eq!(selected.path, "Qwen3-8B-Q4_K_M.gguf");
494    }
495
496    #[test]
497    fn public_default_api_prefers_mlx_weights_over_gguf() {
498        let files = files(&[
499            "Qwen3-8B-Q4_K_M.gguf",
500            "model.safetensors",
501            "model.safetensors.index.json",
502        ]);
503
504        let selected = select_primary_artifact_file(None, &files).unwrap();
505
506        assert_eq!(selected.path, "model.safetensors");
507    }
508
509    #[test]
510    fn default_selection_rejects_sidecars_and_non_first_split_shards() {
511        let files = files(&["mmproj-model-f16.gguf", "Model-Q4_K_M-00002-of-00002.gguf"]);
512
513        let error = select_primary_artifact_file(None, &files).unwrap_err();
514
515        assert_eq!(
516            error.to_string(),
517            "no supported model artifact files found in repository"
518        );
519    }
520
521    #[test]
522    fn public_artifact_set_returns_all_split_gguf_shards() {
523        let files = files(&[
524            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00002-of-00003.gguf",
525            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00001-of-00003.gguf",
526            "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00003-of-00003.gguf",
527            "UD-Q4_K_M/GLM-5.1-UD-Q4_K_M-00001-of-00003.gguf",
528        ]);
529
530        let shards =
531            artifact_files_for_primary("UD-IQ2_M/GLM-5.1-UD-IQ2_M-00001-of-00003.gguf", &files);
532
533        assert_eq!(
534            shards
535                .iter()
536                .map(|file| file.path.as_str())
537                .collect::<Vec<_>>(),
538            vec![
539                "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00001-of-00003.gguf",
540                "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00002-of-00003.gguf",
541                "UD-IQ2_M/GLM-5.1-UD-IQ2_M-00003-of-00003.gguf",
542            ]
543        );
544    }
545
546    #[tokio::test]
547    async fn accepts_revisioned_selector_refs() {
548        let repository = repo(vec!["Model-Q4_K_M.gguf"]);
549
550        let resolved = resolve_model_artifact_ref("org/repo:Q4_K_M@rev-1", &repository)
551            .await
552            .unwrap();
553
554        assert_eq!(resolved.model_id, "org/repo@rev-1:Q4_K_M");
555        assert_eq!(resolved.source_revision, "rev-1");
556        assert_eq!(resolved.canonical_ref, "org/repo@rev-1/Model-Q4_K_M.gguf");
557    }
558
559    #[tokio::test]
560    async fn default_selection_prefers_primary_weights() {
561        let repository = repo(vec![
562            "README.md",
563            "Qwen3-8B-Q4_K_M.gguf",
564            "Qwen3-8B-Q5_K_M.gguf",
565        ]);
566
567        let resolved = resolve_model_artifact_ref("org/repo", &repository)
568            .await
569            .unwrap();
570
571        assert_eq!(resolved.primary_file, "Qwen3-8B-Q4_K_M.gguf");
572        assert_eq!(resolved.format, ModelFormat::Gguf);
573    }
574
575    #[tokio::test]
576    async fn unknown_selector_returns_error() {
577        let repository = repo(vec!["Model-Q4_K_M.gguf"]);
578
579        let error = resolve_model_artifact_ref("org/repo:Q5_K_M", &repository)
580            .await
581            .unwrap_err();
582
583        assert_eq!(
584            error.to_string(),
585            "no model artifact matching selector 'Q5_K_M' in repository"
586        );
587    }
588}