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