Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
138 changes: 138 additions & 0 deletions crates/gateway/src/proxy/early_cancel.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
//! A client that leaves before it gets a response still leaves a row.
//!
//! Once a stream has started, a disconnect is recorded by the stream's
//! tail (`StreamOutcome::ClientCancelled`). Before that — while the key's
//! roles and limits load, the pre-flight stages run, a route is picked,
//! or a whole (non-streamed) answer is awaited — the request is just a
//! future, and when the client goes hyper drops it: nothing after the
//! await point runs, and there was no trace of the request at all.
//!
//! [`EarlyCancel`] is armed by the API-key middleware as soon as the key
//! is known, and disarmed when the handler hands back a response —
//! whatever it is, since every response path writes its own row. If it
//! is dropped still armed, the future was dropped: it writes one
//! `gateway_logs` row with status 499 and no tokens, no cost, no upstream.
//! The handler fills in what it learns on the way ([`EarlyCancelSlot`]):
//! the trace id and the model.

use std::sync::{Arc, Mutex};
use std::time::Instant;

use rust_decimal::Decimal;
use think_watch_common::audit::AuditLogger;

use super::GatewayRequestIdentity;
use super::body_capture::BodyCapture;
use super::log_ctx::emit_gateway_log_with_extra;

/// What is known about the request so far.
#[derive(Default)]
struct Known {
identity: GatewayRequestIdentity,
trace_id: Option<String>,
session_id: Option<String>,
model: Option<String>,
}

/// The handler's handle on the armed guard, carried as a request
/// extension.
#[derive(Clone)]
pub struct EarlyCancelSlot(Arc<Mutex<Known>>);

impl EarlyCancelSlot {
/// The ids the request's other rows carry, once the handler has them.
pub(crate) fn request(&self, trace_id: &str, session_id: Option<&str>) {
if let Ok(mut k) = self.0.lock() {
k.trace_id = Some(trace_id.to_string());
k.session_id = session_id.map(str::to_string);
}
}

/// The model the caller named, after aliasing.
pub(crate) fn model(&self, model: &str) {
if let Ok(mut k) = self.0.lock() {
k.model = Some(model.to_string());
}
}
}

/// Writes the cancelled row if dropped before [`EarlyCancel::disarm`].
pub struct EarlyCancel {
audit: AuditLogger,
known: EarlyCancelSlot,
started: Instant,
armed: bool,
}

impl EarlyCancel {
/// Arm for a request that arrived at `started`, from `identity` (what
/// the middleware has resolved so far).
pub fn arm(audit: AuditLogger, identity: GatewayRequestIdentity, started: Instant) -> Self {
Self {
audit,
known: EarlyCancelSlot(Arc::new(Mutex::new(Known {
identity,
..Default::default()
}))),
started,
armed: true,
}
}

/// The fuller identity, once the middleware has it.
pub fn identity(&self, identity: &GatewayRequestIdentity) {
if let Ok(mut k) = self.known.0.lock() {
k.identity = identity.clone();
}
}

pub fn slot(&self) -> EarlyCancelSlot {
self.known.clone()
}

/// A response exists; it records itself.
pub fn disarm(mut self) {
self.armed = false;
}
}

impl Drop for EarlyCancel {
fn drop(&mut self) {
if !self.armed {
return;
}
let Ok(k) = self.known.0.lock() else {
return;
};
metrics::counter!("gateway_cancelled_before_response_total").increment(1);
let trace_id = k
.trace_id
.clone()
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
let id = &k.identity;
emit_gateway_log_with_extra(
&self.audit,
&trace_id,
k.session_id.as_deref(),
id.user_id.as_deref(),
id.user_email.as_deref(),
id.api_key_id.as_deref(),
id.api_key_lineage_id.as_deref(),
id.ip_address.as_deref(),
k.model.as_deref().unwrap_or("(unknown)"),
None,
None,
0,
0,
Decimal::ZERO,
self.started.elapsed().as_millis() as i64,
499,
Some(serde_json::json!({
// The marker a cancelled stream carries too.
"stream_outcome": "client_cancelled",
"cancelled_before": "response",
})),
BodyCapture::default(),
);
}
}
25 changes: 24 additions & 1 deletion crates/gateway/src/proxy/generate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ use tw_dialect::convert::Session;
use tw_dialect::ir::{Dialect, Target};

