systemprompt_config/services/
provider_catalog.rs1use std::collections::HashMap;
14
15use systemprompt_identifiers::{ModelId, ProviderId, SecretName};
16use systemprompt_models::profile::{
17 ApiSurface, ProviderEntry, ProviderModel, ProviderRegistry, WireProtocol,
18};
19use systemprompt_models::services::ai::{ModelCapabilities, ModelLimits, ModelPricing};
20
21use crate::error::{ConfigError, ConfigResult};
22
23#[derive(Debug, Clone)]
24pub struct ProviderSpec {
25 pub name: ProviderId,
26 pub wire: WireProtocol,
27 pub surface: ApiSurface,
28 pub endpoint: String,
29 pub api_key_secret: SecretName,
30 pub extra_headers: HashMap<String, String>,
31}
32
33#[derive(Debug, Clone)]
34pub struct ModelSpec {
35 pub provider: ProviderId,
36 pub id: ModelId,
37 pub aliases: Vec<ModelId>,
38 pub upstream_model: Option<String>,
39}
40
41#[derive(Debug, Clone, Copy)]
42pub struct ProviderCatalogService;
43
44impl ProviderCatalogService {
45 pub fn upsert_provider(registry: &mut ProviderRegistry, spec: ProviderSpec) {
46 let (models, governance) = registry
47 .find_provider(spec.name.as_str())
48 .map(|p| (p.models.clone(), p.governance))
49 .unwrap_or_default();
50 registry
51 .providers
52 .retain(|p| p.name.as_str() != spec.name.as_str());
53 registry.providers.push(ProviderEntry {
54 name: spec.name,
55 wire: spec.wire,
56 surface: spec.surface,
57 endpoint: spec.endpoint,
58 api_key_secret: spec.api_key_secret,
59 extra_headers: spec.extra_headers,
60 models,
61 governance,
62 });
63 }
64
65 pub fn remove_provider(registry: &mut ProviderRegistry, name: &ProviderId) -> ConfigResult<()> {
66 let before = registry.providers.len();
67 registry
68 .providers
69 .retain(|p| p.name.as_str() != name.as_str());
70 if registry.providers.len() == before {
71 return Err(ConfigError::ProviderNotFound {
72 name: name.to_string(),
73 });
74 }
75 Ok(())
76 }
77
78 pub fn upsert_model(registry: &mut ProviderRegistry, spec: ModelSpec) -> ConfigResult<()> {
79 let provider = registry
80 .providers
81 .iter_mut()
82 .find(|p| p.name.as_str() == spec.provider.as_str())
83 .ok_or_else(|| ConfigError::ProviderNotFound {
84 name: spec.provider.to_string(),
85 })?;
86 provider
87 .models
88 .retain(|m| m.id.as_str() != spec.id.as_str());
89 provider.models.push(ProviderModel {
90 id: spec.id,
91 aliases: spec.aliases,
92 upstream_model: spec.upstream_model,
93 pricing: ModelPricing::default(),
94 capabilities: ModelCapabilities::default(),
95 limits: ModelLimits::default(),
96 governance: None,
97 });
98 Ok(())
99 }
100
101 pub fn remove_model(
102 registry: &mut ProviderRegistry,
103 provider_name: &ProviderId,
104 id: &ModelId,
105 ) -> ConfigResult<()> {
106 let provider = registry
107 .providers
108 .iter_mut()
109 .find(|p| p.name.as_str() == provider_name.as_str())
110 .ok_or_else(|| ConfigError::ProviderNotFound {
111 name: provider_name.to_string(),
112 })?;
113 let before = provider.models.len();
114 provider.models.retain(|m| m.id.as_str() != id.as_str());
115 if provider.models.len() == before {
116 return Err(ConfigError::ModelNotFound {
117 id: id.to_string(),
118 provider: provider_name.to_string(),
119 });
120 }
121 Ok(())
122 }
123}