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