1use std::collections::BTreeMap;
5use std::path::{Path, PathBuf};
6
7use serde::{Deserialize, Serialize};
8
9use crate::persistence::{self, StoreError};
10use crate::records::{ModelRecord, ModelState};
11
12const STORE_FILE: &str = "models.json";
13const LOCK_FILE: &str = "models.json.lock";
14const SCHEMA_VERSION: u32 = 1;
15
16#[derive(Debug, thiserror::Error)]
18pub enum RegistryError {
19 #[error("corrupt registry store: {0}")]
22 CorruptStore(String),
23
24 #[error("registry store schema {found} is newer than supported {supported}")]
27 FutureSchema {
28 found: u32,
30 supported: u32,
32 },
33
34 #[error(transparent)]
36 Store(#[from] StoreError),
37
38 #[error("locking registry store: {0}")]
40 Lock(String),
41}
42
43#[derive(Debug, Deserialize)]
44struct Envelope {
45 schema_version: u32,
46 models: Vec<ModelRecord>,
47}
48
49#[derive(Serialize)]
50struct EnvelopeRef<'a> {
51 schema_version: u32,
52 models: Vec<&'a ModelRecord>,
53}
54
55#[derive(Debug)]
62pub struct Registry {
63 directory: PathBuf,
64 models: BTreeMap<String, ModelRecord>,
65 generation: u64,
66}
67
68impl Registry {
69 pub fn open(directory: &Path) -> Result<Self, RegistryError> {
75 let models = Self::load_models(directory)?;
76 Ok(Self {
77 directory: directory.to_path_buf(),
78 models,
79 generation: 0,
80 })
81 }
82
83 fn load_models(directory: &Path) -> Result<BTreeMap<String, ModelRecord>, RegistryError> {
86 let file = directory.join(STORE_FILE);
87 match persistence::read_json::<Envelope>(&file) {
88 Ok(Some(envelope)) => {
89 if envelope.schema_version > SCHEMA_VERSION {
90 return Err(RegistryError::FutureSchema {
91 found: envelope.schema_version,
92 supported: SCHEMA_VERSION,
93 });
94 }
95 Ok(envelope
96 .models
97 .into_iter()
98 .map(|record| (record.id.clone(), record))
99 .collect())
100 }
101 Ok(None) => Ok(BTreeMap::new()),
102 Err(StoreError::Corrupt { source, .. }) => {
103 Err(RegistryError::CorruptStore(source.to_string()))
104 }
105 Err(other) => Err(RegistryError::Store(other)),
106 }
107 }
108
109 fn reload(&mut self) -> Result<(), RegistryError> {
112 self.models = Self::load_models(&self.directory)?;
113 Ok(())
114 }
115
116 fn lock(&self) -> Result<std::fs::File, RegistryError> {
120 use fs2::FileExt;
121 std::fs::create_dir_all(&self.directory)
122 .map_err(|source| RegistryError::Lock(source.to_string()))?;
123 let path = self.directory.join(LOCK_FILE);
124 let file = std::fs::OpenOptions::new()
125 .create(true)
126 .truncate(false)
127 .read(true)
128 .write(true)
129 .open(&path)
130 .map_err(|source| RegistryError::Lock(source.to_string()))?;
131 file.lock_exclusive()
132 .map_err(|source| RegistryError::Lock(source.to_string()))?;
133 Ok(file)
134 }
135
136 pub fn get(&self, id: &str) -> Option<&ModelRecord> {
138 self.models.get(id)
139 }
140
141 pub fn contains(&self, id: &str) -> bool {
143 self.models.contains_key(id)
144 }
145
146 pub fn len(&self) -> usize {
148 self.models.len()
149 }
150
151 pub fn is_empty(&self) -> bool {
153 self.models.is_empty()
154 }
155
156 pub fn generation(&self) -> u64 {
160 self.generation
161 }
162
163 pub fn list(&self) -> Vec<&ModelRecord> {
165 let mut records: Vec<&ModelRecord> = self.models.values().collect();
166 records.sort_by_cached_key(|record| (record.name.to_lowercase(), record.id.clone()));
167 records
168 }
169
170 pub fn register(&mut self, record: ModelRecord) -> Result<bool, RegistryError> {
173 let _lock = self.lock()?;
174 self.reload()?;
175 if self.models.get(&record.id) == Some(&record) {
176 return Ok(false);
177 }
178 self.models.insert(record.id.clone(), record);
179 self.save()?;
180 Ok(true)
181 }
182
183 pub fn register_all(&mut self, records: Vec<ModelRecord>) -> Result<usize, RegistryError> {
187 let _lock = self.lock()?;
188 self.reload()?;
189 let mut changed = 0;
190 for record in records {
191 if self.models.get(&record.id) != Some(&record) {
192 self.models.insert(record.id.clone(), record);
193 changed += 1;
194 }
195 }
196 if changed > 0 {
197 self.save()?;
198 }
199 Ok(changed)
200 }
201
202 pub fn unregister(&mut self, id: &str) -> Result<Option<ModelRecord>, RegistryError> {
204 let _lock = self.lock()?;
205 self.reload()?;
206 let removed = self.models.remove(id);
207 if removed.is_some() {
208 self.save()?;
209 }
210 Ok(removed)
211 }
212
213 pub fn set_state_if_present(
216 &mut self,
217 id: &str,
218 state: ModelState,
219 ) -> Result<bool, RegistryError> {
220 let _lock = self.lock()?;
221 self.reload()?;
222 let Some(record) = self.models.get_mut(id) else {
223 return Ok(false);
224 };
225 if record.state == state {
226 return Ok(true);
227 }
228 record.state = state;
229 self.save()?;
230 Ok(true)
231 }
232
233 pub fn update(
239 &mut self,
240 ids: &[String],
241 transform: impl Fn(&ModelRecord) -> Option<ModelRecord>,
242 ) -> Result<Vec<ModelRecord>, RegistryError> {
243 let _lock = self.lock()?;
244 self.reload()?;
245 let mut changed = Vec::new();
246 for id in ids {
247 let Some(next) = self.models.get(id).and_then(&transform) else {
248 continue;
249 };
250 if self.models.get(id) == Some(&next) {
251 continue;
252 }
253 if next.id != *id {
254 self.models.remove(id);
255 }
256 self.models.insert(next.id.clone(), next.clone());
257 changed.push(next);
258 }
259 if !changed.is_empty() {
260 self.save()?;
261 }
262 Ok(changed)
263 }
264
265 fn save(&mut self) -> Result<(), RegistryError> {
266 let envelope = EnvelopeRef {
267 schema_version: SCHEMA_VERSION,
268 models: self.models.values().collect(),
269 };
270 persistence::write_json_atomic(&self.directory.join(STORE_FILE), &envelope)?;
271 self.generation += 1;
272 Ok(())
273 }
274}