systemprompt_models/profile/gateway/
config.rs1use 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#[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#[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 requires: None,
204 };
205 route.ensure_id();
206 Some(route)
207 }
208
209 #[must_use]
210 pub fn is_model_exposed(&self, registry: &ProviderRegistry, model: &str) -> bool {
211 if self.find_route(model).is_some() || registry.contains_model(model) {
212 return true;
213 }
214 if self.default_provider.is_some() && self.allow_unlisted_models {
215 tracing::warn!(
216 model,
217 "gateway forwarding an unlisted model to default_provider \
218 (allow_unlisted_models=true): open allowlist posture"
219 );
220 return true;
221 }
222 false
223 }
224
225 pub fn validate(&self, registry: &ProviderRegistry) -> GatewayResult<()> {
226 let mut route_ids: std::collections::HashSet<&str> =
227 std::collections::HashSet::with_capacity(self.routes.len());
228 for route in &self.routes {
229 if !route_ids.insert(route.id.as_str()) {
230 return Err(GatewayProfileError::DuplicateRouteId {
231 id: route.id.as_str().to_owned(),
232 });
233 }
234 }
235 if let Some(provider) = self.default_provider.as_ref()
236 && registry.find_provider(provider.as_str()).is_none()
237 {
238 return Err(GatewayProfileError::DefaultProviderNotInRegistry {
239 provider: provider.as_str().to_owned(),
240 });
241 }
242 for route in &self.routes {
243 if registry.find_provider(route.provider.as_str()).is_none() {
244 return Err(GatewayProfileError::RouteProviderNotInRegistry {
245 route: route.model_pattern.clone(),
246 provider: route.provider.as_str().to_owned(),
247 });
248 }
249 if let Some(when) = route.when.as_ref() {
250 when.validate()?;
251 }
252 self.validate_route_pricing(registry, route)?;
253 validate_route_governance(registry, route)?;
254 }
255 for rule in &self.system_prompt_overrides {
256 rule.validate()?;
257 if let Some(provider) = rule.provider.as_ref()
258 && registry.find_provider(provider.as_str()).is_none()
259 {
260 return Err(GatewayProfileError::OverrideProviderNotInRegistry {
261 provider: provider.as_str().to_owned(),
262 });
263 }
264 }
265 Ok(())
266 }
267
268 fn validate_route_pricing(
269 &self,
270 registry: &ProviderRegistry,
271 route: &GatewayRoute,
272 ) -> GatewayResult<()> {
273 if !self.enabled {
274 return Ok(());
275 }
276 let route_id = route.id.as_str().to_owned();
277 if let Some(pricing) = route.pricing {
278 return if pricing.is_billable() {
279 Ok(())
280 } else {
281 Err(GatewayProfileError::RouteModelUnpriced {
282 route: route_id,
283 model: route.model_pattern.clone(),
284 })
285 };
286 }
287 let Some(entry) = route.resolve(registry) else {
288 return Ok(());
289 };
290 if let Some(upstream) = route.upstream_model.as_deref() {
291 return match entry.find_model(upstream) {
292 Some(model) if model.pricing.is_billable() => Ok(()),
293 Some(model) => Err(GatewayProfileError::RouteModelUnpriced {
294 route: route_id,
295 model: model.id.as_str().to_owned(),
296 }),
297 None => Err(GatewayProfileError::RouteReachesNoPricedModel {
298 route: route_id,
299 pattern: route.model_pattern.clone(),
300 provider: route.provider.as_str().to_owned(),
301 }),
302 };
303 }
304 let mut reached = 0usize;
305 for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
306 reached += 1;
307 if !model.pricing.is_billable() {
308 return Err(GatewayProfileError::RouteModelUnpriced {
309 route: route_id,
310 model: model.id.as_str().to_owned(),
311 });
312 }
313 }
314 if reached == 0 {
315 return Err(GatewayProfileError::RouteReachesNoPricedModel {
316 route: route_id,
317 pattern: route.model_pattern.clone(),
318 provider: route.provider.as_str().to_owned(),
319 });
320 }
321 Ok(())
322 }
323
324 #[must_use]
325 pub fn to_spec(&self) -> GatewayConfigSpec {
326 GatewayConfigSpec {
327 enabled: self.enabled,
328 routes: self.routes.clone(),
329 default_provider: self.default_provider.clone(),
330 allow_unlisted_models: self.allow_unlisted_models,
331 auth_scheme: self.auth_scheme.clone(),
332 inference_path_prefix: self.inference_path_prefix.clone(),
333 system_prompt_overrides: self.system_prompt_overrides.clone(),
334 bridge_releases: self.bridge_releases.clone(),
335 }
336 }
337}
338
339fn validate_route_governance(
340 registry: &ProviderRegistry,
341 route: &GatewayRoute,
342) -> GatewayResult<()> {
343 let Some(requires) = route.requires.as_ref() else {
344 return Ok(());
345 };
346 if requires.declared().is_empty() {
347 return Ok(());
348 }
349 let Some(entry) = route.resolve(registry) else {
350 return Ok(());
351 };
352 let check = |model_id: &str| -> GatewayResult<()> {
353 let unmet = requires.unmet(entry.effective_governance(model_id));
354 if unmet.is_empty() {
355 Ok(())
356 } else {
357 Err(GatewayProfileError::RouteGovernanceUnsatisfied {
358 route: route.id.as_str().to_owned(),
359 model: model_id.to_owned(),
360 requirements: unmet.join(","),
361 })
362 }
363 };
364 if let Some(upstream) = route.upstream_model.as_deref() {
365 return check(upstream);
366 }
367 for model in entry.models.iter().filter(|m| route.matches(m.id.as_str())) {
368 check(model.id.as_str())?;
369 }
370 Ok(())
371}