diff --git a/Cargo.lock b/Cargo.lock index 53e874f0a..3d5723982 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -131,6 +131,7 @@ dependencies = [ "temp-env", "tokio", "tracing", + "uuid", "workspace_root", ] diff --git a/crates/alien-ai-gateway/Cargo.toml b/crates/alien-ai-gateway/Cargo.toml index 70f12042b..bdc9304a8 100644 --- a/crates/alien-ai-gateway/Cargo.toml +++ b/crates/alien-ai-gateway/Cargo.toml @@ -31,6 +31,7 @@ aws-credential-types = { workspace = true } aws-smithy-eventstream = { workspace = true } aws-smithy-types = { workspace = true } http = { workspace = true } +uuid = { workspace = true, features = ["v4"] } [dev-dependencies] httpmock = { workspace = true } diff --git a/crates/alien-ai-gateway/src/lib.rs b/crates/alien-ai-gateway/src/lib.rs index 8def0e15b..339b32ffc 100644 --- a/crates/alien-ai-gateway/src/lib.rs +++ b/crates/alien-ai-gateway/src/lib.rs @@ -10,14 +10,20 @@ mod config; mod creds; mod error; mod router; +mod usage; pub use config::{bindings_from_env, bindings_from_env_map, route_from_remote_ai_lease}; pub use creds::{ AmbientCred, AnthropicApiKeyCred, AwsSigV4Cred, BearerTokenCred, OpenAiApiKeyCred, }; pub use error::{ErrorData, Result}; pub use router::{ - build_router, build_router_with_availability, route_from_direct_anthropic, - route_from_direct_openai, AvailableModels, GatewayRoute, GatewayTarget, + build_router, build_router_with_availability, build_router_with_availability_and_observer, + build_router_with_observer, route_from_direct_anthropic, route_from_direct_openai, + AvailableModels, GatewayRoute, GatewayTarget, +}; +pub use usage::{ + parse_ai_token_usage, AiTokenUsage, AiUsageClientApi, AiUsageEvent, AiUsageObserver, + AiUsageOutcome, AiUsageProvider, }; use std::net::{Ipv4Addr, SocketAddr}; diff --git a/crates/alien-ai-gateway/src/router/mod.rs b/crates/alien-ai-gateway/src/router/mod.rs index cc1bea603..b212259db 100644 --- a/crates/alien-ai-gateway/src/router/mod.rs +++ b/crates/alien-ai-gateway/src/router/mod.rs @@ -22,6 +22,9 @@ use serde_json::{json, Value}; use crate::creds::{AmbientCred, AnthropicApiKeyCred, OpenAiApiKeyCred}; use crate::error::{ErrorData, Result}; +use crate::usage::{ + observe_response, AiUsageClientApi, AiUsageContext, AiUsageObserver, AiUsageProvider, +}; mod bedrock; mod eventstream; @@ -141,6 +144,7 @@ struct AppState { /// Account-specific, read-only control-plane observations supplied by the /// hosted route resolver. `None` keeps embedded gateways catalog-only. available_models: Option, + usage_observer: Option>, } /// Available public model IDs keyed by binding name. @@ -150,7 +154,15 @@ pub type AvailableModels = HashMap>; /// `POST //v1/chat/completions` (OpenAI), `POST //v1/messages` /// (Anthropic), and `GET //v1/models`. pub fn build_router(routes: Vec) -> Router { - build_router_inner(routes, None) + build_router_inner(routes, None, None) +} + +/// Build a router that reports completed requests to a non-blocking observer. +pub fn build_router_with_observer( + routes: Vec, + usage_observer: Arc, +) -> Router { + build_router_inner(routes, None, Some(usage_observer)) } /// Build a router whose model listing and inference paths are restricted by a @@ -159,12 +171,22 @@ pub fn build_router_with_availability( routes: Vec, available_models: AvailableModels, ) -> Router { - build_router_inner(routes, Some(available_models)) + build_router_inner(routes, Some(available_models), None) +} + +/// Build a hosted router with both bounded model availability and usage observation. +pub fn build_router_with_availability_and_observer( + routes: Vec, + available_models: AvailableModels, + usage_observer: Arc, +) -> Router { + build_router_inner(routes, Some(available_models), Some(usage_observer)) } fn build_router_inner( routes: Vec, available_models: Option, + usage_observer: Option>, ) -> Router { let routes: HashMap = routes.into_iter().map(|r| (r.name.clone(), r)).collect(); @@ -172,6 +194,7 @@ fn build_router_inner( routes, client: reqwest::Client::new(), available_models, + usage_observer, }); Router::new() .route( @@ -275,6 +298,23 @@ async fn forward_response(upstream: reqwest::Response) -> Result { }) } +fn usage_client_api(client_api: ClientApi) -> AiUsageClientApi { + match client_api { + ClientApi::OpenAiChatCompletions => AiUsageClientApi::OpenAiChatCompletions, + ClientApi::OpenAiResponses => AiUsageClientApi::OpenAiResponses, + ClientApi::AnthropicMessages => AiUsageClientApi::AnthropicMessages, + } +} + +fn cloud_usage_provider(cloud: Platform) -> AiUsageProvider { + match cloud { + Platform::Aws => AiUsageProvider::AwsBedrock, + Platform::Gcp => AiUsageProvider::GcpVertex, + Platform::Azure => AiUsageProvider::AzureFoundry, + _ => unreachable!("AI cloud routes are available only on AWS, GCP, and Azure"), + } +} + /// Build a JSON POST to `url`, sign it with the ambient credential for `service`, /// and execute it. The handlers differ only in URL, signing service, body, and any /// protocol-required header, so the build + sign + execute + upstream-error @@ -365,7 +405,24 @@ async fn proxy( message: format!("direct Anthropic supports only /{binding}/v1/messages"), })); } - return proxy_direct_anthropic(&state.client, route, payload, &model, &headers).await; + let provider_model = ai_catalog::resolve_direct_anthropic(&model) + .map(|resolved| resolved.upstream_id) + .unwrap_or(model.as_str()); + let descriptor = AiUsageContext::new( + &binding, + AiUsageProvider::Anthropic, + &model, + provider_model, + usage_client_api(client_api), + None, + ); + let response = + proxy_direct_anthropic(&state.client, route, payload, &model, &headers).await?; + return Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )); } GatewayTarget::DirectOpenAi => { ensure_model_available(&state, &binding, &model)?; @@ -376,14 +433,27 @@ async fn proxy( ), })); } - return proxy_direct_openai( + let descriptor = AiUsageContext::new( + &binding, + AiUsageProvider::OpenAi, + &model, + &model, + usage_client_api(client_api), + None, + ); + let response = proxy_direct_openai( &state.client, route, payload, &model, "/v1/chat/completions", ) - .await; + .await?; + return Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )); } GatewayTarget::Cloud(_) => {} } @@ -401,6 +471,15 @@ async fn proxy( })?; ensure_model_available(&state, &binding, &model)?; + let descriptor = AiUsageContext::new( + &binding, + cloud_usage_provider(cloud), + &model, + cm.upstream_id, + usage_client_api(client_api), + route.region.clone(), + ); + if !cm.client_apis.contains(&client_api) { let expected_path = match cm.client_apis.first() { Some(ClientApi::OpenAiChatCompletions) => "v1/chat/completions", @@ -419,21 +498,38 @@ async fn proxy( // endpoint: the model id travels in the URL and the streamed reply is AWS // event-stream framing, so it needs its own request/response shape. if cloud == Platform::Aws && cm.provider_api == ProviderApi::Anthropic { - return proxy_bedrock_anthropic(&state.client, route, cm.upstream_id, payload, &headers) - .await; + let response = + proxy_bedrock_anthropic(&state.client, route, cm.upstream_id, payload, &headers) + .await?; + return Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )); } // GCP serves Claude through Vertex rawPredict: the model id travels in the URL // and streaming is chosen by the URL verb, but the reply is native Anthropic // JSON/SSE — no decoder needed, unlike Bedrock. if cloud == Platform::Gcp && cm.provider_api == ProviderApi::Anthropic { - return proxy_vertex_anthropic(&state.client, route, cm.upstream_id, payload, &headers) - .await; + let response = + proxy_vertex_anthropic(&state.client, route, cm.upstream_id, payload, &headers).await?; + return Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )); } // Azure serves Claude through Foundry's Anthropic endpoint: standard Messages // in both directions, on the `/anthropic/v1` path with the version header. if cloud == Platform::Azure && cm.provider_api == ProviderApi::Anthropic { - return proxy_foundry_anthropic(&state.client, route, cm.upstream_id, payload, &headers) - .await; + let response = + proxy_foundry_anthropic(&state.client, route, cm.upstream_id, payload, &headers) + .await?; + return Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )); } payload["model"] = Value::String(cm.upstream_id.to_string()); @@ -456,7 +552,12 @@ async fn proxy( ) .await?; - forward_response(upstream).await + let response = forward_response(upstream).await?; + Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )) } /// Proxy an OpenAI Responses request (`POST //v1/responses`, used by Codex). @@ -490,8 +591,21 @@ async fn proxy_responses( } GatewayTarget::DirectOpenAi => { ensure_model_available(&state, &binding, &model)?; - return proxy_direct_openai(&state.client, route, payload, &model, "/v1/responses") - .await; + let descriptor = AiUsageContext::new( + &binding, + AiUsageProvider::OpenAi, + &model, + &model, + AiUsageClientApi::OpenAiResponses, + None, + ); + let response = + proxy_direct_openai(&state.client, route, payload, &model, "/v1/responses").await?; + return Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )); } }; let catalog_model = ai_catalog::resolve_for(&model, cloud) @@ -509,6 +623,14 @@ async fn proxy_responses( binding: binding.clone(), }) })?; + let descriptor = AiUsageContext::new( + &binding, + cloud_usage_provider(cloud), + &model, + target.upstream_id, + AiUsageClientApi::OpenAiResponses, + route.region.clone(), + ); payload["model"] = Value::String(target.upstream_id.to_string()); let upstream_body = @@ -538,7 +660,12 @@ async fn proxy_responses( ) .await?; - forward_response(upstream).await + let response = forward_response(upstream).await?; + Ok(observe_response( + response, + state.usage_observer.as_ref(), + descriptor, + )) } /// `GET //v1/models`: the qualified catalog, intersected with the bounded @@ -777,6 +904,8 @@ fn parse_stream_flag(value: Option) -> Result { #[cfg(test)] mod tests { use std::net::Ipv4Addr; + use std::sync::mpsc; + use std::time::Duration; use aws_credential_types::provider::SharedCredentialsProvider; use aws_credential_types::Credentials; @@ -787,6 +916,15 @@ mod tests { use super::*; use crate::creds::{AwsSigV4Cred, BearerTokenCred}; + use crate::usage::{AiTokenUsage, AiUsageEvent, AiUsageOutcome}; + + struct TestUsageObserver(mpsc::Sender); + + impl AiUsageObserver for TestUsageObserver { + fn observe(&self, event: AiUsageEvent) { + let _ = self.0.send(event); + } + } fn test_aws_cred() -> AmbientCred { let creds = Credentials::new( @@ -837,6 +975,100 @@ mod tests { } } + fn direct_openai_route(upstream: &str) -> GatewayRoute { + let mut route = route_from_direct_openai("llm", "sk-test").expect("direct OpenAI route"); + route.upstream_base_override = Some(upstream.to_string()); + route + } + + #[tokio::test] + async fn observes_usage_from_a_completed_response_body() { + let server = MockServer::start_async().await; + let upstream = server + .mock_async(|when, then| { + when.method(POST).path("/v1/chat/completions"); + then.status(200) + .header("content-type", "application/json") + .json_body(json!({ + "id": "response", + "usage": { + "prompt_tokens": 13, + "completion_tokens": 5, + "prompt_tokens_details": { "cached_tokens": 3 } + } + })); + }) + .await; + let (sender, receiver) = mpsc::channel(); + let observer: Arc = Arc::new(TestUsageObserver(sender)); + let url = serve(build_router_with_observer( + vec![direct_openai_route(&server.base_url())], + observer, + )) + .await; + + let response = reqwest::Client::new() + .post(format!("{url}/llm/v1/chat/completions")) + .json(&json!({ "model": "gpt-5-mini", "messages": [] })) + .send() + .await + .expect("proxy response"); + let body = response.bytes().await.expect("response body"); + assert!(body + .windows(b"prompt_tokens".len()) + .any(|part| part == b"prompt_tokens")); + + let event = receiver + .recv_timeout(Duration::from_secs(1)) + .expect("completed usage observation"); + assert_eq!(event.provider, AiUsageProvider::OpenAi); + assert_eq!(event.public_model, "gpt-5-mini"); + assert_eq!(event.client_api, AiUsageClientApi::OpenAiChatCompletions); + assert_eq!(event.outcome, AiUsageOutcome::Success); + assert_eq!(event.status, 200); + assert_eq!(event.tokens.input_tokens, Some(13)); + assert_eq!(event.tokens.output_tokens, Some(5)); + assert_eq!(event.tokens.cache_read_tokens, Some(3)); + upstream.assert_async().await; + } + + #[tokio::test] + async fn observes_sanitized_provider_failures_without_parsing_the_error_body() { + let server = MockServer::start_async().await; + server + .mock_async(|when, then| { + when.method(POST).path("/v1/chat/completions"); + then.status(429) + .header("retry-after", "2") + .json_body(json!({ "secret_provider_detail": "must not escape" })); + }) + .await; + let (sender, receiver) = mpsc::channel(); + let observer: Arc = Arc::new(TestUsageObserver(sender)); + let url = serve(build_router_with_observer( + vec![direct_openai_route(&server.base_url())], + observer, + )) + .await; + + let response = reqwest::Client::new() + .post(format!("{url}/llm/v1/chat/completions")) + .json(&json!({ "model": "gpt-5-mini", "messages": [] })) + .send() + .await + .expect("proxy response"); + assert_eq!(response.status(), 429); + let body = response.text().await.expect("safe error body"); + assert!(!body.contains("secret_provider_detail")); + + let event = receiver + .recv_timeout(Duration::from_secs(1)) + .expect("provider error observation"); + assert_eq!(event.outcome, AiUsageOutcome::ProviderError); + assert_eq!(event.status, 429); + assert_eq!(event.tokens, AiTokenUsage::default()); + } + #[test] fn gcp_vertex_url_regional_vs_global() { // A region prefixes the host; `global` uses the un-prefixed host. The path always diff --git a/crates/alien-ai-gateway/src/usage.rs b/crates/alien-ai-gateway/src/usage.rs new file mode 100644 index 000000000..d3233ebb4 --- /dev/null +++ b/crates/alien-ai-gateway/src/usage.rs @@ -0,0 +1,468 @@ +//! Provider-neutral AI usage events. +//! +//! The gateway reports only request metadata and provider-supplied token counts. +//! It never includes prompts, responses, headers, credentials, or provider error bodies. + +use std::collections::VecDeque; +use std::pin::Pin; +use std::sync::Arc; +use std::time::{Duration, Instant, SystemTime}; + +use axum::body::{Body, Bytes}; +use axum::response::Response; +use futures::StreamExt; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use uuid::Uuid; + +/// Receives completed request observations. Implementations must return quickly; +/// inference must never wait for telemetry delivery. A typical implementation uses +/// a bounded channel and drops the event when that channel is full. +pub trait AiUsageObserver: Send + Sync + 'static { + fn observe(&self, event: AiUsageEvent); +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum AiUsageProvider { + AwsBedrock, + GcpVertex, + AzureFoundry, + Anthropic, + #[serde(rename = "openai")] + OpenAi, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum AiUsageClientApi { + #[serde(rename = "openai-chat-completions")] + OpenAiChatCompletions, + #[serde(rename = "openai-responses")] + OpenAiResponses, + AnthropicMessages, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum AiUsageOutcome { + Success, + ProviderError, + GatewayError, + Cancelled, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AiTokenUsage { + pub input_tokens: Option, + pub output_tokens: Option, + pub cache_read_tokens: Option, + pub cache_write_tokens: Option, + pub reasoning_tokens: Option, +} + +#[derive(Debug, Clone)] +pub struct AiUsageEvent { + pub request_id: String, + pub started_at: SystemTime, + pub duration: Duration, + pub binding: String, + pub provider: AiUsageProvider, + pub public_model: String, + pub provider_model: String, + pub client_api: AiUsageClientApi, + pub provider_region: Option, + pub status: u16, + pub outcome: AiUsageOutcome, + pub tokens: AiTokenUsage, +} + +const MAX_USAGE_RESPONSE_BYTES: usize = 1024 * 1024; + +#[derive(Clone)] +pub(crate) struct AiUsageContext { + request_id: String, + started_at: SystemTime, + started: Instant, + binding: String, + provider: AiUsageProvider, + public_model: String, + provider_model: String, + client_api: AiUsageClientApi, + provider_region: Option, +} + +impl AiUsageContext { + pub(crate) fn new( + binding: &str, + provider: AiUsageProvider, + public_model: &str, + provider_model: &str, + client_api: AiUsageClientApi, + provider_region: Option, + ) -> Self { + Self { + request_id: Uuid::new_v4().to_string(), + started_at: SystemTime::now(), + started: Instant::now(), + binding: binding.to_string(), + provider, + public_model: public_model.to_string(), + provider_model: provider_model.to_string(), + client_api, + provider_region, + } + } +} + +struct ObservedBody { + inner: Pin> + Send>>, + observer: Arc, + context: AiUsageContext, + response_tail: VecDeque, + status: u16, + complete: bool, +} + +impl ObservedBody { + fn retain_tail(&mut self, chunk: &[u8]) { + if chunk.len() >= MAX_USAGE_RESPONSE_BYTES { + self.response_tail.clear(); + self.response_tail.extend( + chunk[chunk.len() - MAX_USAGE_RESPONSE_BYTES..] + .iter() + .copied(), + ); + return; + } + let overflow = self + .response_tail + .len() + .saturating_add(chunk.len()) + .saturating_sub(MAX_USAGE_RESPONSE_BYTES); + self.response_tail.drain(..overflow); + self.response_tail.extend(chunk.iter().copied()); + } + + fn finish(&mut self, outcome: AiUsageOutcome, status: u16) { + if self.complete { + return; + } + self.complete = true; + let tokens = if outcome == AiUsageOutcome::Success { + parse_ai_token_usage( + self.response_tail.make_contiguous(), + self.context.client_api, + ) + } else { + AiTokenUsage::default() + }; + let event = AiUsageEvent { + request_id: self.context.request_id.clone(), + started_at: self.context.started_at, + duration: self.context.started.elapsed(), + binding: self.context.binding.clone(), + provider: self.context.provider, + public_model: self.context.public_model.clone(), + provider_model: self.context.provider_model.clone(), + client_api: self.context.client_api, + provider_region: self.context.provider_region.clone(), + status, + outcome, + tokens, + }; + let observer = Arc::clone(&self.observer); + // A faulty optional observer must not turn successful inference into a + // failed response or abort a response-body task. + let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + observer.observe(event); + })); + } +} + +impl Drop for ObservedBody { + fn drop(&mut self) { + if !self.complete { + self.finish(AiUsageOutcome::Cancelled, 499); + } + } +} + +pub(crate) fn observe_response( + response: Response, + observer: Option<&Arc>, + context: AiUsageContext, +) -> Response { + let Some(observer) = observer else { + return response; + }; + let (parts, body) = response.into_parts(); + let status = parts.status.as_u16(); + let state = ObservedBody { + inner: Box::pin(body.into_data_stream()), + observer: Arc::clone(observer), + context, + response_tail: VecDeque::new(), + status, + complete: false, + }; + let stream = futures::stream::unfold(state, |mut state| async move { + match state.inner.next().await { + Some(Ok(chunk)) => { + state.retain_tail(&chunk); + Some((Ok::<_, axum::Error>(chunk), state)) + } + Some(Err(error)) => { + state.finish(AiUsageOutcome::ProviderError, 502); + Some((Err(error), state)) + } + None => { + let outcome = if (200..300).contains(&state.status) { + AiUsageOutcome::Success + } else { + AiUsageOutcome::ProviderError + }; + let status = state.status; + state.finish(outcome, status); + None + } + } + }); + Response::from_parts(parts, Body::from_stream(stream)) +} + +/// Extract token counts from a complete JSON response or from the JSON payloads +/// carried by an SSE response. Unknown response fields are ignored. Missing usage +/// remains `None`; it is never converted to zero. +pub fn parse_ai_token_usage(body: &[u8], client_api: AiUsageClientApi) -> AiTokenUsage { + if let Ok(value) = serde_json::from_slice::(body) { + return usage_from_value(&value, client_api); + } + + let mut usage = AiTokenUsage::default(); + for line in body.split(|byte| *byte == b'\n') { + let line = trim_ascii(line); + let Some(data) = line.strip_prefix(b"data:") else { + continue; + }; + let data = trim_ascii(data); + if data == b"[DONE]" { + continue; + } + if let Ok(value) = serde_json::from_slice::(data) { + merge_usage(&mut usage, usage_from_value(&value, client_api)); + } + } + usage +} + +fn trim_ascii(mut value: &[u8]) -> &[u8] { + while value.first().is_some_and(u8::is_ascii_whitespace) { + value = &value[1..]; + } + while value.last().is_some_and(u8::is_ascii_whitespace) { + value = &value[..value.len() - 1]; + } + value +} + +fn usage_from_value(value: &Value, client_api: AiUsageClientApi) -> AiTokenUsage { + match client_api { + AiUsageClientApi::OpenAiChatCompletions => openai_usage(value.get("usage")), + AiUsageClientApi::OpenAiResponses => { + let response = value.get("response").unwrap_or(value); + openai_responses_usage(response.get("usage")) + } + AiUsageClientApi::AnthropicMessages => anthropic_usage(value), + } +} + +fn openai_usage(value: Option<&Value>) -> AiTokenUsage { + let Some(value) = value else { + return AiTokenUsage::default(); + }; + AiTokenUsage { + input_tokens: uint(value, "prompt_tokens"), + output_tokens: uint(value, "completion_tokens"), + cache_read_tokens: value + .get("prompt_tokens_details") + .and_then(|details| uint(details, "cached_tokens")), + cache_write_tokens: None, + reasoning_tokens: value + .get("completion_tokens_details") + .and_then(|details| uint(details, "reasoning_tokens")), + } +} + +fn openai_responses_usage(value: Option<&Value>) -> AiTokenUsage { + let Some(value) = value else { + return AiTokenUsage::default(); + }; + AiTokenUsage { + input_tokens: uint(value, "input_tokens"), + output_tokens: uint(value, "output_tokens"), + cache_read_tokens: value + .get("input_tokens_details") + .and_then(|details| uint(details, "cached_tokens")), + cache_write_tokens: None, + reasoning_tokens: value + .get("output_tokens_details") + .and_then(|details| uint(details, "reasoning_tokens")), + } +} + +fn anthropic_usage(value: &Value) -> AiTokenUsage { + let usage = value.get("usage").or_else(|| { + value + .get("message") + .and_then(|message| message.get("usage")) + }); + let Some(usage) = usage else { + return AiTokenUsage::default(); + }; + AiTokenUsage { + input_tokens: uint(usage, "input_tokens"), + output_tokens: uint(usage, "output_tokens"), + cache_read_tokens: uint(usage, "cache_read_input_tokens"), + cache_write_tokens: uint(usage, "cache_creation_input_tokens"), + reasoning_tokens: None, + } +} + +fn uint(value: &Value, key: &str) -> Option { + value.get(key).and_then(Value::as_u64) +} + +fn merge_usage(current: &mut AiTokenUsage, next: AiTokenUsage) { + if next.input_tokens.is_some() { + current.input_tokens = next.input_tokens; + } + if next.output_tokens.is_some() { + current.output_tokens = next.output_tokens; + } + if next.cache_read_tokens.is_some() { + current.cache_read_tokens = next.cache_read_tokens; + } + if next.cache_write_tokens.is_some() { + current.cache_write_tokens = next.cache_write_tokens; + } + if next.reasoning_tokens.is_some() { + current.reasoning_tokens = next.reasoning_tokens; + } +} + +#[cfg(test)] +mod tests { + use std::sync::mpsc; + + use super::*; + + struct TestObserver(mpsc::Sender); + + impl AiUsageObserver for TestObserver { + fn observe(&self, event: AiUsageEvent) { + let _ = self.0.send(event); + } + } + + #[test] + fn dropping_an_incomplete_response_stream_observes_cancellation() { + let (sender, receiver) = mpsc::channel(); + let observer: Arc = Arc::new(TestObserver(sender)); + let state = ObservedBody { + inner: Box::pin(futures::stream::pending()), + observer, + context: AiUsageContext::new( + "llm", + AiUsageProvider::OpenAi, + "gpt-5-mini", + "gpt-5-mini", + AiUsageClientApi::OpenAiChatCompletions, + None, + ), + response_tail: VecDeque::new(), + status: 200, + complete: false, + }; + drop(state); + + let event = receiver.try_recv().expect("cancelled usage observation"); + assert_eq!(event.outcome, AiUsageOutcome::Cancelled); + assert_eq!(event.status, 499); + assert_eq!(event.tokens, AiTokenUsage::default()); + } + + #[test] + fn public_usage_identifiers_match_the_gateway_api() { + assert_eq!( + serde_json::to_string(&AiUsageProvider::OpenAi).unwrap(), + "\"openai\"" + ); + assert_eq!( + serde_json::to_string(&AiUsageClientApi::OpenAiChatCompletions).unwrap(), + "\"openai-chat-completions\"" + ); + assert_eq!( + serde_json::to_string(&AiUsageClientApi::OpenAiResponses).unwrap(), + "\"openai-responses\"" + ); + } + + #[test] + fn extracts_openai_non_streaming_usage() { + let usage = parse_ai_token_usage( + br#"{"usage":{"prompt_tokens":12,"completion_tokens":7,"prompt_tokens_details":{"cached_tokens":4},"completion_tokens_details":{"reasoning_tokens":2}}}"#, + AiUsageClientApi::OpenAiChatCompletions, + ); + assert_eq!( + usage, + AiTokenUsage { + input_tokens: Some(12), + output_tokens: Some(7), + cache_read_tokens: Some(4), + cache_write_tokens: None, + reasoning_tokens: Some(2), + } + ); + } + + #[test] + fn extracts_openai_responses_stream_usage() { + let body = br#"event: response.completed +data: {"type":"response.completed","response":{"usage":{"input_tokens":20,"output_tokens":8,"input_tokens_details":{"cached_tokens":5},"output_tokens_details":{"reasoning_tokens":3}}}} + +data: [DONE] +"#; + let usage = parse_ai_token_usage(body, AiUsageClientApi::OpenAiResponses); + assert_eq!(usage.input_tokens, Some(20)); + assert_eq!(usage.output_tokens, Some(8)); + assert_eq!(usage.cache_read_tokens, Some(5)); + assert_eq!(usage.reasoning_tokens, Some(3)); + } + + #[test] + fn merges_anthropic_stream_usage_without_inventing_missing_counts() { + let body = br#"event: message_start +data: {"type":"message_start","message":{"usage":{"input_tokens":30,"cache_creation_input_tokens":6,"cache_read_input_tokens":9}}} + +event: message_delta +data: {"type":"message_delta","usage":{"output_tokens":11}} +"#; + let usage = parse_ai_token_usage(body, AiUsageClientApi::AnthropicMessages); + assert_eq!(usage.input_tokens, Some(30)); + assert_eq!(usage.output_tokens, Some(11)); + assert_eq!(usage.cache_read_tokens, Some(9)); + assert_eq!(usage.cache_write_tokens, Some(6)); + assert_eq!(usage.reasoning_tokens, None); + } + + #[test] + fn malformed_or_missing_usage_is_unknown_not_zero() { + let malformed = parse_ai_token_usage(b"not json", AiUsageClientApi::AnthropicMessages); + let missing = + parse_ai_token_usage(br#"{"id":"message"}"#, AiUsageClientApi::AnthropicMessages); + assert_eq!(malformed, AiTokenUsage::default()); + assert_eq!(missing, AiTokenUsage::default()); + } +}