use super::body_capture::prepare_body_capture;
use super::early_cancel::EarlyCancelSlot;
use super::headers::{request_id_header, resolve_session_id, resolve_trace_id};
use super::log_ctx::{LogCtx, emit_gateway_error_log, emit_gateway_log};
use super::pipeline::{launch_stream_pump, run_buffered_post_invoke, run_preflight_stages};
Expand Down Expand Up @@ -102,12 +103,14 @@ pub async fn proxy_chat_completion(
State(state): State<GatewayState>,
headers: HeaderMap,
axum::Extension(identity): axum::Extension<GatewayRequestIdentity>,
cancel: Option<axum::Extension<EarlyCancelSlot>>,
body: Bytes,
) -> Result<axum::response::Response, GatewayErrorResponse> {
generate(
state,
headers,
identity,
cancel.map(|c| c.0),
body,
CHAT,
"/v1/chat/completions",
Expand All @@ -121,12 +124,14 @@ pub async fn proxy_anthropic_messages(
State(state): State<GatewayState>,
headers: HeaderMap,
axum::Extension(identity): axum::Extension<GatewayRequestIdentity>,
cancel: Option<axum::Extension<EarlyCancelSlot>>,
body: Bytes,
) -> Result<axum::response::Response, GatewayErrorResponse> {
generate(
state,
headers,
identity,
cancel.map(|c| c.0),
body,
MESSAGES,
"/v1/messages",
Expand All @@ -140,12 +145,14 @@ pub async fn proxy_responses(
State(state): State<GatewayState>,
headers: HeaderMap,
axum::Extension(identity): axum::Extension<GatewayRequestIdentity>,
cancel: Option<axum::Extension<EarlyCancelSlot>>,
body: Bytes,
) -> Result<axum::response::Response, GatewayErrorResponse> {
generate(
state,
headers,
identity,
cancel.map(|c| c.0),
body,
RESPONSES,
"/v1/responses",
Expand All @@ -163,12 +170,14 @@ pub async fn proxy_gemini(
OriginalUri(uri): OriginalUri,
headers: HeaderMap,
axum::Extension(identity): axum::Extension<GatewayRequestIdentity>,
cancel: Option<axum::Extension<EarlyCancelSlot>>,
body: Bytes,
) -> Result<axum::response::Response, GatewayErrorResponse> {
generate(
state,
headers,
identity,
cancel.map(|c| c.0),
body,
GEMINI,
uri.path(),
Expand Down Expand Up @@ -458,24 +467,32 @@ pub(crate) async fn read_whole(
///
/// `path` and `query` are the caller's: Gemini puts the model, whether to
/// stream and which stream form in them.
///
/// `cancel` is the middleware's record of a request whose client leaves
/// before the response exists (see `early_cancel`); `None` on a
/// WebSocket turn, whose connection records its own end.
#[allow(clippy::too_many_arguments)]
pub(crate) async fn generate(
state: GatewayState,
headers: HeaderMap,
identity: GatewayRequestIdentity,
cancel: Option<EarlyCancelSlot>,
body: Bytes,
surface: ClientSurface,
path: &str,
query: Option<&str>,
) -> Result<axum::response::Response, GatewayErrorResponse> {
run(state, headers, identity, body, surface, path, query)
run(state, headers, identity, cancel, body, surface, path, query)
.await
.map_err(|e| e.in_dialect(surface.dialect))
}

#[allow(clippy::too_many_arguments)]
async fn run(
state: GatewayState,
headers: HeaderMap,
identity: GatewayRequestIdentity,
cancel: Option<EarlyCancelSlot>,
body: Bytes,
surface: ClientSurface,
path: &str,
Expand All @@ -484,6 +501,9 @@ async fn run(
let trace_id = resolve_trace_id(&headers);
let session_id = resolve_session_id(&headers);
let request_started_at = std::time::Instant::now();
if let Some(c) = &cancel {
c.request(&trace_id, session_id.as_deref());
}

// A row even for a body we cannot read: an operator chasing a 400
// should find it.
Expand Down Expand Up @@ -527,6 +547,9 @@ async fn run(

// 1. Model aliases
let mapped_model = state.model_mapper.map(&model);
if let Some(c) = &cancel {
c.model(&mapped_model);
}
let ctx = LogCtx::new(
&state.audit,
&identity,
Expand Down
2 changes: 2 additions & 0 deletions crates/gateway/src/proxy/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ use think_watch_common::limits::weight;

mod accounting;
mod body_capture;
mod early_cancel;
pub(crate) mod generate;
mod headers;
mod identity;
Expand All @@ -44,6 +45,7 @@ pub(crate) use log_ctx::emit_gateway_log_with_extra;
pub(crate) use routing::{SelectionRecord, fails as upstream_failed, finalize_health};

// pub re-exports — `server::app` mounts these as route handlers.
pub use early_cancel::{EarlyCancel, EarlyCancelSlot};
pub use generate::{
proxy_anthropic_messages, proxy_chat_completion, proxy_gemini, proxy_responses,
};
Expand Down
1 change: 1 addition & 0 deletions crates/gateway/src/proxy/responses_ws.rs
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,7 @@ async fn serve(
state.clone(),
headers.clone(),
identity.clone(),
None,
Bytes::from(body),
RESPONSES,
"/v1/responses",
Expand Down
Loading
Loading