diff --git a/crates/gateway/src/proxy/early_cancel.rs b/crates/gateway/src/proxy/early_cancel.rs new file mode 100644 index 00000000..42579efc --- /dev/null +++ b/crates/gateway/src/proxy/early_cancel.rs @@ -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, + session_id: Option, + model: Option, +} + +/// The handler's handle on the armed guard, carried as a request +/// extension. +#[derive(Clone)] +pub struct EarlyCancelSlot(Arc>); + +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(), + ); + } +} diff --git a/crates/gateway/src/proxy/generate.rs b/crates/gateway/src/proxy/generate.rs index bc860dd5..6ebf30b7 100644 --- a/crates/gateway/src/proxy/generate.rs +++ b/crates/gateway/src/proxy/generate.rs @@ -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}; @@ -102,12 +103,14 @@ pub async fn proxy_chat_completion( State(state): State, headers: HeaderMap, axum::Extension(identity): axum::Extension, + cancel: Option>, body: Bytes, ) -> Result { generate( state, headers, identity, + cancel.map(|c| c.0), body, CHAT, "/v1/chat/completions", @@ -121,12 +124,14 @@ pub async fn proxy_anthropic_messages( State(state): State, headers: HeaderMap, axum::Extension(identity): axum::Extension, + cancel: Option>, body: Bytes, ) -> Result { generate( state, headers, identity, + cancel.map(|c| c.0), body, MESSAGES, "/v1/messages", @@ -140,12 +145,14 @@ pub async fn proxy_responses( State(state): State, headers: HeaderMap, axum::Extension(identity): axum::Extension, + cancel: Option>, body: Bytes, ) -> Result { generate( state, headers, identity, + cancel.map(|c| c.0), body, RESPONSES, "/v1/responses", @@ -163,12 +170,14 @@ pub async fn proxy_gemini( OriginalUri(uri): OriginalUri, headers: HeaderMap, axum::Extension(identity): axum::Extension, + cancel: Option>, body: Bytes, ) -> Result { generate( state, headers, identity, + cancel.map(|c| c.0), body, GEMINI, uri.path(), @@ -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, body: Bytes, surface: ClientSurface, path: &str, query: Option<&str>, ) -> Result { - 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, body: Bytes, surface: ClientSurface, path: &str, @@ -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. @@ -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, diff --git a/crates/gateway/src/proxy/mod.rs b/crates/gateway/src/proxy/mod.rs index 4fa9612a..7d36b29a 100644 --- a/crates/gateway/src/proxy/mod.rs +++ b/crates/gateway/src/proxy/mod.rs @@ -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; @@ -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, }; diff --git a/crates/gateway/src/proxy/responses_ws.rs b/crates/gateway/src/proxy/responses_ws.rs index e94aae61..8180af0a 100644 --- a/crates/gateway/src/proxy/responses_ws.rs +++ b/crates/gateway/src/proxy/responses_ws.rs @@ -95,6 +95,7 @@ async fn serve( state.clone(), headers.clone(), identity.clone(), + None, Bytes::from(body), RESPONSES, "/v1/responses", diff --git a/crates/server/src/middleware/api_key_auth.rs b/crates/server/src/middleware/api_key_auth.rs index 6ea4b5a9..ff134eca 100644 --- a/crates/server/src/middleware/api_key_auth.rs +++ b/crates/server/src/middleware/api_key_auth.rs @@ -94,6 +94,7 @@ pub fn require_api_key( ) -> impl Fn(State, Request, Next) -> AuthFuture + Clone { move |State(state): State, mut request: Request, next: Next| { Box::pin(async move { + let started = std::time::Instant::now(); let token = presented_key(request.headers(), request.uri().query()) .ok_or(StatusCode::UNAUTHORIZED)? .to_string(); @@ -176,139 +177,166 @@ pub fn require_api_key( return Err(StatusCode::UNAUTHORIZED); } - // Update last_used_at (best-effort, don't block on failure) - let db = state.db.clone(); - let key_id = row.id; - tokio::spawn(async move { - if let Err(e) = - sqlx::query("UPDATE api_keys SET last_used_at = now() WHERE id = $1") - .bind(key_id) - .execute(&db) - .await - { - tracing::warn!("Failed to update api_key last_used_at: {e}"); - } + // From here on a client that leaves before the handler has a + // response still leaves a gateway_logs row (the MCP surface + // records its own). Every return below produces a response + // or an auth refusal, so the guard is disarmed after all of + // them; only a dropped future leaves it armed. + let cancel = (surface == "ai_gateway").then(|| { + think_watch_gateway::proxy::EarlyCancel::arm( + state.audit.clone(), + GatewayRequestIdentity { + user_id: row.user_id.map(|u| u.to_string()), + api_key_id: Some(row.id.to_string()), + api_key_lineage_id: Some(row.lineage_id.to_string()), + ..Default::default() + }, + started, + ) }); + let result: Result = async { + // Update last_used_at (best-effort, don't block on failure) + let db = state.db.clone(); + let key_id = row.id; + tokio::spawn(async move { + if let Err(e) = + sqlx::query("UPDATE api_keys SET last_used_at = now() WHERE id = $1") + .bind(key_id) + .execute(&db) + .await + { + tracing::warn!("Failed to update api_key last_used_at: {e}"); + } + }); - // Compute the user's role-derived constraints and intersect - // with the API-key allow-list. The role union is loaded once - // per request — fast enough at our scale. - // - // We also pull the role NAMES so the MCP access controller - // can gate per-tool access without re-querying the DB, and - // the aggregated `surface_constraints` JSON so the gateway - // hot path has rate limits + budgets without further lookups. - let (role_limits, user_roles, surface_constraints) = if let Some(uid) = row.user_id { - let limits = rbac::compute_user_resource_limits(&state.db, uid) - .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - let names = rbac::load_user_role_names(&state.db, uid) - .await - .unwrap_or_default(); - // Use the api_key-aware variant so per-key - // `rate_limit_rules` / `budget_caps` rows fire on the - // gateway hot path. Falling back to the user-only - // function would silently drop api_key-scope - // overrides — the schema supports them but the - // gateway would never see them. - let constraints = - rbac::compute_effective_surface_constraints(&state.db, uid, row.id) + // Compute the user's role-derived constraints and intersect + // with the API-key allow-list. The role union is loaded once + // per request — fast enough at our scale. + // + // We also pull the role NAMES so the MCP access controller + // can gate per-tool access without re-querying the DB, and + // the aggregated `surface_constraints` JSON so the gateway + // hot path has rate limits + budgets without further lookups. + let (role_limits, user_roles, surface_constraints) = if let Some(uid) = row.user_id { + let limits = rbac::compute_user_resource_limits(&state.db, uid) + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + let names = rbac::load_user_role_names(&state.db, uid) .await .unwrap_or_default(); - (limits, names, constraints) - } else { - // Service-account API keys (no user_id) inherit only - // the per-key constraints, since there's no user to - // resolve roles against. They get an empty role list, - // which means the MCP access controller will deny - // anything that requires a role match, and an empty - // constraint set so no role-inline limits fire. - ( - rbac::UserResourceLimits { - allowed_models: None, - allowed_mcp_tools: None, - }, - Vec::new(), - think_watch_common::limits::SurfaceConstraints::default(), - ) - }; - let merged_models = - intersect_allowlists(row.allowed_models.clone(), role_limits.allowed_models); - let merged_mcp_tools = - intersect_allowlists(row.allowed_mcp_tools.clone(), role_limits.allowed_mcp_tools); - - // Load email for template header resolution ({{user_email}}) - let user_email: Option = if let Some(uid) = row.user_id { - sqlx::query_scalar("SELECT email FROM users WHERE id = $1") - .bind(uid) - .fetch_optional(&state.db) - .await - .ok() - .flatten() - } else { - None - }; + // Use the api_key-aware variant so per-key + // `rate_limit_rules` / `budget_caps` rows fire on the + // gateway hot path. Falling back to the user-only + // function would silently drop api_key-scope + // overrides — the schema supports them but the + // gateway would never see them. + let constraints = + rbac::compute_effective_surface_constraints(&state.db, uid, row.id) + .await + .unwrap_or_default(); + (limits, names, constraints) + } else { + // Service-account API keys (no user_id) inherit only + // the per-key constraints, since there's no user to + // resolve roles against. They get an empty role list, + // which means the MCP access controller will deny + // anything that requires a role match, and an empty + // constraint set so no role-inline limits fire. + ( + rbac::UserResourceLimits { + allowed_models: None, + allowed_mcp_tools: None, + }, + Vec::new(), + think_watch_common::limits::SurfaceConstraints::default(), + ) + }; + let merged_models = + intersect_allowlists(row.allowed_models.clone(), role_limits.allowed_models); + let merged_mcp_tools = + intersect_allowlists(row.allowed_mcp_tools.clone(), role_limits.allowed_mcp_tools); - // Resolve client IP once, share across both identities so - // gateway_logs and mcp_logs see the same value the rest - // of the auth stack uses (honours client_ip_source + - // trusted_proxies via auth_guard::extract_client_ip). - let client_ip = crate::middleware::auth_guard::extract_client_ip( - &state, - request.headers(), - request.extensions(), - ) - .await; + // Load email for template header resolution ({{user_email}}) + let user_email: Option = if let Some(uid) = row.user_id { + sqlx::query_scalar("SELECT email FROM users WHERE id = $1") + .bind(uid) + .fetch_optional(&state.db) + .await + .ok() + .flatten() + } else { + None + }; - let gateway_identity = GatewayRequestIdentity { - user_id: row.user_id.map(|u| u.to_string()), - user_email, - api_key_id: Some(row.id.to_string()), - api_key_lineage_id: Some(row.lineage_id.to_string()), - allowed_models: merged_models.clone(), - surface_constraints: surface_constraints.clone(), - ip_address: client_ip.clone(), - }; + // Resolve client IP once, share across both identities so + // gateway_logs and mcp_logs see the same value the rest + // of the auth stack uses (honours client_ip_source + + // trusted_proxies via auth_guard::extract_client_ip). + let client_ip = crate::middleware::auth_guard::extract_client_ip( + &state, + request.headers(), + request.extensions(), + ) + .await; - // The MCP transport handlers expect their own typed - // extension and require a user_id (sessions are keyed - // by user). Service-account keys without a user_id - // can't talk to MCP — return 401 here rather than - // letting the handler 500 on a missing extension. - if surface == "mcp_gateway" { - let Some(uid) = row.user_id else { - tracing::warn!( - api_key_id = %row.id, - "MCP gateway requires a user-bound API key (service-account keys are not supported)" - ); - return Err(StatusCode::UNAUTHORIZED); - }; - // Reuse the email already loaded for `gateway_identity` - // above — same user_id, same row. The MCP branch used - // to issue a SECOND `SELECT email` query against PG on - // every request which is pure waste; the user-state - // gate at the JOIN above guarantees the user still - // exists, so an absent email here means the user was - // hard-deleted between the JOIN and this point (rare) - // and we should 401 rather than serve the request. - let Some(user_email) = gateway_identity.user_email.clone() else { - return Err(StatusCode::UNAUTHORIZED); - }; - let mcp_identity = McpRequestIdentity { - user_id: uid, + let gateway_identity = GatewayRequestIdentity { + user_id: row.user_id.map(|u| u.to_string()), user_email, - user_roles, + api_key_id: Some(row.id.to_string()), + api_key_lineage_id: Some(row.lineage_id.to_string()), + allowed_models: merged_models.clone(), surface_constraints: surface_constraints.clone(), - allowed_mcp_tools: merged_mcp_tools.clone(), - mcp_account_overrides: row.mcp_account_overrides.clone(), ip_address: client_ip.clone(), }; - request.extensions_mut().insert(mcp_identity); - } - request.extensions_mut().insert(gateway_identity); + // The MCP transport handlers expect their own typed + // extension and require a user_id (sessions are keyed + // by user). Service-account keys without a user_id + // can't talk to MCP — return 401 here rather than + // letting the handler 500 on a missing extension. + if surface == "mcp_gateway" { + let Some(uid) = row.user_id else { + tracing::warn!( + api_key_id = %row.id, + "MCP gateway requires a user-bound API key (service-account keys are not supported)" + ); + return Err(StatusCode::UNAUTHORIZED); + }; + // Reuse the email already loaded for `gateway_identity` + // above — same user_id, same row. The MCP branch used + // to issue a SECOND `SELECT email` query against PG on + // every request which is pure waste; the user-state + // gate at the JOIN above guarantees the user still + // exists, so an absent email here means the user was + // hard-deleted between the JOIN and this point (rare) + // and we should 401 rather than serve the request. + let Some(user_email) = gateway_identity.user_email.clone() else { + return Err(StatusCode::UNAUTHORIZED); + }; + let mcp_identity = McpRequestIdentity { + user_id: uid, + user_email, + user_roles, + surface_constraints: surface_constraints.clone(), + allowed_mcp_tools: merged_mcp_tools.clone(), + mcp_account_overrides: row.mcp_account_overrides.clone(), + ip_address: client_ip.clone(), + }; + request.extensions_mut().insert(mcp_identity); + } - Ok(next.run(request).await) + if let Some(c) = &cancel { + c.identity(&gateway_identity); + request.extensions_mut().insert(c.slot()); + } + request.extensions_mut().insert(gateway_identity); + Ok(next.run(request).await) + } + .await; + if let Some(c) = cancel { + c.disarm(); + } + result }) } } diff --git a/crates/test-support/tests/early_cancel.rs b/crates/test-support/tests/early_cancel.rs new file mode 100644 index 00000000..81a715be --- /dev/null +++ b/crates/test-support/tests/early_cancel.rs @@ -0,0 +1,193 @@ +//! A client that leaves before its response exists still leaves one +//! `gateway_logs` row: status 499, `client_cancelled`, no tokens, no +//! cost. +//! +//! Before, only a stream that had started recorded a disconnect. A +//! client that left while the key's roles loaded, the limits ran, a +//! route was picked or a whole answer was awaited left no trace: hyper +//! dropped the handler and nothing after the await point ran. + +use serde_json::Value; +use think_watch_test_support::prelude::*; +use wiremock::matchers::{method, path}; +use wiremock::{Mock, MockServer, ResponseTemplate}; + +/// A user, a key, and `model` routed to an upstream that takes a minute +/// to answer. +async fn slow_route(app: &TestApp, model: &str) -> (MockServer, Uuid, String) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(json!({"choices": []})) + .set_delay(std::time::Duration::from_secs(60)), + ) + .mount(&server) + .await; + let user = fixtures::create_random_user(&app.db).await.unwrap(); + let provider = fixtures::create_provider( + &app.db, + &unique_name("early-cancel"), + "openai", + &server.uri(), + None, + ) + .await + .unwrap(); + fixtures::create_model_and_route(&app.db, provider.id, model) + .await + .unwrap(); + app.rebuild_gateway_router().await; + let key = fixtures::create_api_key(&app.db, user.user.id, "ec", &["ai_gateway"], None, None) + .await + .unwrap(); + (server, user.user.id, key.plaintext) +} + +/// The user's rows, once `want` of them have landed (or after ~10 s). +async fn rows(app: &TestApp, user_id: Uuid, want: usize) -> Vec<(i64, i64, i64, String)> { + let ch = app.state.clickhouse.as_ref().expect("CH wired up"); + let mut found = Vec::new(); + for _ in 0..100 { + found = ch + .query( + "SELECT ifNull(status_code, -1), ifNull(input_tokens, -1), \ + ifNull(output_tokens, -1), ifNull(detail, '') \ + FROM gateway_logs WHERE user_id = ?", + ) + .bind(user_id.to_string()) + .fetch_all::<(i64, i64, i64, String)>() + .await + .expect("CH query"); + if found.len() >= want { + break; + } + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + } + found +} + +fn assert_cancelled(row: &(i64, i64, i64, String)) { + let (status, input, output, detail) = row; + assert_eq!(*status, 499, "{detail}"); + assert_eq!((*input, *output), (0, 0), "{detail}"); + let detail: Value = serde_json::from_str(detail).unwrap(); + assert_eq!(detail["stream_outcome"], "client_cancelled", "{detail}"); + assert_eq!(detail["cancelled_before"], "response", "{detail}"); +} + +/// Held deterministically: the client goes once the upstream has the +/// request, while the gateway waits for a whole answer. +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_client_that_leaves_while_a_whole_answer_is_awaited_is_recorded() { + let app = TestApp::spawn_with_clickhouse().await; + let (server, user_id, key) = slow_route(&app, "early-cancel-whole").await; + + let url = format!("{}/v1/chat/completions", app.gateway_url); + let call = tokio::spawn(async move { + reqwest::Client::new() + .post(url) + .bearer_auth(key) + .json(&json!({"model": "early-cancel-whole", + "messages": [{"role": "user", "content": "hi"}]})) + .send() + .await + }); + while server + .received_requests() + .await + .unwrap_or_default() + .is_empty() + { + tokio::task::yield_now().await; + } + call.abort(); + + let found = rows(&app, user_id, 1).await; + assert_eq!(found.len(), 1, "{found:?}"); + assert_cancelled(&found[0]); + let detail: Value = serde_json::from_str(&found[0].3).unwrap(); + assert_eq!(detail["model_id"], "early-cancel-whole", "{detail}"); +} + +/// Leaving at once: the request is dropped somewhere in the key's +/// roles, the limits or routing. A client that left before its key was +/// even looked up is nobody yet and writes nothing; every other one +/// leaves exactly one cancelled row. Never a success, never two. +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_client_that_leaves_before_routing_ends_leaves_at_most_one_cancelled_row() { + let app = TestApp::spawn_with_clickhouse().await; + let (server, user_id, key) = slow_route(&app, "early-cancel-quick").await; + + let client = reqwest::Client::builder() + .timeout(std::time::Duration::from_millis(5)) + .build() + .unwrap(); + let mut left = 0; + for _ in 0..5 { + let r = client + .post(format!("{}/v1/chat/completions", app.gateway_url)) + .bearer_auth(&key) + .json(&json!({"model": "early-cancel-quick", "stream": true, + "messages": [{"role": "user", "content": "hi"}]})) + .send() + .await; + if r.is_err() { + left += 1; + } + } + assert!(left > 0, "a 5 ms client never left early"); + + // Let the audit pipeline flush whatever was written. + let found = rows(&app, user_id, left).await; + // At most one row per client that left; and some of them left after + // the key was known — before, none of these were ever recorded. + assert!( + !found.is_empty() && found.len() <= left, + "{left} left: {found:?}" + ); + for row in &found { + let (status, ..) = row; + if *status == 499 { + // A stream that had started carries no `cancelled_before`. + let detail: Value = serde_json::from_str(&row.3).unwrap(); + assert_eq!(detail["stream_outcome"], "client_cancelled", "{detail}"); + } else { + panic!("a request whose client left was logged as {status}: {row:?}"); + } + } + drop(server); +} + +/// A stream that started records its own cancel; the guard is disarmed +/// by then, so there is one row, not two. +#[ignore = "integration test — run via `make test-it`"] +#[tokio::test] +async fn a_started_stream_that_is_left_is_recorded_once() { + let app = TestApp::spawn_with_clickhouse().await; + let (_server, user_id, key) = slow_route(&app, "early-cancel-stream").await; + + let resp = reqwest::Client::new() + .post(format!("{}/v1/chat/completions", app.gateway_url)) + .bearer_auth(&key) + .json(&json!({"model": "early-cancel-stream", "stream": true, + "messages": [{"role": "user", "content": "hi"}]})) + .send() + .await + .expect("the stream starts"); + assert_eq!(resp.status(), 200); + drop(resp); + + let found = rows(&app, user_id, 1).await; + assert_eq!(found.len(), 1, "{found:?}"); + let detail: Value = serde_json::from_str(&found[0].3).unwrap(); + assert_eq!(detail["stream_outcome"], "client_cancelled", "{detail}"); + assert!(detail.get("cancelled_before").is_none(), "{detail}"); + + // No second row turns up later. + tokio::time::sleep(std::time::Duration::from_secs(3)).await; + assert_eq!(rows(&app, user_id, 1).await.len(), 1); +}