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 };
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}