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 abandon;
12pub mod credentials;
13mod error;
14pub mod failover;
15pub mod finalize;
16mod guards;
17mod pricing;
18pub mod resolve;
19pub mod stages;
20
21pub use self::error::{
22 DispatchError, GovernanceDenied, GuardForbidden, GuardUnavailable, PolicyDenied,
23 PromptRepairRequired, QuotaExceeded, SafetyBlocked,
24};
25pub(super) use self::finalize::run_response_safety_scan;
26
27use std::sync::Arc;
28
29use anyhow::{Result, anyhow};
30use axum::body::Body;
31use axum::response::Response;
32use bytes::Bytes;
33use systemprompt_database::DbPool;
34use systemprompt_models::services::{GatewayConfig, ProviderRegistry, QuotaFaultMode};
35
36use self::abandon::AbandonGuard;
37use self::failover::{FailoverSend, send_with_failover};
38use self::finalize::{FinalizeCtx, attach_request_id, finalize};
39use self::guards::{enforce_quota, enforce_request_guards};
40use self::pricing::{dispatch_pricing, trace_dispatch};
41use self::resolve::{ResolvedUpstream, resolve_upstream};
42use self::stages::{GovernedDispatch, PreparedDispatch, ScannedDispatch, UpstreamRelay};
43use super::audit::{GatewayAudit, GatewayRequestContext};
44use super::policy::{GatewayPolicySpec, PolicyResolver};
45use super::protocol::canonical::CanonicalRequest;
46use super::protocol::inbound::InboundAdapter;
47use systemprompt_security::policy::GovernanceEngine;
48
49pub const REQUEST_ID_HEADER: &str = "x-systemprompt-request-id";
50pub const RECOVERY_COUNT_HEADER: &str = "x-systemprompt-recovery-count";
51
52#[derive(Debug, Clone, Copy)]
53pub struct GatewayService;
54
55#[derive(Debug)]
56pub struct DispatchInputs {
57 pub request: CanonicalRequest,
58 pub raw_body: Bytes,
59 pub ctx: GatewayRequestContext,
60 pub inbound: Arc<dyn InboundAdapter>,
61 pub forward_headers: Vec<(String, String)>,
62 pub identity_headers: Vec<(String, String)>,
63 pub governance: Arc<GovernanceEngine>,
64}
65
66impl GatewayService {
67 pub async fn dispatch(
68 config: &GatewayConfig,
69 registry: &ProviderRegistry,
70 db: &DbPool,
71 repos: &super::GatewayRepositories,
72 inputs: DispatchInputs,
73 ) -> Result<Response<Body>, DispatchError> {
74 let DispatchInputs {
75 request,
76 raw_body,
77 ctx,
78 inbound,
79 forward_headers,
80 identity_headers,
81 governance,
82 } = inputs;
83 let policy = dispatch_policy(repos, &ctx, config.quota_fault_mode).await?;
84 let stream_usage = inbound.wants_stream_usage(&raw_body);
85 let ai_request_id = ctx.ai_request_id.clone();
86 let upstream = resolve_upstream(config, registry, &request, &ai_request_id).await?;
87 let pricing = dispatch_pricing(config, registry, &request, &upstream)?;
88
89 trace_dispatch(&ctx, &request, &upstream);
90 let audit = open_audit(repos, &ctx, &request, &raw_body, &identity_headers).await?;
91 let mut guard = AbandonGuard::arm(Arc::clone(&audit));
92 let result = Box::pin(dispatch_opened(OpenedDispatch {
93 config,
94 registry,
95 db,
96 repos,
97 audit,
98 policy,
99 stream_usage,
100 ai_request_id,
101 upstream,
102 pricing,
103 request,
104 raw_body,
105 ctx,
106 inbound,
107 forward_headers,
108 governance,
109 }))
110 .await;
111 guard.disarm();
116 result
117 }
118}
119
120struct OpenedDispatch<'a> {
121 config: &'a GatewayConfig,
122 registry: &'a ProviderRegistry,
123 db: &'a DbPool,
124 repos: &'a super::GatewayRepositories,
125 audit: Arc<GatewayAudit>,
126 policy: GatewayPolicySpec,
127 stream_usage: bool,
128 ai_request_id: systemprompt_identifiers::AiRequestId,
129 upstream: ResolvedUpstream<'a>,
130 pricing: systemprompt_models::services::ModelPricing,
131 request: CanonicalRequest,
132 raw_body: Bytes,
133 ctx: GatewayRequestContext,
134 inbound: Arc<dyn InboundAdapter>,
135 forward_headers: Vec<(String, String)>,
136 governance: Arc<GovernanceEngine>,
137}
138
139async fn dispatch_opened(opened: OpenedDispatch<'_>) -> Result<Response<Body>, DispatchError> {
140 let OpenedDispatch {
141 config,
142 registry,
143 db,
144 repos,
145 audit,
146 policy,
147 stream_usage,
148 ai_request_id,
149 upstream,
150 pricing,
151 request,
152 raw_body,
153 ctx,
154 inbound,
155 forward_headers,
156 governance,
157 } = opened;
158 audit
159 .pin_pricing(pricing)
160 .map_err(DispatchError::PreAudit)?;
161
162 if let Some(descriptor) = upstream.route_match_descriptor.as_deref() {
163 audit.set_route_match(descriptor).await;
164 }
165
166 enforce_quota(db, repos, &policy, &audit, config.quota_fault_mode).await?;
167 enforce_request_guards(db, &ctx.user_id, &upstream, &request, &audit).await?;
168
169 let prepared = PreparedDispatch::build(
170 config,
171 &upstream,
172 request,
173 &audit,
174 UpstreamRelay {
175 raw_body: &raw_body,
176 inbound: inbound.as_ref(),
177 },
178 )
179 .await?;
180 let governed = GovernedDispatch::enforce(prepared, db, &ctx, &audit, &governance).await?;
181 let mut scanned =
182 ScannedDispatch::enforce(governed, repos, &ai_request_id, &policy.safety, &audit).await?;
183
184 let outcome = send_with_failover(
185 &mut scanned,
186 FailoverSend {
187 registry,
188 primary: &upstream,
189 forward_headers: &forward_headers,
190 ai_request_id: &ai_request_id,
191 audit: &audit,
192 },
193 )
194 .await?;
195
196 let mut response = finalize(
197 outcome,
198 FinalizeCtx {
199 audit: Arc::clone(&audit),
200 db: db.clone(),
201 repos: repos.clone(),
202 ai_request_id: ai_request_id.clone(),
203 policy,
204 quota_fault_mode: config.quota_fault_mode,
205 inbound,
206 request_model: scanned.request_model().to_owned(),
207 stream_usage,
208 },
209 )
210 .await;
211 stages::recovery::attach_recovery_count(&mut response, scanned.recovery_count());
212 Ok(attach_request_id(response, &ai_request_id))
213}
214
215async fn dispatch_policy(
216 repos: &super::GatewayRepositories,
217 ctx: &GatewayRequestContext,
218 fault_mode: QuotaFaultMode,
219) -> Result<GatewayPolicySpec, DispatchError> {
220 if ctx.session_id.is_none() {
221 return Err(DispatchError::PreAudit(anyhow!(
222 "gateway dispatch missing authenticated session (session_id)"
223 )));
224 }
225
226 let resolver = PolicyResolver::from_repository(repos.gateway_policies.clone());
227 let policy = resolver
228 .resolve(fault_mode)
229 .await
230 .map_err(|e| DispatchError::PreAudit(anyhow!(PolicyDenied(e.to_string()))))?;
231 Ok(policy)
232}
233
234async fn open_audit(
235 repos: &super::GatewayRepositories,
236 ctx: &GatewayRequestContext,
237 request: &CanonicalRequest,
238 raw_body: &Bytes,
239 identity_headers: &[(String, String)],
240) -> Result<Arc<GatewayAudit>, DispatchError> {
241 let audit = Arc::new(GatewayAudit::new(repos, ctx.clone()));
242 if let Err(error) = audit.open(request, raw_body).await {
243 if let Err(settlement_error) = audit
244 .fail("Gateway admission failed before provider dispatch")
245 .await
246 {
247 tracing::error!(%settlement_error, "Could not record failed gateway admission");
248 }
249 return Err(DispatchError::PreAudit(error));
250 }
251 if !identity_headers.is_empty() {
252 tracing::info!(
253 ai_request_id = %ctx.ai_request_id,
254 user_id = %ctx.user_id,
255 headers = ?identity_headers,
256 "Gateway consumed client identity headers"
257 );
258 }
259 Ok(audit)
260}