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    pub fn validate(&self, registry: &ProviderRegistry) -> GatewayResult<()> {
14        let mut route_ids: std::collections::HashSet<&str> =
15            std::collections::HashSet::with_capacity(self.routes.len());
16        for route in &self.routes {
17            if !route_ids.insert(route.id.as_str()) {
18                return Err(GatewayProfileError::DuplicateRouteId {
19                    id: route.id.as_str().to_owned(),
20                });
21            }
22        }
23        if let Some(provider) = self.default_provider.as_ref()
24            && registry.find_provider(provider.as_str()).is_none()
25        {
26            return Err(GatewayProfileError::DefaultProviderNotInRegistry {
27                provider: provider.as_str().to_owned(),
28            });
29        }
30        for route in &self.routes {
31            if registry.find_provider(route.provider.as_str()).is_none() {
32                return Err(GatewayProfileError::RouteProviderNotInRegistry {
33                    route: route.model_pattern.clone(),
34                    provider: route.provider.as_str().to_owned(),
35                });
36            }
37            if let Some(when) = route.when.as_ref() {
38                when.validate()?;
39            }
40            self.validate_route_pricing(registry, route)?;
41            validate_route_governance(registry, route)?;
42        }
43        for rule in &self.system_prompt_overrides {
44            rule.validate()?;
45            if let Some(provider) = rule.provider.as_ref()
46                && registry.find_provider(provider.as_str()).is_none()
47            {
48                return Err(GatewayProfileError::OverrideProviderNotInRegistry {
49                    provider: provider.as_str().to_owned(),
50                });
51            }
52        }
53        Ok(())
54    }
55
56    fn validate_route_pricing(
57        &self,
58        registry: &ProviderRegistry,
59        route: &GatewayRoute,
60    ) -> GatewayResult<()> {
61        if !self.enabled {
62            return Ok(());
63        }
64        let route_id = route.id.as_str().to_owned();
65        let provider = route.provider.as_str().to_owned();
66        if let Some(pricing) = route.pricing {
67            if !pricing.is_billable() {
68                return Err(GatewayProfileError::RouteModelUnpriced {
69                    route: route_id,
70                    model: route.model_pattern.clone(),
71                });
72            }
73            return check_cache_rate(&pricing, &route_id, &provider, &route.model_pattern);
74        }
75        let Some(entry) = route.resolve(registry) else {
76            return Ok(());
77        };
78        if let Some(upstream) = route.upstream_model.as_deref() {
79            return match entry.find_model(upstream) {
80                Some(model) if model.pricing.is_billable() => {
81                    check_cache_rate(&model.pricing, &route_id, &provider, model.id.as_str())
82                },
83                Some(model) => Err(GatewayProfileError::RouteModelUnpriced {
84                    route: route_id,
85                    model: model.id.as_str().to_owned(),
86                }),
87                None => Err(GatewayProfileError::RouteReachesNoPricedModel {
88                    route: route_id,
89                    pattern: route.model_pattern.clone(),
90                    provider: route.provider.as_str().to_owned(),
91                }),
92            };
93        }
94        let mut reached = 0usize;
95        for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
96            reached += 1;
97            if !model.pricing.is_billable() {
98                return Err(GatewayProfileError::RouteModelUnpriced {
99                    route: route_id,
100                    model: model.id.as_str().to_owned(),
101                });
102            }
103            check_cache_rate(&model.pricing, &route_id, &provider, model.id.as_str())?;
104        }
105        if reached == 0 {
106            return Err(GatewayProfileError::RouteReachesNoPricedModel {
107                route: route_id,
108                pattern: route.model_pattern.clone(),
109                provider: route.provider.as_str().to_owned(),
110            });
111        }
112        Ok(())
113    }
114}
115
116fn check_cache_rate(
117    pricing: &ModelPricing,
118    route: &str,
119    provider: &str,
120    model: &str,
121) -> GatewayResult<()> {
122    if pricing.input_per_million <= 0.0 && pricing.output_per_million <= 0.0 {
123        return Ok(());
124    }
125    if pricing.declares_cache_rate() {
126        Ok(())
127    } else {
128        Err(GatewayProfileError::RouteModelCacheRateUndeclared {
129            route: route.to_owned(),
130            provider: provider.to_owned(),
131            model: model.to_owned(),
132        })
133    }
134}
135
136fn validate_route_governance(
137    registry: &ProviderRegistry,
138    route: &GatewayRoute,
139) -> GatewayResult<()> {
140    let Some(requires) = route.requires.as_ref() else {
141        return Ok(());
142    };
143    if requires.declared().is_empty() {
144        return Ok(());
145    }
146    let Some(entry) = route.resolve(registry) else {
147        return Ok(());
148    };
149    let check = |model_id: &str| -> GatewayResult<()> {
150        let unmet = requires.unmet(entry.effective_governance(model_id));
151        if unmet.is_empty() {
152            Ok(())
153        } else {
154            Err(GatewayProfileError::RouteGovernanceUnsatisfied {
155                route: route.id.as_str().to_owned(),
156                model: model_id.to_owned(),
157                requirements: unmet.join(","),
158            })
159        }
160    };
161    if let Some(upstream) = route.upstream_model.as_deref() {
162        return check(upstream);
163    }
164    for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
165        check(model.id.as_str())?;
166    }
167    Ok(())
168}