Skip to main content

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