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