Skip to main content

studio_worker/
loaders.rs

1//! The worker's in-process model loaders, one per engine, behind the
2//! model host's `ModelRuntime` (see `docs/runtime/model-lifecycle.md`).
3
4use crate::catalog::CatalogModel;
5use crate::host::{LoadedModel, ModelRuntime};
6use std::path::PathBuf;
7use std::sync::Arc;
8
9/// Dispatches a load to the loader for the model's engine.  An engine
10/// with no in-process loader is refused by name, never faked.
11pub struct Loaders {
12    models_root: PathBuf,
13}
14
15impl Loaders {
16    /// `models_root`: where model files are downloaded to.
17    pub fn new(models_root: PathBuf) -> Self {
18        Self { models_root }
19    }
20}
21
22impl ModelRuntime for Loaders {
23    fn can_load(&self, model: &CatalogModel) -> bool {
24        match model.source.engine {
25            #[cfg(all(feature = "llama", not(target_os = "windows")))]
26            crate::types::ModelEngine::LlamaCpp => true,
27            #[cfg(feature = "stt-stream")]
28            crate::types::ModelEngine::Parakeet => true,
29            _ => false,
30        }
31    }
32
33    fn load(&self, model: &CatalogModel) -> anyhow::Result<Arc<dyn LoadedModel>> {
34        match &model.source.engine {
35            #[cfg(all(feature = "llama", not(target_os = "windows")))]
36            crate::types::ModelEngine::LlamaCpp => Ok(Arc::new(
37                crate::engine::llama::load_resident(&self.models_root, model)?,
38            )),
39            #[cfg(feature = "stt-stream")]
40            crate::types::ModelEngine::Parakeet => Ok(Arc::new(
41                crate::engine::parakeet::load_resident(&self.models_root, model)?,
42            )),
43            engine => {
44                let _ = &self.models_root;
45                anyhow::bail!(
46                    "no in-process loader for engine {engine:?} (model {})",
47                    model.id
48                )
49            }
50        }
51    }
52}
53
54#[cfg(test)]
55mod tests {
56    use super::*;
57    use crate::types::{ModelEngine, ModelSource, TaskKind};
58
59    #[test]
60    fn an_engine_without_a_loader_is_refused_by_name() {
61        let model = CatalogModel {
62            id: "m".into(),
63            display_name: "m".into(),
64            kind: TaskKind::Image,
65            vram_gb_estimate: 1.0,
66            description: None,
67            source: ModelSource {
68                engine: ModelEngine::SdCpp,
69                files: vec![],
70                cli_defaults: Default::default(),
71            },
72            enabled: true,
73            origin: "local".into(),
74            exclusive_group: None,
75        };
76        let loaders = Loaders::new(PathBuf::from("/nonexistent"));
77        assert!(!loaders.can_load(&model));
78        let err = loaders.load(&model).err().expect("refused");
79        assert!(
80            err.to_string()
81                .contains("no in-process loader for engine SdCpp"),
82            "{err}"
83        );
84    }
85
86    #[cfg(all(feature = "llama", not(target_os = "windows")))]
87    #[test]
88    fn an_llm_can_be_loaded() {
89        let model = CatalogModel {
90            id: "m".into(),
91            display_name: "m".into(),
92            kind: TaskKind::Llm,
93            vram_gb_estimate: 1.0,
94            description: None,
95            source: ModelSource {
96                engine: ModelEngine::LlamaCpp,
97                files: vec![],
98                cli_defaults: Default::default(),
99            },
100            enabled: true,
101            origin: "local".into(),
102            exclusive_group: None,
103        };
104        assert!(Loaders::new(PathBuf::from("/nonexistent")).can_load(&model));
105    }
106}