mod complete;
pub mod journal;
pub mod message_text;
mod open;
pub mod payload;
use std::sync::{Arc, Mutex};
use std::time::Instant;
use anyhow::Result;
use systemprompt_ai::models::RequestStatus;
use systemprompt_ai::repository::{AiRequestPayloadRepository, AiRequestRepository};
use systemprompt_identifiers::{
AiRequestId, ClientId, ClientSessionId, ContextId, GatewayConversationId, SessionId, TraceId,
UserId,
};
use systemprompt_security::policy::types::AccessScope;
#[derive(Debug, Clone)]
pub struct GatewayAccessLog {
pub method: String,
pub path: String,
pub started: Instant,
}
#[derive(Debug, Clone)]
pub struct GatewayRequestContext {
pub ai_request_id: AiRequestId,
pub user_id: UserId,
pub session_id: Option<SessionId>,
pub context_id: ContextId,
pub gateway_conversation_id: Option<GatewayConversationId>,
pub client_session_id: Option<ClientSessionId>,
pub trace_id: Option<TraceId>,
pub access_scope: AccessScope,
pub client_id: Option<ClientId>,
pub provider: String,
pub model: String,
pub requested_model: Option<String>,
pub max_tokens: Option<u32>,
pub is_streaming: bool,
pub wire_protocol: String,
pub access_log: Option<GatewayAccessLog>,
}
#[expect(
missing_debug_implementations,
reason = "service type holds repository clients that intentionally do not implement Debug"
)]
pub struct GatewayAudit {
settlement: journal::Settlement,
journal_lease: std::sync::OnceLock<std::fs::File>,
pricing_snapshot: std::sync::OnceLock<systemprompt_models::services::ModelPricing>,
requests: Arc<AiRequestRepository>,
payloads: Arc<AiRequestPayloadRepository>,
context_materializer: systemprompt_traits::DynContextMaterializer,
pub ctx: GatewayRequestContext,
served_model: Mutex<Option<String>>,
started_at: Instant,
upstream: Mutex<UpstreamClock>,
}
#[derive(Debug, Default)]
struct UpstreamClock {
started: Option<Instant>,
elapsed_ms: Option<i32>,
}
impl GatewayAudit {
pub fn new(repos: &super::GatewayRepositories, ctx: GatewayRequestContext) -> Self {
Self {
settlement: repos.settlement(),
journal_lease: std::sync::OnceLock::new(),
pricing_snapshot: std::sync::OnceLock::new(),
requests: Arc::clone(&repos.requests),
payloads: Arc::clone(&repos.payloads),
context_materializer: Arc::clone(&repos.context_materializer),
ctx,
served_model: Mutex::new(None),
started_at: Instant::now(),
upstream: Mutex::new(UpstreamClock::default()),
}
}
pub async fn set_served_model(&self, model: &str) {
if model.is_empty() || model == self.ctx.model {
return;
}
if let Ok(mut slot) = self.served_model.lock() {
*slot = Some(model.to_owned());
}
if let Err(e) = self
.requests
.update_model(&self.ctx.ai_request_id, model)
.await
{
tracing::warn!(error = %e, "update_model failed");
}
}
pub async fn set_prepared_body_digest(&self, body: &[u8]) {
let sha256 = payload::digest_hex(body);
if let Err(e) = self
.payloads
.upsert_prepared_sha256(&self.ctx.ai_request_id, &sha256)
.await
{
tracing::warn!(error = %e, ai_request_id = %self.ctx.ai_request_id, "prepared body digest write failed");
}
}
pub async fn set_system_prompt_override(&self, descriptor: &str) {
if let Err(e) = self
.requests
.update_system_prompt_override(&self.ctx.ai_request_id, descriptor)
.await
{
tracing::warn!(error = %e, "update_system_prompt_override failed");
}
}
pub async fn set_route_match(&self, descriptor: &str) {
if let Err(e) = self
.requests
.update_route_match(&self.ctx.ai_request_id, descriptor)
.await
{
tracing::warn!(error = %e, "update_route_match failed");
}
}
pub async fn fail(&self, error: &str) -> Result<()> {
let latency_ms = self.elapsed_ms();
let mut receipt =
journal::Receipt::pending(self.ctx.ai_request_id.clone(), self.ctx.user_id.clone());
receipt.failure = Some(error.to_owned());
if self.journal_lease.get().is_some() {
journal::record(&self.settlement, receipt).await?;
} else {
journal::settle_unadmitted_failure(&self.settlement, &receipt).await?;
}
tracing::warn!(
ai_request_id = %self.ctx.ai_request_id,
user_id = %self.ctx.user_id,
provider = %self.ctx.provider,
model = %self.effective_model(),
requested_model = %self.ctx.model,
wire_protocol = %self.ctx.wire_protocol,
status = RequestStatus::Failed.as_str(),
latency_ms,
tokens_recorded = false,
error,
"Gateway audit: request failed"
);
Ok(())
}
pub fn mark_upstream_start(&self) {
match self.upstream.lock() {
Ok(mut clock) => {
if clock.elapsed_ms.is_none() {
clock.started = Some(Instant::now());
}
},
Err(e) => tracing::warn!(error = %e, "upstream clock mutex poisoned"),
}
}
pub fn mark_upstream_end(&self) {
match self.upstream.lock() {
Ok(mut clock) => {
if clock.elapsed_ms.is_some() {
return;
}
let Some(started) = clock.started else {
return;
};
clock.elapsed_ms = Some(millis_i32(started));
},
Err(e) => tracing::warn!(error = %e, "upstream clock mutex poisoned"),
}
}
pub(crate) fn upstream_elapsed_ms(&self) -> Option<i32> {
match self.upstream.lock() {
Ok(clock) => clock.elapsed_ms,
Err(e) => {
tracing::warn!(error = %e, "upstream clock mutex poisoned");
None
},
}
}
pub(crate) fn elapsed_ms(&self) -> i32 {
millis_i32(self.started_at)
}
}
fn millis_i32(since: Instant) -> i32 {
since.elapsed().as_millis().min(i32::MAX as u128) as i32
}