Skip to main content

systemprompt_api/services/gateway/service/
mod.rs

1//! Gateway dispatch entry point: route resolution, policy and quota checks,
2//! upstream send, and response finalization.
3//!
4//! Copyright (c) systemprompt.io — Business Source License 1.1.
5//! See <https://systemprompt.io> for licensing details.
6#![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
307/// Whether a finding raised at `phase` may deny the request.
308///
309/// A blocked category found in an earlier turn would otherwise deny every
310/// remaining turn of the conversation, including the turns that carry nothing
311/// objectionable — and a tool call the policy layer already refused is replayed
312/// into the scan surface for the rest of the session.
313pub 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}