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
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
286/// Whether a finding raised at `phase` may deny the request.
287///
288/// A blocked category found in an earlier turn would otherwise deny every
289/// remaining turn of the conversation, including the turns that carry nothing
290/// objectionable — and a tool call the policy layer already refused is replayed
291/// into the scan surface for the rest of the session.
292pub 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}