use std::sync::Arc;
use regex::Regex;
use crate::client::ConfigProvider;
use crate::embedded_config::EmbeddedConfigLoader;
use crate::config::{ModelConfig, ServiceConfig, VerificationStatus};
use crate::error::{ClientError, Result};
use crate::types::RequestBuilder;
use crate::export::{
RegistryExport, ServiceExport, ModelExport, RegistryStats, RateLimitsExport,
};
use crate::registry_index::{RegistryIndex, IndexStats, BrokenReference};
use crate::query::ModelQuery;
use crate::validation::{ConfigValidator, ValidationLevel, ValidationReport};
pub struct ModelRegistry {
config_provider: Arc<dyn ConfigProvider + Send + Sync>,
index: RegistryIndex,
}
impl ModelRegistry {
pub fn new() -> Result<Self> {
let embedded_loader = EmbeddedConfigLoader::new()?;
let index = RegistryIndex::build(&embedded_loader);
index.validate_refs()?;
Ok(Self {
config_provider: Arc::new(embedded_loader),
index,
})
}
pub fn new_permissive() -> Result<Self> {
let embedded_loader = EmbeddedConfigLoader::new()?;
let index = RegistryIndex::build(&embedded_loader);
Ok(Self {
config_provider: Arc::new(embedded_loader),
index,
})
}
pub fn with_provider<T: ConfigProvider + Send + Sync + 'static>(provider: T) -> Result<Self> {
let index = RegistryIndex::build(&provider);
index.validate_refs()?;
Ok(Self {
config_provider: Arc::new(provider),
index,
})
}
pub fn with_provider_permissive<T: ConfigProvider + Send + Sync + 'static>(provider: T) -> Self {
let index = RegistryIndex::build(&provider);
Self {
config_provider: Arc::new(provider),
index,
}
}
pub fn from_id(&self, model_id: &str) -> Result<RequestBuilder> {
let _ = self.config_provider.get_model(model_id)?;
RequestBuilder::new(model_id.to_string(), Arc::clone(&self.config_provider))
}
pub fn use_cheapest(&self, pattern: &str) -> Result<RequestBuilder> {
let model_id = self.find_cheapest_model(pattern)?;
self.from_id(&model_id)
}
pub fn use_fastest(&self, pattern: &str) -> Result<RequestBuilder> {
let model_id = self.find_fastest_model(pattern)?;
self.from_id(&model_id)
}
pub fn use_best_quality(&self, pattern: &str) -> Result<RequestBuilder> {
let model_id = self.find_best_quality_model(pattern)?;
self.from_id(&model_id)
}
pub fn list_models(&self) -> Vec<&str> {
self.config_provider.list_models()
}
pub fn list_models_matching(&self, pattern: &str) -> Result<Vec<&str>> {
let regex = self.pattern_to_regex(pattern)?;
let models = self.config_provider.list_models();
Ok(models.into_iter()
.filter(|model_id| regex.is_match(model_id))
.collect())
}
pub fn get_model_info(&self, model_id: &str) -> Result<&ModelConfig> {
self.config_provider.get_model(model_id)
}
pub fn list_families(&self) -> Vec<String> {
self.index.all_families().into_iter().map(String::from).collect()
}
pub fn list_models_in_family(&self, family: &str) -> Vec<&str> {
self.index.models_in_family(family)
.iter()
.map(String::as_str)
.collect()
}
pub fn list_services(&self) -> Vec<&str> {
self.config_provider.list_services()
}
pub fn get_service(&self, name: &str) -> Result<&ServiceConfig> {
self.config_provider.get_service(name)
}
pub fn get_model_with_service(&self, model_id: &str) -> Result<(&ModelConfig, &ServiceConfig)> {
self.config_provider.get_model_with_service(model_id)
}
pub fn list_verified_models(&self) -> Vec<&str> {
self.index.models_with_status(&VerificationStatus::Verified)
.iter()
.map(String::as_str)
.collect()
}
pub fn list_models_by_status(&self, status: VerificationStatus) -> Vec<&str> {
self.index.models_with_status(&status)
.iter()
.map(String::as_str)
.collect()
}
pub fn is_verified(&self, model_id: &str) -> bool {
self.config_provider.get_model(model_id)
.map(|cfg| cfg.model.status == VerificationStatus::Verified)
.unwrap_or(false)
}
pub fn export(&self) -> RegistryExport {
let service_model_counts = self.index.model_counts_by_service();
let model_status_counts = self.index.model_counts_by_status();
let verified_count = model_status_counts.get(&VerificationStatus::Verified).copied().unwrap_or(0);
let unverified_count = model_status_counts.get(&VerificationStatus::Unverified).copied().unwrap_or(0);
let models: Vec<ModelExport> = self.config_provider.list_models()
.into_iter()
.filter_map(|model_id| {
let cfg = self.config_provider.get_model(model_id).ok()?;
Some(ModelExport::from(cfg))
})
.collect();
let services: Vec<ServiceExport> = self.config_provider.list_services()
.into_iter()
.filter_map(|service_name| {
let cfg = self.config_provider.get_service(service_name).ok()?;
let rate_limits = if cfg.rate_limits.requests_per_minute.is_some()
|| cfg.rate_limits.tokens_per_minute.is_some()
|| cfg.rate_limits.concurrent_requests.is_some()
{
Some(RateLimitsExport::from(&cfg.rate_limits))
} else {
None
};
Some(ServiceExport {
name: service_name.to_string(),
base_url: cfg.service.base_url.clone(),
message_format: cfg.message_builder.clone().unwrap_or_default(),
rate_limits,
model_count: service_model_counts.get(service_name).copied().unwrap_or(0),
})
})
.collect();
let families = self.list_families();
RegistryExport {
stats: RegistryStats {
service_count: services.len(),
family_count: families.len(),
model_count: models.len(),
verified_count,
unverified_count,
},
services,
families,
models,
}
}
pub fn export_verified(&self) -> RegistryExport {
self.export().verified_only()
}
pub fn export_by_service(&self, service: &str) -> RegistryExport {
self.export().filter_by_service(service)
}
pub fn export_by_family(&self, family: &str) -> RegistryExport {
self.export().filter_by_family(family)
}
pub fn query(&self) -> ModelQuery<'_, dyn ConfigProvider + Send + Sync> {
ModelQuery::new(self.config_provider.as_ref(), &self.index)
}
pub fn index(&self) -> &RegistryIndex {
&self.index
}
pub fn index_stats(&self) -> IndexStats {
self.index.stats()
}
pub fn models_for_service(&self, service: &str) -> &[String] {
self.index.models_for_service(service)
}
pub fn service_for_model(&self, model_id: &str) -> Option<&str> {
self.index.service_for_model(model_id)
}
pub fn broken_refs(&self) -> &[BrokenReference] {
self.index.broken_refs()
}
pub fn orphan_services(&self) -> &[String] {
self.index.orphan_services()
}
pub fn has_broken_refs(&self) -> bool {
self.index.has_broken_refs()
}
pub fn validate(&self, levels: &[ValidationLevel]) -> ValidationReport {
let validator = ConfigValidator::new(self.config_provider.as_ref(), &self.index);
validator.validate(levels)
}
pub fn validate_model(&self, model_id: &str, levels: &[ValidationLevel]) -> ValidationReport {
let validator = ConfigValidator::new(self.config_provider.as_ref(), &self.index);
validator.validate_model(model_id, levels)
}
pub fn validate_service(&self, service_name: &str, levels: &[ValidationLevel]) -> ValidationReport {
let validator = ConfigValidator::new(self.config_provider.as_ref(), &self.index);
validator.validate_service(service_name, levels)
}
pub fn is_valid(&self) -> bool {
let report = self.validate(&[
ValidationLevel::Schema,
ValidationLevel::CrossRef,
ValidationLevel::Semantic,
]);
report.passed
}
fn find_cheapest_model(&self, pattern: &str) -> Result<String> {
let matching_models = self.list_models_matching(pattern)?;
if matching_models.is_empty() {
return Err(ClientError::Config(crate::error::ConfigError::ModelNotFound(
format!("No models found matching pattern: {}", pattern)
)));
}
let mut cheapest_model = None;
let mut lowest_cost = f64::MAX;
for model_id in matching_models {
if let Ok(model_config) = self.config_provider.get_model(model_id) {
let avg_cost = (model_config.pricing.input_per_1k_tokens * 0.7) +
(model_config.pricing.output_per_1k_tokens * 0.3);
if avg_cost < lowest_cost {
lowest_cost = avg_cost;
cheapest_model = Some(model_id.to_string());
}
}
}
cheapest_model.ok_or_else(|| ClientError::Config(crate::error::ConfigError::ModelNotFound(
format!("No valid models found matching pattern: {}", pattern)
)))
}
fn find_fastest_model(&self, pattern: &str) -> Result<String> {
let matching_models = self.list_models_matching(pattern)?;
if matching_models.is_empty() {
return Err(ClientError::Config(crate::error::ConfigError::ModelNotFound(
format!("No models found matching pattern: {}", pattern)
)));
}
let speed_order = ["haiku", "sonnet", "opus", "mini", "turbo"];
for variant in &speed_order {
for model_id in &matching_models {
if model_id.contains(variant) {
return Ok(model_id.to_string());
}
}
}
Ok(matching_models[0].to_string())
}
fn find_best_quality_model(&self, pattern: &str) -> Result<String> {
let matching_models = self.list_models_matching(pattern)?;
if matching_models.is_empty() {
return Err(ClientError::Config(crate::error::ConfigError::ModelNotFound(
format!("No models found matching pattern: {}", pattern)
)));
}
let quality_order = ["opus", "sonnet", "haiku", "turbo", "mini"];
for variant in &quality_order {
for model_id in &matching_models {
if model_id.contains(variant) {
return Ok(model_id.to_string());
}
}
}
Ok(matching_models[0].to_string())
}
fn pattern_to_regex(&self, pattern: &str) -> Result<Regex> {
let escaped = regex::escape(pattern);
let regex_pattern = escaped.replace(r"\*", ".*").replace(r"\?", ".");
let regex_pattern = format!("^{}$", regex_pattern);
Regex::new(®ex_pattern).map_err(|e| ClientError::Config(
crate::error::ConfigError::InvalidPath(format!("Invalid pattern: {}", e))
))
}
}
impl Default for ModelRegistry {
fn default() -> Self {
Self::new().expect("Failed to create default runtime")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_runtime_creation() {
let runtime = ModelRegistry::new().unwrap();
let models = runtime.list_models();
assert!(!models.is_empty());
}
#[test]
fn test_model_pattern_matching() {
let runtime = ModelRegistry::new().unwrap();
let exact_matches = runtime.list_models_matching("claude-3-haiku-20240307").unwrap();
assert!(exact_matches.contains(&"claude-3-haiku-20240307"));
let haiku_matches = runtime.list_models_matching("claude-3-haiku-*").unwrap();
assert!(!haiku_matches.is_empty());
assert!(haiku_matches.iter().all(|m| m.contains("claude-3-haiku")));
}
#[test]
fn test_cheapest_model_selection() {
let runtime = ModelRegistry::new().unwrap();
let cheapest = runtime.use_cheapest("claude-3-haiku-*").unwrap();
assert!(cheapest.model_id.contains("claude-3-haiku"));
}
#[test]
fn test_model_info_retrieval() {
let runtime = ModelRegistry::new().unwrap();
let models = runtime.list_models();
if let Some(model_id) = models.first() {
let model_info = runtime.get_model_info(model_id).unwrap();
assert_eq!(&model_info.model.id, model_id);
}
}
}