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
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}