Skip to main content

gglib_core/services/
model_registrar.rs

1//! Model registrar service implementation.
2//!
3//! This service implements `ModelRegistrarPort` using the `ModelRepository`
4//! and `GgufParserPort` dependencies. It's used by the download manager
5//! to register completed downloads.
6
7use std::sync::Arc;
8
9use async_trait::async_trait;
10use chrono::Utc;
11
12use super::{HfOrigin, ModelOrigin, build_new_model};
13use crate::domain::{Model, NewModelFile};
14use crate::ports::{
15    CompletedDownload, GgufParserPort, ModelRegistrarPort, ModelRepository, RepositoryError,
16};
17
18/// Repository trait for model files metadata.
19///
20/// We don't depend on `gglib_db` directly - adapters inject the implementation.
21/// This type is re-exported from `gglib_db` for use in adapters.
22#[async_trait]
23pub trait ModelFilesRepositoryPort: Send + Sync {
24    /// Insert a new model file record.
25    async fn insert(&self, model_file: &NewModelFile) -> anyhow::Result<()>;
26}
27
28/// Implementation of the model registrar port.
29///
30/// This service composes over `ModelRepository` for persistence and
31/// `GgufParserPort` for metadata extraction.
32pub struct ModelRegistrar {
33    /// Repository for persisting models.
34    model_repo: Arc<dyn ModelRepository>,
35    /// Parser for extracting GGUF metadata.
36    gguf_parser: Arc<dyn GgufParserPort>,
37    /// Repository for persisting model file metadata.
38    model_files_repo: Option<Arc<dyn ModelFilesRepositoryPort>>,
39}
40
41impl ModelRegistrar {
42    /// Create a new model registrar.
43    ///
44    /// # Arguments
45    ///
46    /// * `model_repo` - Repository for persisting models
47    /// * `gguf_parser` - Parser for extracting GGUF metadata
48    /// * `model_files_repo` - Optional repository for persisting model file metadata
49    pub fn new(
50        model_repo: Arc<dyn ModelRepository>,
51        gguf_parser: Arc<dyn GgufParserPort>,
52        model_files_repo: Option<Arc<dyn ModelFilesRepositoryPort>>,
53    ) -> Self {
54        Self {
55            model_repo,
56            gguf_parser,
57            model_files_repo,
58        }
59    }
60}
61
62#[async_trait]
63impl ModelRegistrarPort for ModelRegistrar {
64    async fn register_model(&self, download: &CompletedDownload) -> Result<Model, RepositoryError> {
65        let file_path = download.db_path();
66
67        // Parse GGUF metadata from the downloaded file
68        let gguf_metadata = self.gguf_parser.parse(file_path).ok();
69
70        let origin = ModelOrigin::HuggingFace(HfOrigin {
71            repo_id: &download.repo_id,
72            commit_sha: &download.commit_sha,
73            hf_tags: &download.hf_tags,
74            quantization_fallback: download.quantization,
75            file_paths: download.file_paths.as_deref(),
76        });
77        let model = build_new_model(
78            file_path,
79            gguf_metadata.as_ref(),
80            self.gguf_parser.as_ref(),
81            &origin,
82            Utc::now(),
83        );
84
85        let registered = self.model_repo.insert(&model).await?;
86
87        // Insert model_files records with OIDs for each shard (if repo is available)
88        if let Some(ref repo) = self.model_files_repo {
89            for (file_index, file_entry) in download.hf_file_entries.iter().enumerate() {
90                if let Some(size) = file_entry.size {
91                    #[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
92                    let model_file = NewModelFile::new(
93                        registered.id,
94                        file_entry.path.clone(),
95                        file_index as i32,
96                        size as i64,
97                        file_entry.oid.clone(),
98                    );
99
100                    if let Err(e) = repo.insert(&model_file).await {
101                        // Soft fail - log but don't propagate error
102                        tracing::warn!(
103                            model_id = registered.id,
104                            file_path = %file_entry.path,
105                            error = %e,
106                            "Failed to insert model_files record - verification features may be unavailable"
107                        );
108                    }
109                }
110            }
111        }
112
113        Ok(registered)
114    }
115}
116
117#[cfg(test)]
118mod tests {
119    use super::*;
120    use crate::domain::{Model, NewModel};
121    use crate::download::Quantization;
122    use crate::ports::NoopGgufParser;
123    use std::path::PathBuf;
124    use std::sync::Mutex;
125
126    /// Mock model repository for testing.
127    struct MockModelRepo {
128        models: Mutex<Vec<Model>>,
129        next_id: Mutex<i64>,
130    }
131
132    impl MockModelRepo {
133        fn new() -> Self {
134            Self {
135                models: Mutex::new(Vec::new()),
136                next_id: Mutex::new(1),
137            }
138        }
139    }
140
141    #[async_trait]
142    impl ModelRepository for MockModelRepo {
143        async fn list(&self) -> Result<Vec<Model>, RepositoryError> {
144            Ok(self.models.lock().unwrap().clone())
145        }
146
147        async fn get_by_id(&self, id: i64) -> Result<Model, RepositoryError> {
148            self.models
149                .lock()
150                .unwrap()
151                .iter()
152                .find(|m| m.id == id)
153                .cloned()
154                .ok_or_else(|| RepositoryError::NotFound(format!("id={id}")))
155        }
156
157        async fn get_by_name(&self, name: &str) -> Result<Model, RepositoryError> {
158            self.models
159                .lock()
160                .unwrap()
161                .iter()
162                .find(|m| m.name == name)
163                .cloned()
164                .ok_or_else(|| RepositoryError::NotFound(format!("name={name}")))
165        }
166
167        async fn insert(&self, model: &NewModel) -> Result<Model, RepositoryError> {
168            let mut id = self.next_id.lock().unwrap();
169            let persisted = Model {
170                id: *id,
171                name: model.name.clone(),
172                model_key: String::new(),
173                file_path: model.file_path.clone(),
174                param_count_b: model.param_count_b,
175                architecture: model.architecture.clone(),
176                quantization: model.quantization.clone(),
177                context_length: model.context_length,
178                expert_count: model.expert_count,
179                expert_used_count: model.expert_used_count,
180                expert_shared_count: model.expert_shared_count,
181                metadata: model.metadata.clone(),
182                added_at: model.added_at,
183                hf_repo_id: model.hf_repo_id.clone(),
184                hf_commit_sha: model.hf_commit_sha.clone(),
185                hf_filename: model.hf_filename.clone(),
186                capabilities: model.capabilities,
187                download_date: model.download_date,
188                last_update_check: model.last_update_check,
189                tags: model.tags.clone(),
190                inference_defaults: model.inference_defaults.clone(),
191                defaults_origin: model.defaults_origin,
192                server_defaults: model.server_defaults.clone(),
193                benchmark_summary: None,
194            };
195            *id += 1;
196            drop(id);
197            self.models.lock().unwrap().push(persisted.clone());
198            Ok(persisted)
199        }
200
201        async fn update(&self, _model: &Model) -> Result<(), RepositoryError> {
202            Ok(())
203        }
204
205        async fn delete(&self, _id: i64) -> Result<(), RepositoryError> {
206            Ok(())
207        }
208    }
209
210    #[tokio::test]
211    async fn test_register_model_basic() {
212        let repo = Arc::new(MockModelRepo::new());
213        let parser = Arc::new(NoopGgufParser);
214        let registrar = ModelRegistrar::new(repo.clone(), parser, None);
215
216        let download = CompletedDownload {
217            primary_path: PathBuf::from("/models/test-model-q4_k_m.gguf"),
218            all_paths: vec![PathBuf::from("/models/test-model-q4_k_m.gguf")],
219            quantization: Quantization::Q4KM,
220            repo_id: "test/model".to_string(),
221            commit_sha: "abc123".to_string(),
222            is_sharded: false,
223            file_paths: None,
224            hf_tags: vec![],
225            hf_file_entries: vec![],
226        };
227
228        let result = registrar.register_model(&download).await;
229        assert!(result.is_ok());
230
231        let model = result.unwrap();
232        assert_eq!(model.name, "model");
233        assert_eq!(model.hf_repo_id, Some("test/model".to_string()));
234        assert_eq!(model.hf_commit_sha, Some("abc123".to_string()));
235        assert_eq!(model.quantization, Some("Q4_K_M".to_string()));
236    }
237
238    #[tokio::test]
239    async fn test_register_sharded_model() {
240        let repo = Arc::new(MockModelRepo::new());
241        let parser = Arc::new(NoopGgufParser);
242        let registrar = ModelRegistrar::new(repo.clone(), parser, None);
243
244        let download = CompletedDownload {
245            primary_path: PathBuf::from("/models/llama-00001-of-00004.gguf"),
246            all_paths: vec![
247                PathBuf::from("/models/llama-00001-of-00004.gguf"),
248                PathBuf::from("/models/llama-00002-of-00004.gguf"),
249                PathBuf::from("/models/llama-00003-of-00004.gguf"),
250                PathBuf::from("/models/llama-00004-of-00004.gguf"),
251            ],
252            quantization: Quantization::Q8_0,
253            repo_id: "test/llama".to_string(),
254            commit_sha: "def456".to_string(),
255            is_sharded: true,
256            file_paths: None,
257            hf_tags: vec![],
258            hf_file_entries: vec![],
259        };
260
261        let result = registrar.register_model(&download).await;
262        assert!(result.is_ok());
263
264        let model = result.unwrap();
265        assert_eq!(model.quantization, Some("Q8_0".to_string()));
266        assert_eq!(model.name, "llama");
267    }
268}