Skip to main content

systemprompt_config/services/
provider_catalog.rs

1//! Provider-registry mutations behind the `admin config catalog` surface.
2//!
3//! [`ProviderCatalogService`] edits the typed
4//! [`ProviderRegistry`] on a profile: declaring or removing upstream providers
5//! and the models each provider serves. Upserts replace an entry in place;
6//! a provider upsert preserves the existing model catalog so connectivity can
7//! be re-declared without re-listing models. Callers revalidate and persist
8//! the profile after a successful mutation.
9//!
10//! Copyright (c) systemprompt.io — Business Source License 1.1.
11//! See <https://systemprompt.io> for licensing details.
12
13use 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}