Skip to main content

systemprompt_models/profile/gateway/
config.rs

1//! Gateway configuration: on-disk spec and resolved runtime form.
2//!
3//! [`GatewayConfigSpec`] is the serde shape accepted under `gateway:` in a
4//! profile; [`GatewayConfig`] is its runtime projection. Routes carry no
5//! embedded provider catalog — every route resolves its provider against
6//! `profile.providers` ([`ProviderRegistry`]) at use time.
7//!
8//! Copyright (c) systemprompt.io — Business Source License 1.1.
9//! See <https://systemprompt.io> for licensing details.
10
11use std::borrow::Cow;
12use std::collections::HashMap;
13
14use serde::{Deserialize, Serialize};
15use systemprompt_identifiers::{ProviderId, RouteId};
16
17use super::super::providers::ProviderRegistry;
18use super::error::{GatewayProfileError, GatewayResult};
19use super::override_rule::SystemPromptRule;
20use super::route::GatewayRoute;
21use crate::wire::canonical::CanonicalRequest;
22
23pub(crate) const DEFAULT_ROUTE_PATTERN: &str = "*";
24
25#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
26#[serde(deny_unknown_fields)]
27pub struct GatewayConfigSpec {
28    #[serde(default)]
29    pub enabled: bool,
30    #[serde(default)]
31    pub routes: Vec<GatewayRoute>,
32    #[serde(default, skip_serializing_if = "Option::is_none")]
33    pub default_provider: Option<ProviderId>,
34    #[serde(default)]
35    pub allow_unlisted_models: bool,
36    #[serde(default = "default_auth_scheme")]
37    pub auth_scheme: String,
38    #[serde(default = "default_inference_path_prefix")]
39    pub inference_path_prefix: String,
40    #[serde(default, skip_serializing_if = "Vec::is_empty")]
41    pub system_prompt_overrides: Vec<SystemPromptRule>,
42    #[serde(default, skip_serializing_if = "Option::is_none")]
43    pub bridge_releases: Option<BridgeReleasesSpec>,
44}
45
46/// Release feed for the desktop bridge self-updater.
47///
48/// The bridge cannot reach these assets itself — the repository is private —
49/// so the gateway resolves and proxies them. Keeping the resolution here is
50/// also what makes staged rollouts a config change rather than a client
51/// release.
52#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
53#[serde(deny_unknown_fields)]
54pub struct BridgeReleasesSpec {
55    pub repo: String,
56    #[serde(default, skip_serializing_if = "Option::is_none")]
57    pub token_env: Option<String>,
58    #[serde(default = "default_tag_prefix")]
59    pub tag_prefix: String,
60    #[serde(default, skip_serializing_if = "Option::is_none")]
61    pub pinned_version: Option<String>,
62    #[serde(default)]
63    pub assets: std::collections::BTreeMap<String, String>,
64}
65
66fn default_tag_prefix() -> String {
67    "bridge-v".to_owned()
68}
69
70impl Default for GatewayConfigSpec {
71    fn default() -> Self {
72        Self {
73            enabled: false,
74            routes: Vec::new(),
75            default_provider: None,
76            allow_unlisted_models: false,
77            auth_scheme: default_auth_scheme(),
78            inference_path_prefix: default_inference_path_prefix(),
79            system_prompt_overrides: Vec::new(),
80            bridge_releases: None,
81        }
82    }
83}
84
85fn default_auth_scheme() -> String {
86    "bearer".to_owned()
87}
88
89fn default_inference_path_prefix() -> String {
90    "/v1".to_owned()
91}
92
93impl GatewayConfigSpec {
94    #[must_use]
95    pub fn resolve(self) -> GatewayConfig {
96        let Self {
97            enabled,
98            routes,
99            default_provider,
100            allow_unlisted_models,
101            auth_scheme,
102            inference_path_prefix,
103            system_prompt_overrides,
104            bridge_releases,
105        } = self;
106
107        GatewayConfig {
108            enabled,
109            routes,
110            default_provider,
111            allow_unlisted_models,
112            auth_scheme,
113            inference_path_prefix,
114            system_prompt_overrides,
115            bridge_releases,
116        }
117    }
118}
119
120/// Runtime gateway configuration: the post-resolution shape every non-loader
121/// caller sees.
122///
123/// Not `Deserialize`: the only legal construction paths are
124/// [`GatewayConfigSpec::resolve`] for the production loader and direct
125/// struct-literal construction in tests.
126#[derive(Debug, Clone)]
127pub struct GatewayConfig {
128    pub enabled: bool,
129    pub routes: Vec<GatewayRoute>,
130    pub default_provider: Option<ProviderId>,
131    pub allow_unlisted_models: bool,
132    pub auth_scheme: String,
133    pub inference_path_prefix: String,
134    pub system_prompt_overrides: Vec<SystemPromptRule>,
135    pub bridge_releases: Option<BridgeReleasesSpec>,
136}
137
138impl Default for GatewayConfig {
139    fn default() -> Self {
140        Self {
141            enabled: false,
142            routes: Vec::new(),
143            default_provider: None,
144            allow_unlisted_models: false,
145            auth_scheme: default_auth_scheme(),
146            inference_path_prefix: default_inference_path_prefix(),
147            system_prompt_overrides: Vec::new(),
148            bridge_releases: None,
149        }
150    }
151}
152
153impl GatewayConfig {
154    pub fn find_route(&self, model: &str) -> Option<&GatewayRoute> {
155        self.routes.iter().find(|route| route.matches(model))
156    }
157
158    pub fn candidate_routes<'a>(
159        &'a self,
160        registry: &ProviderRegistry,
161    ) -> impl Iterator<Item = Cow<'a, GatewayRoute>> {
162        self.routes
163            .iter()
164            .map(Cow::Borrowed)
165            .chain(self.synthesize_default_route(registry).map(Cow::Owned))
166    }
167
168    #[must_use]
169    pub fn resolve_route<'a>(
170        &'a self,
171        registry: &ProviderRegistry,
172        request: &CanonicalRequest,
173    ) -> Option<Cow<'a, GatewayRoute>> {
174        self.candidate_routes(registry)
175            .find(|route| route.matches_request(request))
176    }
177
178    #[must_use]
179    pub fn dispatchable_route_ids(&self, registry: &ProviderRegistry) -> Vec<RouteId> {
180        let mut ids: Vec<RouteId> = Vec::new();
181        let mut seen: std::collections::HashSet<RouteId> = std::collections::HashSet::new();
182        for route in self.candidate_routes(registry) {
183            let mut route = route.into_owned();
184            route.ensure_id();
185            if seen.insert(route.id.clone()) {
186                ids.push(route.id);
187            }
188        }
189        ids
190    }
191
192    fn synthesize_default_route(&self, registry: &ProviderRegistry) -> Option<GatewayRoute> {
193        let provider = self.default_provider.as_ref()?;
194        registry.find_provider(provider.as_str())?;
195        let mut route = GatewayRoute {
196            id: RouteId::new(""),
197            model_pattern: DEFAULT_ROUTE_PATTERN.to_owned(),
198            provider: provider.clone(),
199            upstream_model: None,
200            extra_headers: HashMap::new(),
201            pricing: None,
202            when: None,
203        };
204        route.ensure_id();
205        Some(route)
206    }
207
208    #[must_use]
209    pub fn is_model_exposed(&self, registry: &ProviderRegistry, model: &str) -> bool {
210        if self.find_route(model).is_some() || registry.contains_model(model) {
211            return true;
212        }
213        if self.default_provider.is_some() && self.allow_unlisted_models {
214            tracing::warn!(
215                model,
216                "gateway forwarding an unlisted model to default_provider \
217                 (allow_unlisted_models=true): open allowlist posture"
218            );
219            return true;
220        }
221        false
222    }
223
224    pub fn validate(&self, registry: &ProviderRegistry) -> GatewayResult<()> {
225        let mut route_ids: std::collections::HashSet<&str> =
226            std::collections::HashSet::with_capacity(self.routes.len());
227        for route in &self.routes {
228            if !route_ids.insert(route.id.as_str()) {
229                return Err(GatewayProfileError::DuplicateRouteId {
230                    id: route.id.as_str().to_owned(),
231                });
232            }
233        }
234        if let Some(provider) = self.default_provider.as_ref()
235            && registry.find_provider(provider.as_str()).is_none()
236        {
237            return Err(GatewayProfileError::DefaultProviderNotInRegistry {
238                provider: provider.as_str().to_owned(),
239            });
240        }
241        for route in &self.routes {
242            if registry.find_provider(route.provider.as_str()).is_none() {
243                return Err(GatewayProfileError::RouteProviderNotInRegistry {
244                    route: route.model_pattern.clone(),
245                    provider: route.provider.as_str().to_owned(),
246                });
247            }
248            if let Some(when) = route.when.as_ref() {
249                when.validate()?;
250            }
251            self.validate_route_pricing(registry, route)?;
252        }
253        for rule in &self.system_prompt_overrides {
254            rule.validate()?;
255            if let Some(provider) = rule.provider.as_ref()
256                && registry.find_provider(provider.as_str()).is_none()
257            {
258                return Err(GatewayProfileError::OverrideProviderNotInRegistry {
259                    provider: provider.as_str().to_owned(),
260                });
261            }
262        }
263        Ok(())
264    }
265
266    fn validate_route_pricing(
267        &self,
268        registry: &ProviderRegistry,
269        route: &GatewayRoute,
270    ) -> GatewayResult<()> {
271        if !self.enabled {
272            return Ok(());
273        }
274        let route_id = route.id.as_str().to_owned();
275        if let Some(pricing) = route.pricing {
276            return if pricing.is_billable() {
277                Ok(())
278            } else {
279                Err(GatewayProfileError::RouteModelUnpriced {
280                    route: route_id,
281                    model: route.model_pattern.clone(),
282                })
283            };
284        }
285        let Some(entry) = route.resolve(registry) else {
286            return Ok(());
287        };
288        if let Some(upstream) = route.upstream_model.as_deref() {
289            return match entry.find_model(upstream) {
290                Some(model) if model.pricing.is_billable() => Ok(()),
291                Some(model) => Err(GatewayProfileError::RouteModelUnpriced {
292                    route: route_id,
293                    model: model.id.as_str().to_owned(),
294                }),
295                None => Err(GatewayProfileError::RouteReachesNoPricedModel {
296                    route: route_id,
297                    pattern: route.model_pattern.clone(),
298                    provider: route.provider.as_str().to_owned(),
299                }),
300            };
301        }
302        let mut reached = 0usize;
303        for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
304            reached += 1;
305            if !model.pricing.is_billable() {
306                return Err(GatewayProfileError::RouteModelUnpriced {
307                    route: route_id,
308                    model: model.id.as_str().to_owned(),
309                });
310            }
311        }
312        if reached == 0 {
313            return Err(GatewayProfileError::RouteReachesNoPricedModel {
314                route: route_id,
315                pattern: route.model_pattern.clone(),
316                provider: route.provider.as_str().to_owned(),
317            });
318        }
319        Ok(())
320    }
321
322    #[must_use]
323    pub fn to_spec(&self) -> GatewayConfigSpec {
324        GatewayConfigSpec {
325            enabled: self.enabled,
326            routes: self.routes.clone(),
327            default_provider: self.default_provider.clone(),
328            allow_unlisted_models: self.allow_unlisted_models,
329            auth_scheme: self.auth_scheme.clone(),
330            inference_path_prefix: self.inference_path_prefix.clone(),
331            system_prompt_overrides: self.system_prompt_overrides.clone(),
332            bridge_releases: self.bridge_releases.clone(),
333        }
334    }
335}