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 abandon;
12pub mod credentials;
13mod error;
14pub mod failover;
15pub mod finalize;
16mod guards;
17mod pricing;
18pub mod resolve;
19pub mod stages;
20
21pub use self::error::{
22    DispatchError, GovernanceDenied, GuardForbidden, GuardUnavailable, PolicyDenied,
23    PromptRepairRequired, QuotaExceeded, SafetyBlocked,
24};
25pub(super) use self::finalize::run_response_safety_scan;
26
27use std::sync::Arc;
28
29use anyhow::{Result, anyhow};
30use axum::body::Body;
31use axum::response::Response;
32use bytes::Bytes;
33use systemprompt_database::DbPool;
34use systemprompt_models::services::{GatewayConfig, ProviderRegistry, QuotaFaultMode};
35
36use self::abandon::AbandonGuard;
37use self::failover::{FailoverSend, send_with_failover};
38use self::finalize::{FinalizeCtx, attach_request_id, finalize};
39use self::guards::{enforce_quota, enforce_request_guards};
40use self::pricing::{dispatch_pricing, trace_dispatch};
41use self::resolve::{ResolvedUpstream, resolve_upstream};
42use self::stages::{GovernedDispatch, PreparedDispatch, ScannedDispatch, UpstreamRelay};
43use super::audit::{GatewayAudit, GatewayRequestContext};
44use super::policy::{GatewayPolicySpec, PolicyResolver};
45use super::protocol::canonical::CanonicalRequest;
46use super::protocol::inbound::InboundAdapter;
47use systemprompt_security::policy::GovernanceEngine;
48
49pub const REQUEST_ID_HEADER: &str = "x-systemprompt-request-id";
50pub const RECOVERY_COUNT_HEADER: &str = "x-systemprompt-recovery-count";
51
52#[derive(Debug, Clone, Copy)]
53pub struct GatewayService;
54
55#[derive(Debug)]
56pub struct DispatchInputs {
57    pub request: CanonicalRequest,
58    pub raw_body: Bytes,
59    pub ctx: GatewayRequestContext,
60    pub inbound: Arc<dyn InboundAdapter>,
61    pub forward_headers: Vec<(String, String)>,
62    pub identity_headers: Vec<(String, String)>,
63    pub governance: Arc<GovernanceEngine>,
64}
65
66impl GatewayService {
67    pub async fn dispatch(
68        config: &GatewayConfig,
69        registry: &ProviderRegistry,
70        db: &DbPool,
71        repos: &super::GatewayRepositories,
72        inputs: DispatchInputs,
73    ) -> Result<Response<Body>, DispatchError> {
74        let DispatchInputs {
75            request,
76            raw_body,
77            ctx,
78            inbound,
79            forward_headers,
80            identity_headers,
81            governance,
82        } = inputs;
83        let policy = dispatch_policy(repos, &ctx, config.quota_fault_mode).await?;
84        let stream_usage = inbound.wants_stream_usage(&raw_body);
85        let ai_request_id = ctx.ai_request_id.clone();
86        let upstream = resolve_upstream(config, registry, &request, &ai_request_id).await?;
87        let pricing = dispatch_pricing(config, registry, &request, &upstream)?;
88
89        trace_dispatch(&ctx, &request, &upstream);
90        let audit = open_audit(repos, &ctx, &request, &raw_body, &identity_headers).await?;
91        let mut guard = AbandonGuard::arm(Arc::clone(&audit));
92        let result = Box::pin(dispatch_opened(OpenedDispatch {
93            config,
94            registry,
95            db,
96            repos,
97            audit,
98            policy,
99            stream_usage,
100            ai_request_id,
101            upstream,
102            pricing,
103            request,
104            raw_body,
105            ctx,
106            inbound,
107            forward_headers,
108            governance,
109        }))
110        .await;
111        // Why: every `Err` from the opened dispatch has already recorded itself
112        // on the audit row, and an `Ok` has handed the row to its completion
113        // task or stream tap. The guard is for the third outcome — the future
114        // being dropped before it returns either.
115        guard.disarm();
116        result
117    }
118}
119
120struct OpenedDispatch<'a> {
121    config: &'a GatewayConfig,
122    registry: &'a ProviderRegistry,
123    db: &'a DbPool,
124    repos: &'a super::GatewayRepositories,
125    audit: Arc<GatewayAudit>,
126    policy: GatewayPolicySpec,
127    stream_usage: bool,
128    ai_request_id: systemprompt_identifiers::AiRequestId,
129    upstream: ResolvedUpstream<'a>,
130    pricing: systemprompt_models::services::ModelPricing,
131    request: CanonicalRequest,
132    raw_body: Bytes,
133    ctx: GatewayRequestContext,
134    inbound: Arc<dyn InboundAdapter>,
135    forward_headers: Vec<(String, String)>,
136    governance: Arc<GovernanceEngine>,
137}
138
139async fn dispatch_opened(opened: OpenedDispatch<'_>) -> Result<Response<Body>, DispatchError> {
140    let OpenedDispatch {
141        config,
142        registry,
143        db,
144        repos,
145        audit,
146        policy,
147        stream_usage,
148        ai_request_id,
149        upstream,
150        pricing,
151        request,
152        raw_body,
153        ctx,
154        inbound,
155        forward_headers,
156        governance,
157    } = opened;
158    audit
159        .pin_pricing(pricing)
160        .map_err(DispatchError::PreAudit)?;
161
162    if let Some(descriptor) = upstream.route_match_descriptor.as_deref() {
163        audit.set_route_match(descriptor).await;
164    }
165
166    enforce_quota(db, repos, &policy, &audit, config.quota_fault_mode).await?;
167    enforce_request_guards(db, &ctx.user_id, &upstream, &request, &audit).await?;
168
169    let prepared = PreparedDispatch::build(
170        config,
171        &upstream,
172        request,
173        &audit,
174        UpstreamRelay {
175            raw_body: &raw_body,
176            inbound: inbound.as_ref(),
177        },
178    )
179    .await?;
180    let governed = GovernedDispatch::enforce(prepared, db, &ctx, &audit, &governance).await?;
181    let mut scanned =
182        ScannedDispatch::enforce(governed, repos, &ai_request_id, &policy.safety, &audit).await?;
183
184    let outcome = send_with_failover(
185        &mut scanned,
186        FailoverSend {
187            registry,
188            primary: &upstream,
189            forward_headers: &forward_headers,
190            ai_request_id: &ai_request_id,
191            audit: &audit,
192        },
193    )
194    .await?;
195
196    let mut response = finalize(
197        outcome,
198        FinalizeCtx {
199            audit: Arc::clone(&audit),
200            db: db.clone(),
201            repos: repos.clone(),
202            ai_request_id: ai_request_id.clone(),
203            policy,
204            quota_fault_mode: config.quota_fault_mode,
205            inbound,
206            request_model: scanned.request_model().to_owned(),
207            stream_usage,
208        },
209    )
210    .await;
211    stages::recovery::attach_recovery_count(&mut response, scanned.recovery_count());
212    Ok(attach_request_id(response, &ai_request_id))
213}
214
215async fn dispatch_policy(
216    repos: &super::GatewayRepositories,
217    ctx: &GatewayRequestContext,
218    fault_mode: QuotaFaultMode,
219) -> Result<GatewayPolicySpec, DispatchError> {
220    if ctx.session_id.is_none() {
221        return Err(DispatchError::PreAudit(anyhow!(
222            "gateway dispatch missing authenticated session (session_id)"
223        )));
224    }
225
226    let resolver = PolicyResolver::from_repository(repos.gateway_policies.clone());
227    let policy = resolver
228        .resolve(fault_mode)
229        .await
230        .map_err(|e| DispatchError::PreAudit(anyhow!(PolicyDenied(e.to_string()))))?;
231    Ok(policy)
232}
233
234async fn open_audit(
235    repos: &super::GatewayRepositories,
236    ctx: &GatewayRequestContext,
237    request: &CanonicalRequest,
238    raw_body: &Bytes,
239    identity_headers: &[(String, String)],
240) -> Result<Arc<GatewayAudit>, DispatchError> {
241    let audit = Arc::new(GatewayAudit::new(repos, ctx.clone()));
242    if let Err(error) = audit.open(request, raw_body).await {
243        if let Err(settlement_error) = audit
244            .fail("Gateway admission failed before provider dispatch")
245            .await
246        {
247            tracing::error!(%settlement_error, "Could not record failed gateway admission");
248        }
249        return Err(DispatchError::PreAudit(error));
250    }
251    if !identity_headers.is_empty() {
252        tracing::info!(
253            ai_request_id = %ctx.ai_request_id,
254            user_id = %ctx.user_id,
255            headers = ?identity_headers,
256            "Gateway consumed client identity headers"
257        );
258    }
259    Ok(audit)
260}