1#![expect(
7 clippy::clone_on_ref_ptr,
8 reason = "Arc::clone usage is intentional and ergonomic in this gateway dispatch path"
9)]
10
11mod finalize;
12mod resolve;
13
14pub(super) use self::finalize::run_response_safety_scan;
15
16#[cfg(feature = "test-api")]
17pub mod test_api {
18 pub use super::blocks_at_phase;
19 pub use super::finalize::{apply_system_prompt_override, attach_request_id, dedupe_findings};
20}
21
22use std::sync::Arc;
23
24use anyhow::{Result, anyhow};
25use axum::body::Body;
26use axum::response::Response;
27use bytes::Bytes;
28use systemprompt_ai::{PHASE_REQUEST, PHASE_REQUEST_HISTORY, SafetyConfig, SafetyHistoryMode};
29use systemprompt_database::DbPool;
30use systemprompt_identifiers::{AiRequestId, UserId};
31use systemprompt_models::profile::{GatewayConfig, ProviderRegistry};
32
33use self::finalize::{
34 FinalizeCtx, apply_system_prompt_override, attach_request_id, finalize, run_request_safety_scan,
35};
36use self::resolve::{ResolvedUpstream, resolve_upstream};
37use super::audit::{GatewayAudit, GatewayRequestContext};
38use super::policy::{PolicyResolver, QuotaWindow};
39use super::protocol::canonical::CanonicalRequest;
40use super::protocol::inbound::InboundAdapter;
41use super::protocol::outbound::{OutboundCtx, OutboundOutcome};
42use super::quota;
43
44pub const REQUEST_ID_HEADER: &str = "x-systemprompt-request-id";
45
46#[derive(Debug, Clone, Copy)]
47pub struct GatewayService;
48
49#[derive(Debug)]
50pub struct DispatchInputs {
51 pub request: CanonicalRequest,
52 pub raw_body: Bytes,
53 pub ctx: GatewayRequestContext,
54 pub inbound: Arc<dyn InboundAdapter>,
55}
56
57#[derive(Debug, thiserror::Error)]
58pub enum DispatchError {
59 #[error(transparent)]
60 PreAudit(anyhow::Error),
61 #[error(transparent)]
62 Recorded(anyhow::Error),
63}
64
65#[derive(Debug, thiserror::Error)]
66#[error("{0}")]
67pub struct PolicyDenied(pub String);
68
69#[derive(Debug, thiserror::Error)]
70#[error("{message}")]
71pub struct QuotaExceeded {
72 pub message: String,
73 pub retry_after_seconds: i32,
74}
75
76#[derive(Debug, thiserror::Error)]
77#[error("{message}")]
78pub struct GuardForbidden {
79 pub message: String,
80}
81
82#[derive(Debug, thiserror::Error)]
83#[error("{message}")]
84pub struct SafetyBlocked {
85 pub category: String,
86 pub message: String,
87}
88
89impl GatewayService {
90 pub async fn dispatch(
91 config: &GatewayConfig,
92 registry: &ProviderRegistry,
93 db: &DbPool,
94 inputs: DispatchInputs,
95 ) -> Result<Response<Body>, DispatchError> {
96 let DispatchInputs {
97 mut request,
98 raw_body,
99 ctx,
100 inbound,
101 } = inputs;
102 if ctx.session_id.is_none() {
103 return Err(DispatchError::PreAudit(anyhow!(
104 "gateway dispatch missing conversation binding (session_id)"
105 )));
106 }
107
108 let ai_request_id = ctx.ai_request_id.clone();
109 let upstream = resolve_upstream(config, registry, &request, &ai_request_id).await?;
110
111 tracing::info!(
112 ai_request_id = %ai_request_id,
113 user_id = %ctx.user_id,
114 model = %request.model,
115 provider = %upstream.route.provider,
116 upstream = %upstream.provider.endpoint,
117 wire_protocol = %ctx.wire_protocol,
118 streaming = request.stream,
119 "Gateway request dispatched"
120 );
121
122 let resolver = PolicyResolver::new(db).map_err(DispatchError::PreAudit)?;
123 let policy = resolver.resolve().await;
124
125 let audit = Arc::new(
126 GatewayAudit::new(db, ctx.clone())
127 .map_err(|e| DispatchError::PreAudit(anyhow!("audit init failed: {e}")))?,
128 );
129
130 if let Err(e) = audit.open(&request, &raw_body).await {
131 tracing::error!(error = %e, "audit open failed — proceeding without audit row");
132 }
133
134 if let Some(descriptor) = upstream.route_match_descriptor.as_deref() {
135 audit.set_route_match(descriptor).await;
136 }
137
138 enforce_quota(db, &ctx.user_id, &policy.quota_windows, &audit).await?;
139 enforce_request_guards(db, &ctx.user_id, &upstream, &request, &audit).await?;
140 enforce_request_safety(db, &ai_request_id, &request, &policy.safety, &audit).await?;
141
142 let outcome = send_to_upstream(config, &upstream, &mut request, &audit).await?;
143
144 let response = finalize(
145 outcome,
146 FinalizeCtx {
147 audit: Arc::clone(&audit),
148 db: db.clone(),
149 ai_request_id: ai_request_id.clone(),
150 policy,
151 inbound,
152 request_model: request.model.clone(),
153 },
154 )
155 .await;
156 Ok(attach_request_id(response, &ai_request_id))
157 }
158}
159
160async fn send_to_upstream(
161 config: &GatewayConfig,
162 upstream: &ResolvedUpstream<'_>,
163 request: &mut CanonicalRequest,
164 audit: &GatewayAudit,
165) -> Result<OutboundOutcome, DispatchError> {
166 let upstream_model = upstream
167 .route
168 .effective_upstream_model(&request.model)
169 .to_owned();
170 if let Some(descriptor) =
171 apply_system_prompt_override(config, &upstream.provider.name, &upstream_model, request)
172 .await
173 {
174 audit.set_system_prompt_override(&descriptor).await;
175 }
176 let model_limits = upstream
177 .provider
178 .find_model(&upstream_model)
179 .map(|m| m.limits);
180 let outbound_ctx = OutboundCtx {
181 route: upstream.route.as_ref(),
182 endpoint: &upstream.provider.endpoint,
183 api_key: upstream.api_key,
184 request,
185 upstream_model: &upstream_model,
186 model_limits,
187 };
188
189 match upstream.adapter.send(outbound_ctx).await {
190 Ok(o) => Ok(o),
191 Err(e) => {
192 audit_upstream_failure(audit, upstream.provider.name.as_str(), &request.model, &e)
193 .await;
194 Err(DispatchError::Recorded(e))
195 },
196 }
197}
198
199async fn enforce_quota(
200 db: &DbPool,
201 user_id: &UserId,
202 quota_windows: &[QuotaWindow],
203 audit: &GatewayAudit,
204) -> Result<(), DispatchError> {
205 let reservation = quota::precheck_and_reserve(db, user_id, quota_windows)
206 .await
207 .map_err(DispatchError::Recorded)?;
208 let Some(decision) = reservation else {
209 return Ok(());
210 };
211 if decision.allow {
212 return Ok(());
213 }
214 let msg = decision.message;
215 if let Err(e) = audit.fail(&msg).await {
216 tracing::warn!(error = %e, "quota audit fail failed");
217 }
218 Err(DispatchError::Recorded(
219 QuotaExceeded {
220 message: msg,
221 retry_after_seconds: decision.window_seconds,
222 }
223 .into(),
224 ))
225}
226
227async fn enforce_request_guards(
228 db: &DbPool,
229 user_id: &UserId,
230 upstream: &ResolvedUpstream<'_>,
231 request: &CanonicalRequest,
232 audit: &GatewayAudit,
233) -> Result<(), DispatchError> {
234 let Some(pool) = db.pool() else {
235 return Ok(());
236 };
237 let guard_request = systemprompt_extension::GatewayGuardRequest {
238 user_id: user_id.as_str(),
239 model: &request.model,
240 route_id: Some(upstream.route.id.as_str()),
241 provider: upstream.route.provider.as_str(),
242 streaming: request.stream,
243 };
244 let Err(deny) = systemprompt_extension::run_gateway_guards(&pool, &guard_request).await else {
245 return Ok(());
246 };
247 tracing::warn!(
248 user_id = %user_id,
249 model = %request.model,
250 route_id = %upstream.route.id,
251 kind = ?deny.kind,
252 reason = %deny.message,
253 "Gateway request denied by request guard"
254 );
255 if let Err(e) = audit.fail(&deny.message).await {
256 tracing::warn!(error = %e, "request-guard audit fail failed");
257 }
258 let inner: anyhow::Error = match deny.kind {
259 systemprompt_extension::GatewayDenyKind::Forbidden => GuardForbidden {
260 message: deny.message,
261 }
262 .into(),
263 systemprompt_extension::GatewayDenyKind::Quota => QuotaExceeded {
264 message: deny.message,
265 retry_after_seconds: deny.retry_after_seconds,
266 }
267 .into(),
268 };
269 Err(DispatchError::Recorded(inner))
270}
271
272async fn enforce_request_safety(
273 db: &DbPool,
274 ai_request_id: &AiRequestId,
275 request: &CanonicalRequest,
276 safety: &SafetyConfig,
277 audit: &GatewayAudit,
278) -> Result<(), DispatchError> {
279 let findings = run_request_safety_scan(db, ai_request_id, request, safety).await;
280 let Some(finding) = findings.iter().find(|f| {
281 safety.block_categories.contains(&f.category) && blocks_at_phase(f.phase, safety.history)
282 }) else {
283 return Ok(());
284 };
285 let msg = format!(
286 "request blocked by safety policy: category '{}'",
287 finding.category
288 );
289 tracing::warn!(
290 ai_request_id = %ai_request_id,
291 category = %finding.category,
292 scanner = %finding.scanner,
293 "Gateway blocked request by safety policy"
294 );
295 if let Err(e) = audit.fail(&msg).await {
296 tracing::warn!(error = %e, "safety-block audit fail failed");
297 }
298 Err(DispatchError::Recorded(
299 SafetyBlocked {
300 category: finding.category.clone(),
301 message: msg,
302 }
303 .into(),
304 ))
305}
306
307pub fn blocks_at_phase(phase: &str, history: SafetyHistoryMode) -> bool {
314 match phase {
315 PHASE_REQUEST => true,
316 PHASE_REQUEST_HISTORY => history == SafetyHistoryMode::Block,
317 _ => false,
318 }
319}
320
321async fn audit_upstream_failure(
322 audit: &GatewayAudit,
323 provider: &str,
324 model: &str,
325 error: &anyhow::Error,
326) {
327 tracing::warn!(
328 provider = %provider,
329 model = %model,
330 error = %error,
331 "gateway upstream call failed"
332 );
333 if let Err(audit_err) = audit.fail(&error.to_string()).await {
334 tracing::warn!(error = %audit_err, "upstream audit fail failed");
335 }
336}