1#![expect(
7 clippy::clone_on_ref_ptr,
8 reason = "Arc::clone usage is intentional and ergonomic in this gateway dispatch path"
9)]
10
11pub mod credentials;
12mod error;
13pub mod finalize;
14mod pricing;
15pub mod resolve;
16pub mod stages;
17
18pub use self::error::{
19 DispatchError, GovernanceDenied, GuardForbidden, GuardUnavailable, PolicyDenied,
20 PromptRepairRequired, QuotaExceeded, SafetyBlocked,
21};
22pub(super) use self::finalize::run_response_safety_scan;
23
24use std::sync::Arc;
25
26use anyhow::{Result, anyhow};
27use axum::body::Body;
28use axum::response::Response;
29use bytes::Bytes;
30use systemprompt_database::DbPool;
31use systemprompt_identifiers::UserId;
32use systemprompt_models::services::{GatewayConfig, ProviderRegistry};
33
34use self::finalize::{FinalizeCtx, attach_request_id, finalize};
35use self::pricing::{dispatch_pricing, trace_dispatch};
36use self::resolve::{ResolvedUpstream, resolve_upstream};
37use self::stages::{
38 GovernedDispatch, PreparedDispatch, ScannedDispatch, UpstreamRelay, record_quota_warning,
39};
40use super::audit::{GatewayAudit, GatewayRequestContext};
41use super::policy::{GatewayPolicySpec, PolicyResolver};
42use super::protocol::canonical::CanonicalRequest;
43use super::protocol::inbound::InboundAdapter;
44use super::quota;
45
46pub const REQUEST_ID_HEADER: &str = "x-systemprompt-request-id";
47pub const RECOVERY_COUNT_HEADER: &str = "x-systemprompt-recovery-count";
48
49#[derive(Debug, Clone, Copy)]
50pub struct GatewayService;
51
52#[derive(Debug)]
53pub struct DispatchInputs {
54 pub request: CanonicalRequest,
55 pub raw_body: Bytes,
56 pub ctx: GatewayRequestContext,
57 pub inbound: Arc<dyn InboundAdapter>,
58 pub forward_headers: Vec<(String, String)>,
59 pub identity_headers: Vec<(String, String)>,
60}
61
62impl GatewayService {
63 pub async fn dispatch(
64 config: &GatewayConfig,
65 registry: &ProviderRegistry,
66 db: &DbPool,
67 repos: &super::GatewayRepositories,
68 inputs: DispatchInputs,
69 ) -> Result<Response<Body>, DispatchError> {
70 let DispatchInputs {
71 request,
72 raw_body,
73 ctx,
74 inbound,
75 forward_headers,
76 identity_headers,
77 } = inputs;
78 let (policy, evaluation_session) = dispatch_policy(repos, &ctx).await?;
79 let stream_usage = inbound.wants_stream_usage(&raw_body);
80 let ai_request_id = ctx.ai_request_id.clone();
81 let upstream = resolve_upstream(config, registry, &request, &ai_request_id).await?;
82 let pricing = dispatch_pricing(config, registry, &request, &upstream, evaluation_session)?;
83
84 trace_dispatch(&ctx, &request, &upstream);
85 let audit = open_audit(repos, &ctx, &request, &raw_body, &identity_headers).await?;
86 if evaluation_session {
87 audit
88 .pin_evaluation_pricing(pricing)
89 .map_err(DispatchError::PreAudit)?;
90 }
91
92 if let Some(descriptor) = upstream.route_match_descriptor.as_deref() {
93 audit.set_route_match(descriptor).await;
94 }
95
96 enforce_quota(db, repos, &ctx, &policy, &audit).await?;
97 enforce_request_guards(db, &ctx.user_id, &upstream, &request, &audit).await?;
98
99 let prepared = PreparedDispatch::build(
100 config,
101 &upstream,
102 request,
103 &audit,
104 UpstreamRelay {
105 raw_body: &raw_body,
106 inbound: inbound.as_ref(),
107 },
108 )
109 .await?;
110 let governed = GovernedDispatch::enforce(prepared, db, &ctx, &audit).await?;
111 let scanned =
112 ScannedDispatch::enforce(governed, repos, &ai_request_id, &policy.safety, &audit)
113 .await?;
114
115 let evaluation = scanned.admit_evaluation(repos, &ctx, &pricing).await?;
116 let retry_policy = if evaluation {
117 super::protocol::outbound::retry::RetryPolicy::none()
118 } else {
119 super::protocol::outbound::retry::current_policy()
120 };
121 let outcome = super::protocol::outbound::retry::with_policy(
122 retry_policy,
123 scanned.send(&upstream, &forward_headers, &audit),
124 )
125 .await?;
126
127 let mut response = finalize(
128 outcome,
129 FinalizeCtx {
130 audit: Arc::clone(&audit),
131 db: db.clone(),
132 repos: repos.clone(),
133 ai_request_id: ai_request_id.clone(),
134 policy,
135 inbound,
136 request_model: scanned.request_model().to_owned(),
137 stream_usage,
138 },
139 )
140 .await;
141 stages::recovery::attach_recovery_count(&mut response, scanned.recovery_count());
142 Ok(attach_request_id(response, &ai_request_id))
143 }
144}
145
146async fn dispatch_policy(
147 repos: &super::GatewayRepositories,
148 ctx: &GatewayRequestContext,
149) -> Result<(GatewayPolicySpec, bool), DispatchError> {
150 if ctx.session_id.is_none() {
151 return Err(DispatchError::PreAudit(anyhow!(
152 "gateway dispatch missing conversation binding (session_id)"
153 )));
154 }
155
156 let resolver = PolicyResolver::from_repository(repos.gateway_policies.clone());
157 let policy = resolver.resolve().await;
158 let evaluation_session = super::evaluation::preflight(repos, ctx, &policy)
159 .await
160 .map_err(DispatchError::PreAudit)?;
161 Ok((policy, evaluation_session))
162}
163
164async fn open_audit(
165 repos: &super::GatewayRepositories,
166 ctx: &GatewayRequestContext,
167 request: &CanonicalRequest,
168 raw_body: &Bytes,
169 identity_headers: &[(String, String)],
170) -> Result<Arc<GatewayAudit>, DispatchError> {
171 let audit = Arc::new(GatewayAudit::new(repos, ctx.clone()));
172 if let Err(e) = audit.open(request, raw_body).await {
173 tracing::error!(error = %e, "audit open failed — proceeding without audit row");
174 }
175 if !identity_headers.is_empty() {
176 tracing::info!(
177 ai_request_id = %ctx.ai_request_id,
178 user_id = %ctx.user_id,
179 headers = ?identity_headers,
180 "Gateway consumed client identity headers"
181 );
182 }
183 Ok(audit)
184}
185
186async fn enforce_quota(
187 db: &DbPool,
188 repos: &super::GatewayRepositories,
189 ctx: &GatewayRequestContext,
190 policy: &GatewayPolicySpec,
191 audit: &GatewayAudit,
192) -> Result<(), DispatchError> {
193 let reservation = quota::precheck_and_reserve(
194 db,
195 &repos.quota_buckets,
196 &ctx.user_id,
197 &policy.quota_windows,
198 )
199 .await
200 .map_err(DispatchError::Recorded)?;
201 let Some(decision) = reservation else {
202 return Ok(());
203 };
204 if decision.allow {
205 return Ok(());
206 }
207 if policy.quota_mode.is_warn() {
208 tracing::warn!(
209 ai_request_id = %ctx.ai_request_id,
210 user_id = %ctx.user_id,
211 window_seconds = decision.window_seconds,
212 reason = %decision.message,
213 "Gateway quota window exhausted in warn mode; allowing the request"
214 );
215 record_quota_warning(db, ctx, &decision.message).await;
216 return Ok(());
217 }
218 let msg = decision.message;
219 if let Err(e) = audit.fail(&msg).await {
220 tracing::warn!(error = %e, "quota audit fail failed");
221 }
222 Err(DispatchError::Recorded(
223 QuotaExceeded {
224 message: msg,
225 retry_after_seconds: decision.window_seconds,
226 }
227 .into(),
228 ))
229}
230
231async fn enforce_request_guards(
232 db: &DbPool,
233 user_id: &UserId,
234 upstream: &ResolvedUpstream<'_>,
235 request: &CanonicalRequest,
236 audit: &GatewayAudit,
237) -> Result<(), DispatchError> {
238 let Some(pool) = db.pool() else {
239 return Ok(());
240 };
241 let guard_request = systemprompt_extension::GatewayGuardRequest {
242 user_id: user_id.as_str(),
243 model: &request.model,
244 route_id: Some(upstream.route.id.as_str()),
245 provider: upstream.route.provider.as_str(),
246 streaming: request.stream,
247 };
248 let Err(deny) = systemprompt_extension::run_gateway_guards(&pool, &guard_request).await else {
249 return Ok(());
250 };
251 tracing::warn!(
252 user_id = %user_id,
253 model = %request.model,
254 route_id = %upstream.route.id,
255 kind = ?deny.kind,
256 reason = %deny.message,
257 "Gateway request denied by request guard"
258 );
259 if let Err(e) = audit.fail(&deny.message).await {
260 tracing::warn!(error = %e, "request-guard audit fail failed");
261 }
262 let inner: anyhow::Error = match deny.kind {
263 systemprompt_extension::GatewayDenyKind::Unavailable => GuardUnavailable {
264 message: deny.message,
265 retry_after_seconds: deny.retry_after_seconds,
266 }
267 .into(),
268 systemprompt_extension::GatewayDenyKind::Forbidden => GuardForbidden {
269 message: deny.message,
270 }
271 .into(),
272 _ => QuotaExceeded {
275 message: deny.message,
276 retry_after_seconds: deny.retry_after_seconds,
277 }
278 .into(),
279 };
280 Err(DispatchError::Recorded(inner))
281}