Skip to main content

systemprompt_models/services/gateway/config/
validate.rs

1//! Cross-checks of gateway routes against the provider registry.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use crate::services::ai::ModelPricing;
7use crate::services::gateway::config::GatewayConfig;
8use crate::services::gateway::error::{GatewayProfileError, GatewayResult};
9use crate::services::gateway::route::GatewayRoute;
10use crate::services::providers::ProviderRegistry;
11
12impl GatewayConfig {
13    #[must_use]
14    pub fn unresolved_secret_refs(&self, has_secret: impl Fn(&str) -> bool) -> Vec<String> {
15        let mut unresolved = Vec::new();
16        if let Some(name) = self
17            .bridge_releases
18            .as_ref()
19            .and_then(|spec| spec.token_secret.as_deref())
20            && !has_secret(name)
21        {
22            unresolved.push(format!("bridge_releases.token_secret={name}"));
23        }
24        unresolved
25    }
26
27    pub fn validate(&self, registry: &ProviderRegistry) -> GatewayResult<()> {
28        let mut route_ids: std::collections::HashSet<&str> =
29            std::collections::HashSet::with_capacity(self.routes.len());
30        for route in &self.routes {
31            if !route_ids.insert(route.id.as_str()) {
32                return Err(GatewayProfileError::DuplicateRouteId {
33                    id: route.id.as_str().to_owned(),
34                });
35            }
36        }
37        if let Some(provider) = self.default_provider.as_ref()
38            && registry.find_provider(provider.as_str()).is_none()
39        {
40            return Err(GatewayProfileError::DefaultProviderNotInRegistry {
41                provider: provider.as_str().to_owned(),
42            });
43        }
44        for route in &self.routes {
45            if registry.find_provider(route.provider.as_str()).is_none() {
46                return Err(GatewayProfileError::RouteProviderNotInRegistry {
47                    route: route.model_pattern.clone(),
48                    provider: route.provider.as_str().to_owned(),
49                });
50            }
51            if let Some(when) = route.when.as_ref() {
52                when.validate()?;
53            }
54            self.validate_route_pricing(registry, route)?;
55            validate_route_governance(registry, route)?;
56            self.validate_route_fallback(registry, route)?;
57        }
58        for rule in &self.system_prompt_overrides {
59            rule.validate()?;
60            if let Some(provider) = rule.provider.as_ref()
61                && registry.find_provider(provider.as_str()).is_none()
62            {
63                return Err(GatewayProfileError::OverrideProviderNotInRegistry {
64                    provider: provider.as_str().to_owned(),
65                });
66            }
67        }
68        Ok(())
69    }
70
71    // Why: the fallback is validated as the route the fallback provider will
72    // actually serve, so a failover can never land on a model the primary
73    // route's pricing and governance checks would have refused at boot.
74    fn validate_route_fallback(
75        &self,
76        registry: &ProviderRegistry,
77        route: &GatewayRoute,
78    ) -> GatewayResult<()> {
79        let route_id = route.id.as_str().to_owned();
80        let Some(view) = route.fallback_view() else {
81            if route.fallback_upstream_model.is_some() {
82                return Err(GatewayProfileError::RouteFallbackModelWithoutProvider {
83                    route: route_id,
84                });
85            }
86            return Ok(());
87        };
88        if view.provider == route.provider {
89            return Err(GatewayProfileError::RouteFallbackIsPrimary {
90                route: route_id,
91                provider: view.provider.as_str().to_owned(),
92            });
93        }
94        if view.resolve(registry).is_none() {
95            return Err(GatewayProfileError::RouteFallbackProviderNotInRegistry {
96                route: route_id,
97                provider: view.provider.as_str().to_owned(),
98            });
99        }
100        self.validate_route_pricing(registry, &view)?;
101        validate_route_governance(registry, &view)
102    }
103
104    fn validate_route_pricing(
105        &self,
106        registry: &ProviderRegistry,
107        route: &GatewayRoute,
108    ) -> GatewayResult<()> {
109        if !self.enabled {
110            return Ok(());
111        }
112        let route_id = route.id.as_str().to_owned();
113        let provider = route.provider.as_str().to_owned();
114        if let Some(pricing) = route.pricing {
115            if !pricing.is_billable() {
116                return Err(GatewayProfileError::RouteModelUnpriced {
117                    route: route_id,
118                    model: route.model_pattern.clone(),
119                });
120            }
121            return check_cache_rate(&pricing, &route_id, &provider, &route.model_pattern);
122        }
123        let Some(entry) = route.resolve(registry) else {
124            return Ok(());
125        };
126        if let Some(upstream) = route.upstream_model.as_deref() {
127            return match entry.find_model(upstream) {
128                Some(model) if model.pricing.is_billable() => {
129                    check_cache_rate(&model.pricing, &route_id, &provider, model.id.as_str())
130                },
131                Some(model) => Err(GatewayProfileError::RouteModelUnpriced {
132                    route: route_id,
133                    model: model.id.as_str().to_owned(),
134                }),
135                None => Err(GatewayProfileError::RouteReachesNoPricedModel {
136                    route: route_id,
137                    pattern: route.model_pattern.clone(),
138                    provider: route.provider.as_str().to_owned(),
139                }),
140            };
141        }
142        let mut reached = 0usize;
143        for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
144            reached += 1;
145            if !model.pricing.is_billable() {
146                return Err(GatewayProfileError::RouteModelUnpriced {
147                    route: route_id,
148                    model: model.id.as_str().to_owned(),
149                });
150            }
151            check_cache_rate(&model.pricing, &route_id, &provider, model.id.as_str())?;
152        }
153        if reached == 0 {
154            return Err(GatewayProfileError::RouteReachesNoPricedModel {
155                route: route_id,
156                pattern: route.model_pattern.clone(),
157                provider: route.provider.as_str().to_owned(),
158            });
159        }
160        Ok(())
161    }
162}
163
164fn check_cache_rate(
165    pricing: &ModelPricing,
166    route: &str,
167    provider: &str,
168    model: &str,
169) -> GatewayResult<()> {
170    if pricing.input_per_million <= 0.0 && pricing.output_per_million <= 0.0 {
171        return Ok(());
172    }
173    if pricing.declares_cache_rate() {
174        Ok(())
175    } else {
176        Err(GatewayProfileError::RouteModelCacheRateUndeclared {
177            route: route.to_owned(),
178            provider: provider.to_owned(),
179            model: model.to_owned(),
180        })
181    }
182}
183
184fn validate_route_governance(
185    registry: &ProviderRegistry,
186    route: &GatewayRoute,
187) -> GatewayResult<()> {
188    let Some(requires) = route.requires.as_ref() else {
189        return Ok(());
190    };
191    if requires.declared().is_empty() {
192        return Ok(());
193    }
194    let Some(entry) = route.resolve(registry) else {
195        return Ok(());
196    };
197    let check = |model_id: &str| -> GatewayResult<()> {
198        let unmet = requires.unmet(entry.effective_governance(model_id));
199        if unmet.is_empty() {
200            Ok(())
201        } else {
202            Err(GatewayProfileError::RouteGovernanceUnsatisfied {
203                route: route.id.as_str().to_owned(),
204                model: model_id.to_owned(),
205                requirements: unmet.join(","),
206            })
207        }
208    };
209    if let Some(upstream) = route.upstream_model.as_deref() {
210        return check(upstream);
211    }
212    for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
213        check(model.id.as_str())?;
214    }
215    Ok(())
216}