use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use serde::{Deserialize, Serialize};
use crate::persistence::{self, StoreError};
use crate::records::{ModelRecord, ModelState};
const STORE_FILE: &str = "models.json";
const LOCK_FILE: &str = "models.json.lock";
const SCHEMA_VERSION: u32 = 1;
#[derive(Debug, thiserror::Error)]
pub enum RegistryError {
#[error("corrupt registry store: {0}")]
CorruptStore(String),
#[error("registry store schema {found} is newer than supported {supported}")]
FutureSchema {
found: u32,
supported: u32,
},
#[error(transparent)]
Store(#[from] StoreError),
#[error("locking registry store: {0}")]
Lock(String),
}
#[derive(Debug, Deserialize)]
struct Envelope {
schema_version: u32,
models: Vec<ModelRecord>,
}
#[derive(Serialize)]
struct EnvelopeRef<'a> {
schema_version: u32,
models: Vec<&'a ModelRecord>,
}
#[derive(Debug)]
pub struct Registry {
directory: PathBuf,
models: BTreeMap<String, ModelRecord>,
generation: u64,
}
impl Registry {
pub fn open(directory: &Path) -> Result<Self, RegistryError> {
let models = Self::load_models(directory)?;
Ok(Self {
directory: directory.to_path_buf(),
models,
generation: 0,
})
}
fn load_models(directory: &Path) -> Result<BTreeMap<String, ModelRecord>, RegistryError> {
let file = directory.join(STORE_FILE);
match persistence::read_json::<Envelope>(&file) {
Ok(Some(envelope)) => {
if envelope.schema_version > SCHEMA_VERSION {
return Err(RegistryError::FutureSchema {
found: envelope.schema_version,
supported: SCHEMA_VERSION,
});
}
Ok(envelope
.models
.into_iter()
.map(|mut record| {
record.adopt_legacy_footprint();
(record.id.clone(), record)
})
.collect())
}
Ok(None) => Ok(BTreeMap::new()),
Err(StoreError::Corrupt { source, .. }) => {
Err(RegistryError::CorruptStore(source.to_string()))
}
Err(other) => Err(RegistryError::Store(other)),
}
}
fn reload(&mut self) -> Result<(), RegistryError> {
self.models = Self::load_models(&self.directory)?;
Ok(())
}
pub fn refresh(&mut self) -> Result<bool, RegistryError> {
let models = Self::load_models(&self.directory)?;
if models == self.models {
return Ok(false);
}
self.models = models;
self.generation += 1;
Ok(true)
}
fn lock(&self) -> Result<std::fs::File, RegistryError> {
use fs2::FileExt;
std::fs::create_dir_all(&self.directory)
.map_err(|source| RegistryError::Lock(source.to_string()))?;
let path = self.directory.join(LOCK_FILE);
let file = std::fs::OpenOptions::new()
.create(true)
.truncate(false)
.read(true)
.write(true)
.open(&path)
.map_err(|source| RegistryError::Lock(source.to_string()))?;
file.lock_exclusive()
.map_err(|source| RegistryError::Lock(source.to_string()))?;
Ok(file)
}
pub fn get(&self, id: &str) -> Option<&ModelRecord> {
self.models.get(id)
}
pub fn contains(&self, id: &str) -> bool {
self.models.contains_key(id)
}
pub fn len(&self) -> usize {
self.models.len()
}
pub fn is_empty(&self) -> bool {
self.models.is_empty()
}
pub fn generation(&self) -> u64 {
self.generation
}
pub fn list(&self) -> Vec<&ModelRecord> {
let mut records: Vec<&ModelRecord> = self.models.values().collect();
records.sort_by_cached_key(|record| (record.name.to_lowercase(), record.id.clone()));
records
}
pub fn register(&mut self, record: ModelRecord) -> Result<bool, RegistryError> {
let _lock = self.lock()?;
self.reload()?;
if self.models.get(&record.id) == Some(&record) {
return Ok(false);
}
self.models.insert(record.id.clone(), record);
self.save()?;
Ok(true)
}
pub fn register_all(&mut self, records: Vec<ModelRecord>) -> Result<usize, RegistryError> {
let _lock = self.lock()?;
self.reload()?;
let mut changed = 0;
for record in records {
if self.models.get(&record.id) != Some(&record) {
self.models.insert(record.id.clone(), record);
changed += 1;
}
}
if changed > 0 {
self.save()?;
}
Ok(changed)
}
pub fn unregister(&mut self, id: &str) -> Result<Option<ModelRecord>, RegistryError> {
let _lock = self.lock()?;
self.reload()?;
let removed = self.models.remove(id);
if removed.is_some() {
self.save()?;
}
Ok(removed)
}
pub fn set_state_if_present(
&mut self,
id: &str,
state: ModelState,
) -> Result<bool, RegistryError> {
let _lock = self.lock()?;
self.reload()?;
let Some(record) = self.models.get_mut(id) else {
return Ok(false);
};
if record.state == state {
return Ok(true);
}
record.state = state;
self.save()?;
Ok(true)
}
pub fn update(
&mut self,
ids: &[String],
transform: impl Fn(&ModelRecord) -> Option<ModelRecord>,
) -> Result<Vec<ModelRecord>, RegistryError> {
let _lock = self.lock()?;
self.reload()?;
let mut changed = Vec::new();
for id in ids {
let Some(next) = self.models.get(id).and_then(&transform) else {
continue;
};
if self.models.get(id) == Some(&next) {
continue;
}
if next.id != *id {
self.models.remove(id);
}
self.models.insert(next.id.clone(), next.clone());
changed.push(next);
}
if !changed.is_empty() {
self.save()?;
}
Ok(changed)
}
fn save(&mut self) -> Result<(), RegistryError> {
let envelope = EnvelopeRef {
schema_version: SCHEMA_VERSION,
models: self.models.values().collect(),
};
persistence::write_json_atomic(&self.directory.join(STORE_FILE), &envelope)?;
self.generation += 1;
Ok(())
}
}