use std::sync::Arc;
use std::sync::OnceLock;
use ferrin_spec::DynLanguageModel;
use ferrin_spec::LanguageModelRef;
use ferrin_spec::ModelRef;
use super::ProviderRegistry;
use crate::error::Error;
static DEFAULT_REGISTRY: OnceLock<Arc<ProviderRegistry>> = OnceLock::new();
pub fn set_default_registry(registry: Arc<ProviderRegistry>) -> Result<(), Error> {
DEFAULT_REGISTRY
.set(registry)
.map_err(|_| Error::invalid_argument("registry", "default registry is already set"))
}
#[must_use]
pub fn default_registry() -> Option<Arc<ProviderRegistry>> {
DEFAULT_REGISTRY.get().cloned()
}
pub(crate) fn resolve_model<D: ?Sized>(
model: &ModelRef<D>,
lookup: impl FnOnce(&ProviderRegistry, &str) -> Result<ModelRef<D>, Error>,
) -> Result<Arc<D>, Error> {
if let Some(model) = model.model() {
return Ok(Arc::clone(model));
}
let id = model.unresolved_id().unwrap_or_default();
let Some(registry) = default_registry() else {
return Err(Error::NoDefaultRegistry {
model_id: id.to_owned(),
});
};
lookup(®istry, id)?
.into_model()
.map_err(|id| Error::NoDefaultRegistry { model_id: id })
}
pub(crate) fn resolve_language_model(
model: &LanguageModelRef,
) -> Result<Arc<dyn DynLanguageModel>, Error> {
resolve_model(model, ProviderRegistry::language_model)
}