1use 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#[async_trait]
23pub trait ModelFilesRepositoryPort: Send + Sync {
24 async fn insert(&self, model_file: &NewModelFile) -> anyhow::Result<()>;
26}
27
28pub struct ModelRegistrar {
33 model_repo: Arc<dyn ModelRepository>,
35 gguf_parser: Arc<dyn GgufParserPort>,
37 model_files_repo: Option<Arc<dyn ModelFilesRepositoryPort>>,
39}
40
41impl ModelRegistrar {
42 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 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 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 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 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}