From 82d35649dbc880d1bc8ffc536a45f4195bcb9ed6 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 01:44:24 +0530 Subject: [PATCH 01/17] feat(bud-auth): realtime session settings, per-modality rates, and live-session reach checks FRD-023 WP-RT2.3 and Q-7. - voice_table config.realtime parses into RealtimeSettings {session_type, defaults, limits, policy}, leniently: a malformed block or field is dropped with a warning, never the endpoint. RealtimePolicy's accessors apply the secure defaults (no MCP tools, no stored prompts). - VoicePricing.rates: a realtime token price is its per-modality rates; a token price without rates is refused (served unpriced) rather than priced at zero. Keys are the closed REALTIME_RATE_KEYS set. - Principal.expires_at carries a JWT caller's verified exp (client-secret lifetime cap). - BudPlane::hash_reaches / subject_reaches / subject_aliases: the revalidation a live realtime session and an ek_bud_ parent make, by snapshot hash or JWT subject. - user_projects:{sub} events evict that subject's cached grants and project_models:* events clear them all, so JWT revocation reaches live sessions without waiting OIDC_AUTHZ_TTL_SECS. - tests/fixtures/realtime_voice_entry.json is byte-identical to budapp's publisher fixture. Co-Authored-By: Claude Opus 5.5 --- bud-auth/src/credentials.rs | 71 +++++++ bud-auth/src/endpoint_config.rs | 130 +++++++++++- bud-auth/src/jwt.rs | 13 ++ bud-auth/src/lib.rs | 7 +- bud-auth/src/runtime.rs | 200 ++++++++++++++++++ .../tests/fixtures/realtime_voice_entry.json | 62 ++++++ bud-auth/tests/realtime_contract.rs | 175 +++++++++++++++ 7 files changed, 654 insertions(+), 4 deletions(-) create mode 100644 bud-auth/tests/fixtures/realtime_voice_entry.json create mode 100644 bud-auth/tests/realtime_contract.rs diff --git a/bud-auth/src/credentials.rs b/bud-auth/src/credentials.rs index 67afe32f..fdd8dceb 100644 --- a/bud-auth/src/credentials.rs +++ b/bud-auth/src/credentials.rs @@ -235,8 +235,30 @@ pub struct VoicePricing { /// How many units `cost_per_unit` covers. `0` is carried as published and prices nothing, /// rather than dividing by it. pub per_units: u64, + /// FRD-023 §5.10: a realtime deployment's per-modality rates, each per `per_units` tokens + /// (`transcription_per_minute` per minute). Keys are restricted to [`REALTIME_RATE_KEYS`]. + /// Empty for every other unit and every non-realtime deployment. + pub rates: BTreeMap, } +/// The rate keys a realtime price may carry (CONTRACTS C1). A closed set: a key outside it is +/// ignored with a warning rather than guessed at, and a component whose rate is absent is +/// recorded as UNPRICED, never as free. +pub const REALTIME_RATE_KEYS: &[&str] = &[ + "input_text", + "input_audio", + "input_image", + "cached_input_text", + "cached_input_audio", + "cached_input_image", + "output_text", + "output_audio", + "transcription_per_minute", + "transcription_input_audio", + "transcription_input_text", + "transcription_output_text", +]; + /// Read a published `pricing` block, or `None` with a warning when it cannot be used. /// /// Tolerant of the shapes a Python publisher produces: numbers or numeric strings, a missing @@ -269,8 +291,15 @@ pub fn parse_pricing(endpoint_id: &str, raw: &serde_json::Value) -> Option u.to_ascii_lowercase(), _ => return refuse("no unit"), }; + let rates = parse_rates(endpoint_id, fields.get("rates")); + // FRD-023: a TOKEN price on the audio plane is a realtime price, and is ONLY its rates. With + // none it would price every response at zero, so it is refused like any unusable price. + if unit == "token" && rates.is_empty() { + return refuse("a token price on a voice endpoint needs per-modality rates"); + } let cost_per_unit = match fields.get("cost_per_unit").and_then(number) { Some(c) if c.is_finite() && c >= 0.0 => c, + None if unit == "token" => 0.0, _ => return refuse("cost_per_unit is not a non-negative number"), }; let per_units = match fields.get("per_units") { @@ -293,9 +322,50 @@ pub fn parse_pricing(endpoint_id: &str, raw: &serde_json::Value) -> Option) -> BTreeMap { + let mut out = BTreeMap::new(); + let fields = match raw { + None | Some(serde_json::Value::Null) => return out, + Some(serde_json::Value::Object(fields)) => fields, + Some(_) => { + tracing::warn!(endpoint_id = %endpoint_id, "voice_table pricing.rates is not an object; ignored"); + return out; + } + }; + for (key, value) in fields { + if !REALTIME_RATE_KEYS.contains(&key.as_str()) { + tracing::warn!( + endpoint_id = %endpoint_id, + field = %key, + "voice_table pricing.rates carries a key this build does not price; ignored" + ); + continue; + } + let rate = match value { + serde_json::Value::Number(n) => n.as_f64(), + serde_json::Value::String(s) => s.trim().parse::().ok(), + _ => None, + }; + match rate { + Some(r) if r.is_finite() && r >= 0.0 => { + out.insert(key.clone(), r); + } + _ => tracing::warn!( + endpoint_id = %endpoint_id, + field = %key, + "voice_table pricing rate is not a non-negative number; ignored" + ), + } + } + out +} + /// A voice endpoint after hydration: the credential is already plaintext. /// /// `PartialEq` but not `Eq`: the config block carries vendor float knobs (pitch, stability, @@ -741,6 +811,7 @@ mod tests { cost_per_unit: 0.0001, currency: Some("USD".into()), per_units: 1, + rates: BTreeMap::new(), }) ); } diff --git a/bud-auth/src/endpoint_config.rs b/bud-auth/src/endpoint_config.rs index d3035c0c..31e83fc4 100644 --- a/bud-auth/src/endpoint_config.rs +++ b/bud-auth/src/endpoint_config.rs @@ -185,6 +185,121 @@ pub struct TranslationSettings { pub partials: Option, } +/// Input-transcription defaults for a realtime session (OpenAI GA `audio.input.transcription`). +#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)] +pub struct RealtimeTranscription { + #[serde(default)] + pub model: Option, + #[serde(default)] + pub language: Option, + #[serde(default)] + pub prompt: Option, +} + +/// What WaaV sends the vendor in its first `session.update`, unless the client overrides it +/// (FRD-023 §5.6: request > deployment > vendor default). +#[derive(Debug, Clone, Default, PartialEq, Deserialize)] +pub struct RealtimeDefaults { + #[serde(default)] + pub voice: Option, + #[serde(default)] + pub instructions: Option, + #[serde(default)] + pub output_modalities: Option>, + /// Passed to the vendor as written: `{"type": "server_vad" | "semantic_vad", …}` or `null`. + /// budapp validated its shape; the vendor owns its semantics. + #[serde(default)] + pub turn_detection: Option, + #[serde(default)] + pub input_transcription: Option, + #[serde(default)] + pub noise_reduction: Option, + #[serde(default)] + pub max_output_tokens: Option, + #[serde(default)] + pub speed: Option, +} + +/// Per-deployment session limits. Absent means the gateway's ceiling applies. +#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)] +pub struct RealtimeLimits { + #[serde(default)] + pub max_session_seconds: Option, + #[serde(default)] + pub idle_timeout_seconds: Option, +} + +/// What a client may change on a realtime session (FRD-023 D-11, S-5). +/// +/// Every field is `Option` like the rest of this module, but the ACCESSORS apply the secure +/// default: a stored prompt, an MCP connector and a trace belong to the VENDOR ORG, which every +/// project sharing the credential shares, so they are off unless the deployment turns them on. +#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)] +pub struct RealtimePolicy { + #[serde(default)] + pub allow_client_instructions: Option, + #[serde(default)] + pub allow_mcp_tools: Option, + #[serde(default)] + pub allow_prompt_references: Option, + #[serde(default)] + pub allow_image_input: Option, + /// Transcription models a client may select. Absent = any; present = exactly these. + #[serde(default)] + pub input_transcription_models: Option>, +} + +impl RealtimePolicy { + pub fn allows_client_instructions(&self) -> bool { + self.allow_client_instructions.unwrap_or(true) + } + + pub fn allows_mcp_tools(&self) -> bool { + self.allow_mcp_tools.unwrap_or(false) + } + + pub fn allows_prompt_references(&self) -> bool { + self.allow_prompt_references.unwrap_or(false) + } + + pub fn allows_image_input(&self) -> bool { + self.allow_image_input.unwrap_or(true) + } + + pub fn allows_transcription_model(&self, model: &str) -> bool { + match &self.input_transcription_models { + None => true, + Some(list) => list.iter().any(|m| m == model), + } + } +} + +/// A realtime (speech-to-speech) deployment's session settings (FRD-023 §5.3). +#[derive(Debug, Clone, Default, PartialEq, Deserialize)] +pub struct RealtimeSettings { + /// `realtime` or `transcription`, derived by budapp from the model's modality. + #[serde(default)] + pub session_type: Option, + #[serde(default)] + pub defaults: Option, + #[serde(default)] + pub limits: Option, + #[serde(default)] + pub policy: Option, +} + +impl RealtimeSettings { + /// The policy block, or an all-default one (every accessor at its secure default). + pub fn policy(&self) -> RealtimePolicy { + self.policy.clone().unwrap_or_default() + } + + /// Whether this deployment serves transcription-only sessions. + pub fn is_transcription(&self) -> bool { + self.session_type.as_deref() == Some("transcription") + } +} + /// The whole `config` object on a voice entry. #[derive(Debug, Clone, Default, PartialEq, Deserialize)] pub struct VoiceEndpointSettings { @@ -194,11 +309,17 @@ pub struct VoiceEndpointSettings { pub stt: Option, #[serde(default)] pub translation: Option, + /// FRD-023: the realtime session block, on `realtime_session` deployments only. + #[serde(default)] + pub realtime: Option, } impl VoiceEndpointSettings { pub fn is_empty(&self) -> bool { - self.tts.is_none() && self.stt.is_none() && self.translation.is_none() + self.tts.is_none() + && self.stt.is_none() + && self.translation.is_none() + && self.realtime.is_none() } /// The tts block, or an all-`None` one. Saves every call site an `unwrap_or_default` clone. @@ -269,7 +390,9 @@ const KNOWN_STT: &[&str] = &[ const KNOWN_TRANSLATION: &[&str] = &["target_languages", "translate_to_english", "partials"]; -const KNOWN_SECTIONS: &[&str] = &["tts", "stt", "translation"]; +const KNOWN_REALTIME: &[&str] = &["session_type", "defaults", "limits", "policy"]; + +const KNOWN_SECTIONS: &[&str] = &["tts", "stt", "translation", "realtime"]; /// Parse a `config` object, warning about anything this build does not model. /// @@ -286,6 +409,7 @@ pub fn parse_endpoint_settings( ("tts", KNOWN_TTS), ("stt", KNOWN_STT), ("translation", KNOWN_TRANSLATION), + ("realtime", KNOWN_REALTIME), ] { if let Some(serde_json::Value::Object(fields)) = sections.get(section) { warn_unmodelled( @@ -318,6 +442,8 @@ pub fn parse_endpoint_settings( stt: section("stt").and_then(|v| parse_section(endpoint_id, "stt", v)), translation: section("translation") .and_then(|v| parse_section(endpoint_id, "translation", v)), + realtime: section("realtime") + .and_then(|v| parse_section(endpoint_id, "realtime", v)), } } } diff --git a/bud-auth/src/jwt.rs b/bud-auth/src/jwt.rs index 4ddd7d90..884631ed 100644 --- a/bud-auth/src/jwt.rs +++ b/bud-auth/src/jwt.rs @@ -504,6 +504,19 @@ impl JwtVerifier { ); } + /// Drop one subject's cached authorization (FRD-023 Q-7): `user_projects:{sub}` changed, so + /// the next check re-reads it instead of serving the old grants for up to `authz_ttl`. + pub fn evict_authz(&self, sub: &str) { + self.authz_cache.remove(sub); + } + + /// Drop every cached authorization: a `project_models:*` blob changed, and which subjects it + /// feeds is not known without reading every `user_projects:*`. Membership changes are rare + /// and a miss costs one store read, so a full clear is the cheap correct answer. + pub fn clear_authz(&self) { + self.authz_cache.clear(); + } + pub fn config(&self) -> &JwtConfig { &self.cfg } diff --git a/bud-auth/src/lib.rs b/bud-auth/src/lib.rs index 04e160b4..5b998d2e 100644 --- a/bud-auth/src/lib.rs +++ b/bud-auth/src/lib.rs @@ -27,9 +27,12 @@ pub mod store; pub mod types; pub use authz::{AuthzTier, Resolution}; -pub use credentials::{CredentialDecryptor, CredentialError, VoiceEndpoint, VoicePricing}; +pub use credentials::{ + CredentialDecryptor, CredentialError, REALTIME_RATE_KEYS, VoiceEndpoint, VoicePricing, +}; pub use endpoint_config::{ - Pronunciation, SttSettings, TranslationSettings, TtsSettings, VoiceDescriptor, + Pronunciation, RealtimeDefaults, RealtimeLimits, RealtimePolicy, RealtimeSettings, + RealtimeTranscription, SttSettings, TranslationSettings, TtsSettings, VoiceDescriptor, VoiceEndpointSettings, }; pub use guards::{Denied, EscalationPermit, KeyShape, MissGuardConfig, MissGuards}; diff --git a/bud-auth/src/runtime.rs b/bud-auth/src/runtime.rs index 5d162673..2ed3ae39 100644 --- a/bud-auth/src/runtime.rs +++ b/bud-auth/src/runtime.rs @@ -31,6 +31,9 @@ pub struct Principal { pub user_id: Option, /// How the caller proved identity. pub via: PrincipalKind, + /// A JWT caller's verified `exp` (unix seconds). `None` for an API key, whose snapshot entry + /// carries no expiry: an expired or deleted key leaves the snapshot instead (FRD-023 §5.8). + pub expires_at: Option, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -185,6 +188,19 @@ impl BudPlane { self.origin.elapsed().as_millis() as u64 + 1, Ordering::Relaxed, ); + // FRD-023 Q-7: a JWT caller's grants are cached per `sub` for `OIDC_AUTHZ_TTL_SECS`. Without + // eviction, removing a user from a project reached WaaV only when that entry aged out, so + // a live realtime session outlived the revocation by up to the TTL. + if let Some(jwt) = &self.jwt { + if let Some(sub) = key.strip_prefix(crate::authz::USER_PROJECTS_PREFIX) { + jwt.evict_authz(sub); + return Ok(()); + } + if key.starts_with(crate::authz::PROJECT_MODELS_PREFIX) { + jwt.clear_authz(); + return Ok(()); + } + } // The overlay is one key, written whole: re-read it on a set, drop it on a delete. if key == hydrate::PUBLISHED_MODEL_INFO_KEY { match event { @@ -236,6 +252,7 @@ impl BudPlane { api_key_id: md.as_ref().and_then(|m| m.api_key_id.clone()), user_id: md.as_ref().and_then(|m| m.user_id.clone()), via: PrincipalKind::ApiKey, + expires_at: None, }); } @@ -273,6 +290,7 @@ impl BudPlane { api_key_id: None, user_id: Some(identity.sub), via: PrincipalKind::Jwt, + expires_at: (identity.exp > 0).then_some(identity.exp), }) } Err(_) => Err(AuthFailure::JwtRejected), @@ -307,6 +325,7 @@ impl BudPlane { api_key_id: md.as_ref().and_then(|m| m.api_key_id.clone()), user_id: md.as_ref().and_then(|m| m.user_id.clone()), via: PrincipalKind::ApiKey, + expires_at: None, }) } Err(_) => Err(AuthFailure::Unauthorized), @@ -733,6 +752,187 @@ impl BudPlane { let entry = jwt.cached_authz(&identity.sub)?; matching(&entry.aliases) } + + /// Whether an API key, known only by its snapshot hash, still reaches `endpoint_id`. + /// + /// FRD-023 revalidation and `ek_bud_` parents: a live session and a client secret hold the + /// HASH, never the raw key, so the lookup is the same `HashMap::get` a direct connect makes. + /// `client_key` is whether the raw key was a `bud_client_*` key, recorded when the raw key was + /// last in hand — such a key also reaches the published overlay. + pub fn hash_reaches( + &self, + hashed: &str, + endpoint_id: &str, + client_key: bool, + ) -> Option { + let own = self.auth.resolve(hashed)?; + let matching = |aliases: &AliasMap| { + aliases + .values() + .find(|m| m.endpoint_id.as_deref() == Some(endpoint_id)) + .cloned() + }; + if client_key && let Some(meta) = matching(&self.published.load_full()) { + return Some(meta); + } + matching(&own) + } + + /// The `__metadata__` of an API key known by its hash. + pub fn hash_metadata(&self, hashed: &str) -> Option { + self.auth.metadata(hashed).map(|m| (*m).clone()) + } + + /// What a JWT subject may reach, through the per-`sub` authz cache; a miss is one store read + /// of `user_projects:{sub}` (FRD-023 D-17). `None` when JWT acceptance is not configured. + pub async fn subject_aliases(&self, sub: &str) -> Option> { + let jwt = self.jwt.as_ref()?; + if let Some(entry) = jwt.cached_authz(sub) { + return Some(entry.aliases); + } + let r = authz::resolve(self.store.as_ref(), sub, self.published.load_full()).await; + if r.tier.is_cacheable() { + jwt.store_authz(sub, Arc::clone(&r.aliases)); + } + Some(r.aliases) + } + + /// Whether a JWT subject still reaches `endpoint_id`. + pub async fn subject_reaches(&self, sub: &str, endpoint_id: &str) -> Option { + let aliases = self.subject_aliases(sub).await?; + aliases + .values() + .find(|m| m.endpoint_id.as_deref() == Some(endpoint_id)) + .cloned() + } + + /// The JWT verifier, when JWT acceptance is configured. + pub fn jwt(&self) -> Option<&Arc> { + self.jwt.as_ref() + } +} + +#[cfg(test)] +mod realtime_reach_tests { + //! FRD-023 D-17 / §5.8: what a live session and an `ek_bud_` parent re-check, by hash or sub. + use super::*; + use crate::store::MemoryStore; + + fn key_blob() -> String { + r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1","model_id":"m1"},"__metadata__":{"api_key_id":"ak1","user_id":"u1","api_key_project_id":"p1"}}"#.to_string() + } + + async fn plane_with(keys: &[(&str, String)], jwt: Option>) -> (Arc, BudPlane) { + let store = Arc::new(MemoryStore::new()); + for (k, v) in keys { + store.set(k, v); + } + let plane = BudPlane::new(Arc::clone(&store) as Arc, jwt); + plane.boot().await.unwrap(); + (store, plane) + } + + struct NoKeys; + #[async_trait::async_trait] + impl crate::jwt::JwksSource for NoKeys { + async fn fetch(&self) -> Result { + Ok(r#"{"keys":[]}"#.to_string()) + } + } + + fn verifier() -> Arc { + let cfg = crate::jwt::JwtConfig::from_lookup(|k| match k { + "OIDC_ISSUER" => Some("https://kc.example/realms/bud".into()), + "OIDC_ALLOWED_CLIENTS" => Some("bud-playground".into()), + _ => None, + }) + .unwrap(); + Arc::new(JwtVerifier::new(cfg, Arc::new(NoKeys))) + } + + #[tokio::test] + async fn a_hash_reaches_its_endpoint_until_the_key_is_revoked() { + let hashed = hash_api_key("bud_rt"); + let key = format!("api_key:{hashed}"); + let (store, plane) = plane_with(&[(&key, key_blob())], None).await; + + let entry = plane.hash_reaches(&hashed, "ep-rt", false).expect("reaches"); + assert_eq!(entry.project_id.as_deref(), Some("p1")); + assert!(plane.hash_reaches(&hashed, "ep-other", false).is_none()); + + store.remove(&key); + plane.on_key_event(&key, KeyEvent::Del).await.unwrap(); + assert!( + plane.hash_reaches(&hashed, "ep-rt", false).is_none(), + "a revoked key still reached its endpoint; a live session would never close" + ); + } + + #[tokio::test] + async fn only_a_client_key_reaches_the_published_overlay_by_hash() { + let hashed = hash_api_key("bud_client_x"); + let (_s, plane) = plane_with( + &[ + (&format!("api_key:{hashed}"), r#"{"__metadata__":{"api_key_id":"ak"}}"#.to_string()), + ( + crate::hydrate::PUBLISHED_MODEL_INFO_KEY, + r#"{"pub-rt":{"endpoint_id":"ep-pub","project_id":"p9"}}"#.to_string(), + ), + ], + None, + ) + .await; + assert!(plane.hash_reaches(&hashed, "ep-pub", true).is_some()); + assert!(plane.hash_reaches(&hashed, "ep-pub", false).is_none()); + } + + #[tokio::test] + async fn a_user_projects_event_evicts_the_subjects_cached_grants() { + let jwt = verifier(); + let (store, plane) = plane_with( + &[ + ("user_projects:sub-1", r#"{"user_id":"u1","projects":["p1"]}"#.to_string()), + ("project_models:p1", r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1"}}"#.to_string()), + ], + Some(Arc::clone(&jwt)), + ) + .await; + assert!(plane.subject_reaches("sub-1", "ep-rt").await.is_some()); + assert!(jwt.cached_authz("sub-1").is_some(), "the resolution is cached"); + + // The user is removed from the project. + store.set("user_projects:sub-1", r#"{"user_id":"u1","projects":[]}"#); + plane + .on_key_event("user_projects:sub-1", KeyEvent::Set) + .await + .unwrap(); + assert!( + jwt.cached_authz("sub-1").is_none(), + "Q-7: the cached grants survived the change" + ); + assert!(plane.subject_reaches("sub-1", "ep-rt").await.is_none()); + } + + #[tokio::test] + async fn a_project_models_event_clears_every_cached_grant() { + let jwt = verifier(); + let (store, plane) = plane_with( + &[ + ("user_projects:sub-1", r#"{"user_id":"u1","projects":["p1"]}"#.to_string()), + ("project_models:p1", r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1"}}"#.to_string()), + ], + Some(Arc::clone(&jwt)), + ) + .await; + assert!(plane.subject_reaches("sub-1", "ep-rt").await.is_some()); + + store.set("project_models:p1", "{}"); + plane + .on_key_event("project_models:p1", KeyEvent::Set) + .await + .unwrap(); + assert!(plane.subject_reaches("sub-1", "ep-rt").await.is_none()); + } } #[cfg(test)] diff --git a/bud-auth/tests/fixtures/realtime_voice_entry.json b/bud-auth/tests/fixtures/realtime_voice_entry.json new file mode 100644 index 00000000..43ae4bad --- /dev/null +++ b/bud-auth/tests/fixtures/realtime_voice_entry.json @@ -0,0 +1,62 @@ +{ + "ep-rt-contract": { + "vendor": "openai", + "endpoints": [ + "realtime_session" + ], + "model": "gpt-realtime-2.1", + "config": { + "realtime": { + "session_type": "realtime", + "defaults": { + "voice": "marin", + "instructions": "You are a helpful assistant.", + "output_modalities": [ + "audio" + ], + "turn_detection": { + "type": "semantic_vad", + "eagerness": "auto" + }, + "input_transcription": { + "model": "gpt-4o-mini-transcribe", + "language": "en" + }, + "noise_reduction": "near_field", + "max_output_tokens": 4096, + "speed": 1.0 + }, + "limits": { + "max_session_seconds": 3600, + "idle_timeout_seconds": 300 + }, + "policy": { + "allow_client_instructions": true, + "allow_mcp_tools": false, + "allow_prompt_references": false, + "allow_image_input": true, + "input_transcription_models": [ + "gpt-4o-mini-transcribe" + ] + } + } + }, + "pricing": { + "unit": "token", + "per_units": 1000000, + "currency": "USD", + "rates": { + "input_text": 4.0, + "input_audio": 32.0, + "input_image": 5.0, + "cached_input_text": 0.4, + "cached_input_audio": 0.4, + "cached_input_image": 0.5, + "output_text": 24.0, + "output_audio": 64.0, + "transcription_per_minute": 0.003 + } + }, + "max_concurrent": 20 + } +} diff --git a/bud-auth/tests/realtime_contract.rs b/bud-auth/tests/realtime_contract.rs new file mode 100644 index 00000000..03019cfc --- /dev/null +++ b/bud-auth/tests/realtime_contract.rs @@ -0,0 +1,175 @@ +//! TC-PUB-15 — the realtime half of the `voice_table` wire contract (FRD-023 §5.3, CONTRACTS C1). +//! +//! `tests/fixtures/realtime_voice_entry.json` is **byte-identical** to +//! `bud-runtime/services/budapp/tests/fixtures/realtime_voice_entry.json`. budapp's publisher test +//! asserts it BUILDS exactly that entry; this suite asserts WaaV PARSES it into the settings the +//! relay enforces. Either half alone lets the two sides drift, and drift here is silent: a +//! `realtime` block WaaV cannot read degrades to vendor defaults and an open policy. + +use bud_auth::credentials::{CredentialDecryptor, parse_voice_blob}; + +const FIXTURE: &str = include_str!("fixtures/realtime_voice_entry.json"); + +fn parse(json: &str) -> bud_auth::VoiceEndpoint { + let map = parse_voice_blob(json, &CredentialDecryptor::disabled()).expect("blob parses"); + map.into_values().next().expect("one endpoint") +} + +#[test] +fn the_fixture_parses_into_a_realtime_endpoint() { + let ep = parse(FIXTURE); + assert_eq!(ep.vendor, "openai"); + assert!(ep.serves("realtime_session")); + assert_eq!(ep.model.as_deref(), Some("gpt-realtime-2.1")); + assert_eq!(ep.policy.max_concurrent, Some(20)); +} + +#[test] +fn the_realtime_block_round_trips() { + let ep = parse(FIXTURE); + let rt = ep.config.realtime.as_ref().expect("realtime block parsed"); + assert_eq!(rt.session_type.as_deref(), Some("realtime")); + + let d = rt.defaults.as_ref().expect("defaults"); + assert_eq!(d.voice.as_deref(), Some("marin")); + assert_eq!(d.instructions.as_deref(), Some("You are a helpful assistant.")); + assert_eq!(d.output_modalities.as_deref(), Some(&["audio".to_string()][..])); + assert_eq!( + d.turn_detection, + Some(serde_json::json!({"type": "semantic_vad", "eagerness": "auto"})) + ); + let tr = d.input_transcription.as_ref().expect("input_transcription"); + assert_eq!(tr.model.as_deref(), Some("gpt-4o-mini-transcribe")); + assert_eq!(tr.language.as_deref(), Some("en")); + assert_eq!(d.noise_reduction.as_deref(), Some("near_field")); + assert_eq!(d.max_output_tokens, Some(4096)); + assert_eq!(d.speed, Some(1.0)); + + let l = rt.limits.as_ref().expect("limits"); + assert_eq!(l.max_session_seconds, Some(3600)); + assert_eq!(l.idle_timeout_seconds, Some(300)); + + let p = rt.policy.clone().unwrap_or_default(); + assert!(p.allows_client_instructions()); + assert!(!p.allows_mcp_tools()); + assert!(!p.allows_prompt_references()); + assert!(p.allows_image_input()); + assert!(p.allows_transcription_model("gpt-4o-mini-transcribe")); + assert!(!p.allows_transcription_model("gpt-4o-transcribe")); +} + +#[test] +fn the_token_price_carries_its_rates_and_no_cost_per_unit() { + let ep = parse(FIXTURE); + let pricing = ep.pricing.expect("a token price with rates is usable"); + assert_eq!(pricing.unit, "token"); + assert_eq!(pricing.per_units, 1_000_000); + assert_eq!(pricing.rates.get("input_audio"), Some(&32.0)); + assert_eq!(pricing.rates.get("output_audio"), Some(&64.0)); + assert_eq!(pricing.rates.get("cached_input_audio"), Some(&0.4)); + assert_eq!(pricing.rates.get("transcription_per_minute"), Some(&0.003)); + assert_eq!(pricing.rates.len(), 9); +} + +/// FRD-023 §5.3: "a malformed block drops that block with a warning, never the endpoint". +#[test] +fn a_malformed_realtime_block_drops_the_block_not_the_endpoint() { + let blob = serde_json::json!({"ep-1": { + "vendor": "openai", + "endpoints": ["realtime_session"], + "model": "gpt-realtime-2.1", + "config": {"realtime": "not an object"} + }}) + .to_string(); + let ep = parse(&blob); + assert!(ep.serves("realtime_session"), "the endpoint must survive"); + assert!(ep.config.realtime.is_none()); +} + +#[test] +fn a_malformed_field_inside_realtime_drops_only_that_field() { + let blob = serde_json::json!({"ep-1": { + "vendor": "openai", + "endpoints": ["realtime_session"], + "config": {"realtime": { + "session_type": "realtime", + "defaults": {"speed": "fast"}, + "policy": {"allow_mcp_tools": true} + }} + }}) + .to_string(); + let ep = parse(&blob); + let rt = ep.config.realtime.expect("block kept"); + assert_eq!(rt.session_type.as_deref(), Some("realtime")); + assert!(rt.defaults.is_none(), "the malformed defaults block is dropped"); + assert!(rt.policy.unwrap().allows_mcp_tools(), "the good sibling survives"); +} + +#[test] +fn a_token_price_without_rates_is_unusable_not_zero() { + // A token price with no rates would price every response at zero. The endpoint is served + // unpriced instead, and the turn says so (FRD-023 §5.10). + let blob = serde_json::json!({"ep-1": { + "vendor": "openai", + "endpoints": ["realtime_session"], + "pricing": {"unit": "token", "per_units": 1000000, "currency": "USD"} + }}) + .to_string(); + assert!(parse(&blob).pricing.is_none()); +} + +#[test] +fn unknown_rate_keys_are_ignored_not_fatal() { + let blob = serde_json::json!({"ep-1": { + "vendor": "openai", + "endpoints": ["realtime_session"], + "pricing": {"unit": "token", "per_units": 1000000, + "rates": {"input_audio": 32, "output_audio": 64, "input_video": 1, "output_text": "24"}} + }}) + .to_string(); + let pricing = parse(&blob).pricing.expect("usable"); + assert_eq!(pricing.rates.get("input_video"), None); + assert_eq!(pricing.rates.get("output_text"), Some(&24.0), "numeric strings are numbers"); + assert_eq!(pricing.rates.len(), 3); +} + +#[test] +fn a_minute_price_on_session_time_needs_no_rates() { + let blob = serde_json::json!({"ep-1": { + "vendor": "openai", + "endpoints": ["realtime_session"], + "pricing": {"unit": "minute", "cost_per_unit": 0.06, "per_units": 1} + }}) + .to_string(); + let pricing = parse(&blob).pricing.expect("usable"); + assert_eq!(pricing.unit, "minute"); + assert!(pricing.rates.is_empty()); +} + +#[test] +fn an_endpoint_without_realtime_config_has_none() { + let blob = serde_json::json!({"ep-1": {"vendor": "deepgram", "endpoints": ["text_to_speech"]}}) + .to_string(); + assert!(parse(&blob).config.realtime.is_none()); +} + +/// S-5: a deployment that says nothing about vendor-org resources does not expose them. The +/// fixture sets every flag explicitly, so this is the case that pins the DEFAULTS. +#[test] +fn absent_policy_fields_take_the_secure_defaults() { + let blob = serde_json::json!({"ep-1": { + "vendor": "openai", + "endpoints": ["realtime_session"], + "config": {"realtime": {"session_type": "realtime", "policy": {}}} + }}) + .to_string(); + let policy = parse(&blob).config.realtime.expect("block").policy(); + assert!(!policy.allows_mcp_tools(), "MCP tools reach the vendor org"); + assert!(!policy.allows_prompt_references(), "stored prompts belong to the vendor org"); + assert!(policy.allows_client_instructions()); + assert!(policy.allows_image_input()); + assert!(policy.allows_transcription_model("anything"), "no allowlist = any model"); + + let none = bud_auth::RealtimePolicy::default(); + assert!(!none.allows_mcp_tools() && !none.allows_prompt_references()); +} From 661f2f613b40bda67e98bcfd6d9b6e962c06be87 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 03:13:20 +0530 Subject: [PATCH 02/17] fix(gateway): no process vendor keys on Bud-mode socket paths (FRD-023 RT0) Under the Bud control plane every vendor credential comes from the deployment (voice_table). Four socket paths still reached for the PROCESS's keys, latent only because the Bud chart sets none; the first operator to add one to "make realtime work" would have exposed it: - X-1 /ws conversation_config: a client-chosen base_url with api_key omitted sent the platform's OPENAI_API_KEY to that host. base_url, api_key, reasoning_base_url and reasoning_api_key are now refused in Bud mode, and LlmClient never falls back to the environment (neither ${VAR} nor the default env var) while the process is in Bud mode. - X-2 /ws DAG: an inline dag_config.definition is refused, and a DAG node's own credential (literal or ${VAR}) is refused in Bud mode. - X-3 native /realtime: a config message is refused with deployment_required; Bud deployments are served by /v1/realtime?model=. - X-4 /ws STT/TTS legs: a provider-only leg no longer falls back to config.get_api_key. - F-2: a second native /realtime config is refused instead of replacing the provider without disconnect(). process_in_bud_mode() is set once when BudMode starts, so a code path with no AppState in reach cannot forget the rule. Two existing tests asserted the X-4 fallback; they now assert its refusal (TC-SEC-06). Every guard (TC-SEC-01..06, 08) was seen failing with its check removed. Co-Authored-By: Claude Opus 5.5 --- gateway/src/auth/bud_mode.rs | 24 +++ gateway/src/core/llm/mod.rs | 171 +++++++++++++++-- gateway/src/dag/nodes/llm.rs | 3 + gateway/src/dag/nodes/provider.rs | 80 ++++++++ gateway/src/handlers/realtime/handler.rs | 186 ++++++++++++++++++ gateway/src/handlers/ws/config_handler.rs | 218 +++++++++++++++++++++- gateway/src/lib.rs | 3 + gateway/src/test_support.rs | 94 ++++++++++ 8 files changed, 752 insertions(+), 27 deletions(-) create mode 100644 gateway/src/test_support.rs diff --git a/gateway/src/auth/bud_mode.rs b/gateway/src/auth/bud_mode.rs index b60190c3..ecb51859 100644 --- a/gateway/src/auth/bud_mode.rs +++ b/gateway/src/auth/bud_mode.rs @@ -21,6 +21,29 @@ use bud_auth::{AuthFailure, BudPlane, ControlPlaneStore, JwtConfig, JwtVerifier, use crate::auth::context::Auth; +/// Whether THIS PROCESS serves Bud deployments (FRD-023 RT0, S-1). +/// +/// Set once, when [`BudMode::start`] succeeds, and never cleared: a WaaV process is in Bud mode +/// for its whole life or not at all. It exists for the code paths that fetch a vendor key from the +/// process environment and have no `AppState` in reach — the LLM client's `${VAR}` / default-env +/// fallback and DAG node credentials. In Bud mode every one of those would spend (or, through a +/// client-chosen `base_url`, exfiltrate) a platform key on behalf of an arbitrary tenant, so they +/// consult this flag rather than trusting each call site to pass one down. +static PROCESS_IN_BUD_MODE: std::sync::atomic::AtomicBool = + std::sync::atomic::AtomicBool::new(false); + +/// See [`PROCESS_IN_BUD_MODE`]. +pub fn process_in_bud_mode() -> bool { + PROCESS_IN_BUD_MODE.load(std::sync::atomic::Ordering::SeqCst) +} + +/// Mark the process as serving Bud deployments. Called by [`BudMode::start`]; public so an +/// integration test can put a test binary into the same state. +#[doc(hidden)] +pub fn mark_process_in_bud_mode() { + PROCESS_IN_BUD_MODE.store(true, std::sync::atomic::Ordering::SeqCst); +} + /// Everything needed to stand the plane up, read from the environment. pub struct BudModeConfig { pub redis_url: String, @@ -143,6 +166,7 @@ impl BudMode { let stats = plane.boot().await.map_err(|e| { format!("initial control-plane hydration failed ({e}); refusing to start with an empty auth map") })?; + mark_process_in_bud_mode(); tracing::info!( api_keys = stats.api_keys, skipped = stats.skipped, diff --git a/gateway/src/core/llm/mod.rs b/gateway/src/core/llm/mod.rs index 22fabcd2..544217eb 100644 --- a/gateway/src/core/llm/mod.rs +++ b/gateway/src/core/llm/mod.rs @@ -518,6 +518,53 @@ pub struct LlmClientConfig { /// applies the floor in `ConversationConfig::to_client_config`). #[serde(default, skip_serializing_if = "Option::is_none")] pub reasoning_effort: Option, + /// Whether a missing key may fall back to the process environment. Never (de)serialised: a + /// client-supplied config must not be able to turn it on. Bud mode forces it off regardless + /// (FRD-023 RT0). + #[serde(skip, default = "default_allow_env_fallback")] + pub allow_env_fallback: bool, +} + +/// The key-resolution rule, over an injected environment lookup so it can be tested without +/// touching the process environment. `env_allowed == false` means the lookup is never called. +pub(crate) fn resolve_llm_api_key( + per_call: Option<&str>, + config_key: Option<&str>, + env_allowed: bool, + default_env_key: &str, + env: impl Fn(&str) -> Option, +) -> Option { + if let Some(key) = per_call { + return Some(key.to_string()); + } + if let Some(key) = config_key { + if key.starts_with("${") && key.ends_with('}') { + if !env_allowed { + warn!( + "Refused a ${{VAR}} API-key reference: vendor keys never come from the environment in Bud mode" + ); + return None; + } + let var_name = &key[2..key.len() - 1]; + if !ALLOWED_ENV_VARS.contains(&var_name) { + warn!( + var_name = %var_name, + "Blocked access to non-whitelisted environment variable" + ); + return None; + } + return env(var_name); + } + return Some(key.to_string()); + } + if !env_allowed { + return None; + } + env(default_env_key) +} + +fn default_allow_env_fallback() -> bool { + true } impl Default for LlmClientConfig { @@ -550,6 +597,7 @@ impl Default for LlmClientConfig { extra: HashMap::new(), provider_kind: None, reasoning_effort: None, + allow_env_fallback: true, } } } @@ -836,27 +884,20 @@ impl LlmClient { /// Priority: per-call key > config key (literal or `${ENV_VAR}`) > the /// active vendor's default env var (`OPENAI_API_KEY` / /// `ANTHROPIC_API_KEY` / `GOOGLE_AI_API_KEY`). + /// + /// FRD-023 RT0 (X-1): in Bud mode the process environment is NEVER read — neither a + /// `${VAR}` reference nor the default env var. A client-chosen `base_url` with the key + /// omitted used to send the platform's `OPENAI_API_KEY` to that host. pub fn resolve_api_key(&self, per_call: Option<&str>) -> Option { - if let Some(key) = per_call { - return Some(key.to_string()); - } - - if let Some(key) = &self.config.api_key { - if key.starts_with("${") && key.ends_with('}') { - let var_name = &key[2..key.len() - 1]; - if !ALLOWED_ENV_VARS.contains(&var_name) { - warn!( - var_name = %var_name, - "Blocked access to non-whitelisted environment variable" - ); - return None; - } - return std::env::var(var_name).ok(); - } - return Some(key.clone()); - } - - std::env::var(self.adapter.default_env_key()).ok() + let env_allowed = + self.config.allow_env_fallback && !crate::auth::bud_mode::process_in_bud_mode(); + resolve_llm_api_key( + per_call, + self.config.api_key.as_deref(), + env_allowed, + self.adapter.default_env_key(), + |name| std::env::var(name).ok(), + ) } /// Render a vendor request via the adapter and apply the operator's extra @@ -1891,3 +1932,93 @@ mod tests { assert_eq!(find_utf8_boundary(bytes), 0); } } + +#[cfg(test)] +mod frd023_key_tests { + //! TC-SEC-02: in Bud mode no path that builds an `LlmClient` reads a key from the environment. + use super::resolve_llm_api_key; + use std::cell::RefCell; + + fn canary(reads: &RefCell>) -> impl Fn(&str) -> Option + '_ { + move |name| { + reads.borrow_mut().push(name.to_string()); + Some("sk-canary".to_string()) + } + } + + #[test] + fn tc_sec_02_no_env_fallback_when_env_is_not_allowed() { + let reads = RefCell::new(Vec::new()); + assert_eq!( + resolve_llm_api_key(None, None, false, "OPENAI_API_KEY", canary(&reads)), + None + ); + assert_eq!( + resolve_llm_api_key( + None, + Some("${OPENAI_API_KEY}"), + false, + "OPENAI_API_KEY", + canary(&reads) + ), + None + ); + assert!( + reads.borrow().is_empty(), + "the environment was read: {:?}", + reads.borrow() + ); + } + + #[test] + fn tc_sec_02_an_explicit_key_still_wins() { + let reads = RefCell::new(Vec::new()); + assert_eq!( + resolve_llm_api_key( + Some("bud_caller"), + None, + false, + "OPENAI_API_KEY", + canary(&reads) + ) + .as_deref(), + Some("bud_caller") + ); + assert_eq!( + resolve_llm_api_key( + None, + Some("literal"), + false, + "OPENAI_API_KEY", + canary(&reads) + ) + .as_deref(), + Some("literal") + ); + assert!(reads.borrow().is_empty()); + } + + #[test] + fn standalone_mode_keeps_its_env_fallback() { + let reads = RefCell::new(Vec::new()); + assert_eq!( + resolve_llm_api_key(None, None, true, "OPENAI_API_KEY", canary(&reads)).as_deref(), + Some("sk-canary") + ); + assert_eq!(reads.borrow().as_slice(), ["OPENAI_API_KEY".to_string()]); + } + + #[test] + fn allow_env_fallback_cannot_be_set_from_a_client_config() { + let cfg: super::LlmClientConfig = serde_json::from_value(serde_json::json!({ + "base_url": "https://api.openai.com/v1", "model": "m", "allow_env_fallback": false + })) + .unwrap(); + assert!( + cfg.allow_env_fallback, + "the field is not client-settable in either direction" + ); + let out = serde_json::to_value(&cfg).unwrap(); + assert!(out.get("allow_env_fallback").is_none()); + } +} diff --git a/gateway/src/dag/nodes/llm.rs b/gateway/src/dag/nodes/llm.rs index 8eff6797..8c3d7f0b 100644 --- a/gateway/src/dag/nodes/llm.rs +++ b/gateway/src/dag/nodes/llm.rs @@ -204,6 +204,9 @@ impl From for LlmClientConfig { extra: c.extra, provider_kind: c.provider_kind, reasoning_effort: c.reasoning_effort, + // Env fallback stays the client default; Bud mode refuses it process-wide + // (`resolve_api_key`, FRD-023 RT0). + allow_env_fallback: true, } } } diff --git a/gateway/src/dag/nodes/provider.rs b/gateway/src/dag/nodes/provider.rs index 61cc7ab3..69cfc1fd 100644 --- a/gateway/src/dag/nodes/provider.rs +++ b/gateway/src/dag/nodes/provider.rs @@ -34,7 +34,30 @@ use crate::dag::error::{DAGError, DAGResult}; /// Without this, STT/TTS provider nodes built `STTConfig`/`TTSConfig` with an EMPTY `api_key`, so a /// DAG could never authenticate to a real vendor — the node failed with "API key is required". pub(crate) fn resolve_node_credential(config: &serde_json::Value, field: &str) -> Option { + resolve_node_credential_with(config, field, crate::auth::bud_mode::process_in_bud_mode()) +} + +/// [`resolve_node_credential`] with Bud mode explicit, so the rule is testable without a +/// process-wide flag. +/// +/// FRD-023 RT0 (X-2): in Bud mode a DAG node carries NO credential of its own — neither a literal +/// key (a tenant's BYOK that bypasses attribution, quota and billing) nor a `${VAR}` reference +/// (the platform's key, spent on behalf of whoever wrote the node). Vendor credentials come from +/// the Bud deployment the node addresses. +pub(crate) fn resolve_node_credential_with( + config: &serde_json::Value, + field: &str, + bud_mode: bool, +) -> Option { let raw = config.get(field)?.as_str()?; + if bud_mode { + warn!( + field = %field, + "DAG node config: refused a node-level credential; in Bud mode vendor credentials \ + come from the addressed deployment (FRD-023 RT0)" + ); + return None; + } if let Some(var) = raw.strip_prefix("${").and_then(|s| s.strip_suffix('}')) { let looks_like_credential = !var.is_empty() && var @@ -70,6 +93,11 @@ fn resolve_configured_node_credential( return Ok(None); } + if crate::auth::bud_mode::process_in_bud_mode() { + return Err(bud_mode_node_credential_error( + node_id, provider, kind, field, + )); + } match resolve_node_credential(config, field) { Some(value) if !value.trim().is_empty() => Ok(Some(value)), _ => Err(DAGError::MissingConfiguration(format!( @@ -79,6 +107,20 @@ fn resolve_configured_node_credential( } } +/// The refusal a Bud-mode DAG node with its own credential gets (FRD-023 RT0, TC-SEC-04). +pub(crate) fn bud_mode_node_credential_error( + node_id: &str, + provider: &str, + kind: &str, + field: &str, +) -> DAGError { + DAGError::MissingConfiguration(format!( + "{kind} provider node '{node_id}' ({provider}) sets config.{field}, which this gateway \ + does not accept: in Bud mode vendor credentials come from a Bud deployment, never from a \ + DAG node or the process environment (FRD-023 RT0)" + )) +} + /// Callback bridge for TTS provider to DAG node /// /// This struct implements the `AudioCallback` trait and bridges @@ -2433,3 +2475,41 @@ mod session_realtime_tests { ); } } + +#[cfg(test)] +mod frd023_node_credential_tests { + //! TC-SEC-04: in Bud mode a DAG node's own credential — literal or `${VAR}` — is refused. + use super::*; + + #[test] + fn tc_sec_04_bud_mode_refuses_literal_and_env_credentials() { + let literal = serde_json::json!({"api_key": "sk-literal"}); + let env_ref = serde_json::json!({"api_key": "${OPENAI_API_KEY}"}); + assert_eq!( + resolve_node_credential_with(&literal, "api_key", true), + None + ); + assert_eq!( + resolve_node_credential_with(&env_ref, "api_key", true), + None + ); + } + + #[test] + fn standalone_keeps_literal_credentials() { + let literal = serde_json::json!({"api_key": "sk-literal"}); + assert_eq!( + resolve_node_credential_with(&literal, "api_key", false).as_deref(), + Some("sk-literal") + ); + } + + #[test] + fn tc_sec_04_the_refusal_names_bud_mode_not_a_missing_key() { + let err = bud_mode_node_credential_error("n1", "deepgram", "STT", "api_key").to_string(); + assert!( + err.contains("Bud deployment") && err.contains("FRD-023"), + "{err}" + ); + } +} diff --git a/gateway/src/handlers/realtime/handler.rs b/gateway/src/handlers/realtime/handler.rs index d855dbeb..aefdedda 100644 --- a/gateway/src/handlers/realtime/handler.rs +++ b/gateway/src/handlers/realtime/handler.rs @@ -589,6 +589,43 @@ async fn handle_config( app_state: &Arc, trace_parent: &str, ) -> bool { + // F-2 (FRD-023 WP-RT0.5): one config per session, as on `/ws`. A second config used to replace + // the provider without `disconnect()`, leaking the first upstream socket and its tasks. + if realtime_provider.is_some() { + warn!("Rejecting a second config on an already-configured realtime session"); + send_realtime_with_policy( + message_tx, + RealtimeMessageRoute::Outgoing(RealtimeOutgoingMessage::Error { + code: Some("session_already_configured".to_string()), + message: "Session already configured — open a new connection to reconfigure \ + (one config message per session)" + .to_string(), + }), + ) + .await; + return true; + } + + // FRD-023 RT0 (X-3): under the Bud control plane this native path would otherwise spend the + // PROCESS's vendor key for whichever tenant connected. Bud deployments are served by the + // OpenAI-compatible `/v1/realtime?model=`, which takes the deployment's own + // credential from `voice_table`. + if app_state.bud_mode.is_some() { + warn!("Refusing a native realtime config in Bud mode"); + send_realtime_with_policy( + message_tx, + RealtimeMessageRoute::Outgoing(RealtimeOutgoingMessage::Error { + code: Some("deployment_required".to_string()), + message: "This gateway serves Bud deployments: connect to \ + /v1/realtime?model= (OpenAI Realtime \ + protocol) instead of sending a native config (FRD-023)." + .to_string(), + }), + ) + .await; + return true; + } + // P3: resolve a server-side ALIAS into the session config BEFORE the provider / // credential is selected. Definitions are server-config-only (SSRF-safe); explicit // client fields win. Unknown alias is non-fatal (proceed + advisory). This mirrors @@ -1389,3 +1426,152 @@ mod tests { } } } + +#[cfg(test)] +mod frd023_native_tests { + //! TC-SEC-05 / TC-SEC-08 on the native `/realtime` path. + use super::*; + use crate::core::realtime::*; + use std::sync::atomic::{AtomicBool, Ordering}; + + struct ConnectedRt(Arc); + + #[async_trait::async_trait] + impl BaseRealtime for ConnectedRt { + fn new(_c: RealtimeConfig) -> RealtimeResult { + unreachable!() + } + async fn connect(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn disconnect(&mut self) -> RealtimeResult<()> { + self.0.store(true, Ordering::SeqCst); + Ok(()) + } + fn is_ready(&self) -> bool { + true + } + fn get_connection_state(&self) -> ConnectionState { + ConnectionState::Connected + } + async fn send_audio(&mut self, _a: bytes::Bytes) -> RealtimeResult<()> { + Ok(()) + } + async fn send_text(&mut self, _t: &str) -> RealtimeResult<()> { + Ok(()) + } + async fn create_response(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn cancel_response(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn commit_audio_buffer(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn clear_audio_buffer(&mut self) -> RealtimeResult<()> { + Ok(()) + } + fn on_transcript(&mut self, _c: TranscriptCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_audio(&mut self, _c: AudioOutputCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_error(&mut self, _c: RealtimeErrorCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_function_call(&mut self, _c: FunctionCallCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_speech_event(&mut self, _c: SpeechEventCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_response_done(&mut self, _c: ResponseDoneCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_reconnection(&mut self, _c: ReconnectionCallback) -> RealtimeResult<()> { + Ok(()) + } + async fn update_session(&mut self, _c: RealtimeConfig) -> RealtimeResult<()> { + Ok(()) + } + async fn submit_function_result(&mut self, _id: &str, _r: &str) -> RealtimeResult<()> { + Ok(()) + } + fn get_provider_info(&self) -> serde_json::Value { + serde_json::json!({}) + } + } + + fn config(provider: &str) -> RealtimeSessionConfig { + serde_json::from_value(serde_json::json!({"provider": provider})).expect("config") + } + + fn error_code(rx: &mut mpsc::Receiver) -> Option { + match rx.try_recv() { + Ok(RealtimeMessageRoute::Outgoing(RealtimeOutgoingMessage::Error { code, .. })) => code, + _ => None, + } + } + + #[tokio::test] + #[serial_test::serial] + async fn tc_sec_05_native_config_in_bud_mode_never_reads_process_keys() { + let mut cfg = crate::test_support::minimal_config(); + cfg.openai_api_key = Some("sk-canary".to_string()); + let mut app_state = AppState::new(cfg).await; + let bud = crate::test_support::bud_state(&[]).await; + Arc::get_mut(&mut app_state).unwrap().bud_mode = bud.bud_mode.clone(); + + let (tx, mut rx) = mpsc::channel(8); + let mut provider: Option> = None; + let mut session_id = None; + + handle_config( + config("openai"), + &mut provider, + &mut session_id, + &tx, + &app_state, + "", + ) + .await; + + assert_eq!(error_code(&mut rx).as_deref(), Some("deployment_required")); + assert!( + provider.is_none(), + "no upstream provider may be built in Bud mode" + ); + assert!(session_id.is_none()); + } + + #[tokio::test] + #[serial_test::serial] + async fn tc_sec_08_a_second_config_is_refused_and_the_first_provider_kept() { + let app_state = AppState::new(crate::test_support::minimal_config()).await; + let disconnected = Arc::new(AtomicBool::new(false)); + let mut provider: Option> = + Some(Box::new(ConnectedRt(Arc::clone(&disconnected)))); + let mut session_id = Some("sess-1".to_string()); + let (tx, mut rx) = mpsc::channel(8); + + handle_config( + config("openai"), + &mut provider, + &mut session_id, + &tx, + &app_state, + "", + ) + .await; + + assert_eq!( + error_code(&mut rx).as_deref(), + Some("session_already_configured") + ); + assert!(provider.is_some(), "the first provider is still in place"); + assert!(!disconnected.load(Ordering::SeqCst), "and still connected"); + assert_eq!(session_id.as_deref(), Some("sess-1")); + } +} diff --git a/gateway/src/handlers/ws/config_handler.rs b/gateway/src/handlers/ws/config_handler.rs index b0fb9906..31abbb25 100644 --- a/gateway/src/handlers/ws/config_handler.rs +++ b/gateway/src/handlers/ws/config_handler.rs @@ -188,6 +188,17 @@ pub async fn handle_config_message( } } + // FRD-023 RT0 (X-1, X-2): in Bud mode a client never chooses where a platform-side call goes + // or which credential it carries. Refused before anything is built, so nothing is dialled. + if app_state.bud_mode.is_some() + && let Some(refusal) = + bud_mode_config_refusal(conversation_ws_config.as_ref(), dag_ws_config.as_ref()) + { + warn!(reason = %refusal, "Refusing a Bud-mode /ws config"); + send_error(message_tx, refusal).await; + return true; + } + // P3: resolve a server-side ALIAS into the session config BEFORE any provider // construction. The alias supplies DEFAULTS; explicit client fields above always // win (handled inside `splice_alias`). Definitions are server-config-only, so the @@ -618,6 +629,55 @@ pub async fn handle_config_message( true } +/// What a Bud-mode `/ws` config may not carry (FRD-023 RT0, FR-WS-2, FR-WS-3), or `None`. +/// +/// * `conversation_config.base_url` / `api_key` / `reasoning_base_url` / `reasoning_api_key` — the +/// voice agent's LLM leg is a Bud chat deployment reached through budgateway with the caller's +/// own credential (RT6). A client-chosen host with the key omitted used to receive the +/// platform's `OPENAI_API_KEY` (X-1). +/// * an inline `dag_config.definition` — its nodes could carry literal keys or `${VAR}` +/// references (X-2). Server templates are the Bud-mode DAG. +pub(crate) fn bud_mode_config_refusal( + conversation: Option<&ConversationWebSocketConfig>, + dag: Option<&DAGWebSocketConfig>, +) -> Option { + if let Some(conv) = conversation { + let mut fields: Vec<&str> = Vec::new(); + if !conv.base_url.trim().is_empty() { + fields.push("base_url"); + } + if conv + .api_key + .as_deref() + .is_some_and(|k| !k.trim().is_empty()) + { + fields.push("api_key"); + } + if conv.reasoning_base_url.is_some() { + fields.push("reasoning_base_url"); + } + if conv.reasoning_api_key.is_some() { + fields.push("reasoning_api_key"); + } + if !fields.is_empty() { + return Some(format!( + "conversation_config.{} is not accepted by this gateway. The voice agent's LLM leg \ + is a Bud chat deployment reached through the Bud gateway with your own \ + credential: remove the field and name the deployment in `model` (FRD-023 RT6).", + fields.join(", conversation_config.") + )); + } + } + if dag.is_some_and(|d| d.definition.is_some()) { + return Some( + "dag_config.definition is not accepted by this gateway: an inline DAG can carry vendor \ + credentials of its own. Use a server template (dag_config.template) (FRD-023 RT0)." + .to_string(), + ); + } + None +} + /// Initialize the built-in conversation loop for a session (plan W-O2). /// /// Constructs a [`ConversationOrchestrator`] (validating the client-supplied LLM @@ -1275,6 +1335,26 @@ async fn resolve_provider_api_key( /v1/audio/speech with `model` set to your Bud endpoint name." ) } + // FRD-023 RT0 (X-4): under the Bud control plane a socket leg NEVER uses a key from the + // process configuration — that key is the platform's, and every tenant would spend it + // unattributed. The leg's credential comes from the deployment it addresses (RT6). + None if !allow_client_keys => { + warn!( + provider = %provider, + role = %role, + "Refused a provider-only socket leg: no process vendor keys in Bud mode" + ); + format!( + "{role}_config names provider '{provider}' but no Bud deployment. This gateway \ + serves Bud deployments only: set {role}_config.model to the name of your {kind} \ + deployment and its credential is used (FRD-023 RT6).", + kind = if role == "stt" { + "transcription" + } else { + "text-to-speech" + }, + ) + } None => match config.get_api_key(provider) { Ok(key) => return Some(key), Err(error_msg) => error_msg, @@ -3452,8 +3532,11 @@ mod tests { assert!(next_error(&mut rx).is_none(), "BYOK is not an error"); } + /// TC-SEC-06 (FRD-023 X-4). This test used to assert the opposite — that an empty client key + /// under Bud mode fell back to the SERVER's vendor key. That fallback is the exposure: every + /// tenant spent the platform key, unattributed. #[tokio::test] - async fn test_empty_client_api_key_falls_back_under_bud_mode() { + async fn tc_sec_06_empty_client_key_under_bud_mode_never_reaches_the_process_key() { let (tx, mut rx) = mpsc::channel(4); let empty = String::new(); @@ -3467,22 +3550,49 @@ mod tests { ) .await; - assert_eq!(key.as_deref(), Some("dg-server-key")); + assert_eq!(key, None, "the process's deepgram key was used in Bud mode"); + let message = next_error(&mut rx).expect("the refusal must reach the client"); assert!( - next_error(&mut rx).is_none(), - "an empty key bypasses nothing and must not fail the session" + message.contains("stt_config.model"), + "must say how to address a deployment: {message}" ); + assert!(!message.contains("dg-server-key")); } #[tokio::test] - async fn test_absent_client_api_key_falls_back_under_bud_mode() { + async fn tc_sec_06_absent_client_key_under_bud_mode_never_reaches_the_process_key() { + for role in ["stt", "tts"] { + let (tx, mut rx) = mpsc::channel(4); + + let key = resolve_provider_api_key( + None, + "deepgram", + role, + false, + &config_with_deepgram_key(), + &tx, + ) + .await; + + assert_eq!( + key, None, + "the process's vendor key was used for a Bud-mode {role} leg" + ); + let message = next_error(&mut rx).expect("the refusal must reach the client"); + assert!(message.contains("FRD-023"), "{message}"); + assert!(!message.contains("dg-server-key")); + } + } + + #[tokio::test] + async fn test_absent_client_key_falls_back_to_server_config_in_standalone_mode() { let (tx, mut rx) = mpsc::channel(4); let key = resolve_provider_api_key( None, "deepgram", "stt", - false, + true, &config_with_deepgram_key(), &tx, ) @@ -3500,7 +3610,7 @@ mod tests { None, "elevenlabs", "tts", - false, + true, &config_with_deepgram_key(), &tx, ) @@ -3671,3 +3781,97 @@ mod tests { assert_eq!(tts.audio_out_chunk_ms, Some(20), "15ms → 20ms opus frame"); } } + +#[cfg(test)] +mod frd023_bud_mode_tests { + //! TC-SEC-01 / TC-SEC-03: a Bud-mode `/ws` config never chooses a platform-side host or key. + use super::*; + use crate::handlers::ws::state::ConnectionState; + + fn conversation(extra: serde_json::Value) -> ConversationWebSocketConfig { + let mut base = serde_json::json!({"base_url": "", "model": "chat-deployment"}); + for (k, v) in extra.as_object().unwrap() { + base[k] = v.clone(); + } + serde_json::from_value(base).expect("conversation config") + } + + #[test] + fn tc_sec_01_client_llm_endpoint_and_keys_are_refused() { + for (field, value) in [ + ("base_url", serde_json::json!("https://attacker.example/v1")), + ("api_key", serde_json::json!("sk-caller")), + ( + "reasoning_base_url", + serde_json::json!("https://attacker.example/v1"), + ), + ("reasoning_api_key", serde_json::json!("sk-caller")), + ] { + let conv = conversation(serde_json::json!({ field: value })); + let refusal = bud_mode_config_refusal(Some(&conv), None) + .unwrap_or_else(|| panic!("{field} was accepted in Bud mode")); + assert!(refusal.contains(field), "must name {field}: {refusal}"); + assert!(refusal.contains("FRD-023 RT6"), "{refusal}"); + assert!(!refusal.contains("sk-caller"), "a key must not be echoed"); + } + } + + #[test] + fn a_conversation_naming_only_a_deployment_is_not_refused() { + assert!( + bud_mode_config_refusal(Some(&conversation(serde_json::json!({}))), None).is_none() + ); + } + + #[test] + fn tc_sec_03_inline_dag_definitions_are_refused() { + let dag: DAGWebSocketConfig = + serde_json::from_value(serde_json::json!({"definition": {"nodes": []}})).unwrap(); + let refusal = bud_mode_config_refusal(None, Some(&dag)).expect("refused"); + assert!(refusal.contains("dag_config.definition"), "{refusal}"); + + let template: DAGWebSocketConfig = + serde_json::from_value(serde_json::json!({"template": "support-agent"})).unwrap(); + assert!(bud_mode_config_refusal(None, Some(&template)).is_none()); + } + + #[tokio::test] + #[serial_test::serial] + async fn tc_sec_01_handle_config_refuses_before_building_anything() { + let app_state = crate::test_support::bud_state(&[]).await; + let state = Arc::new(RwLock::new(ConnectionState::new())); + let (tx, mut rx) = mpsc::channel(8); + let conv = conversation(serde_json::json!({"base_url": "https://attacker.example/v1"})); + + let keep_open = handle_config_message( + None, + Some(false), + None, + None, + None, + None, + Some(conv), + None, + &state, + &tx, + &app_state, + ) + .await; + + assert!( + keep_open, + "a refused config is an error frame, not a dropped socket" + ); + match rx.try_recv() { + Ok(MessageRoute::Outgoing(OutgoingMessage::Error { message })) => { + assert!(message.contains("FRD-023 RT6"), "{message}") + } + other => panic!("expected an error frame, got {other:?}"), + } + let guard = state.read().await; + assert!( + guard.stream_id.is_none() && guard.voice_manager.is_none(), + "nothing was built" + ); + } +} diff --git a/gateway/src/lib.rs b/gateway/src/lib.rs index ee5ba53c..b0d47766 100644 --- a/gateway/src/lib.rs +++ b/gateway/src/lib.rs @@ -28,6 +28,9 @@ pub mod routes; pub mod state; pub mod utils; +#[cfg(test)] +pub(crate) mod test_support; + // Re-export commonly used items for convenience pub use config::ServerConfig; pub use core::*; diff --git a/gateway/src/test_support.rs b/gateway/src/test_support.rs new file mode 100644 index 00000000..ed14c594 --- /dev/null +++ b/gateway/src/test_support.rs @@ -0,0 +1,94 @@ +//! Shared helpers for the crate's own unit tests. + +use std::sync::Arc; + +use crate::config::ServerConfig; +use crate::state::AppState; + +/// A credential-free config for a real `AppState` (the same literal `state` tests use). +pub(crate) fn minimal_config() -> ServerConfig { + ServerConfig { + host: "localhost".to_string(), + port: 3001, + tls: None, + livekit_url: "ws://localhost:7880".to_string(), + livekit_public_url: "http://localhost:7880".to_string(), + livekit_api_key: None, + livekit_api_secret: None, + deepgram_api_key: None, + elevenlabs_api_key: None, + google_credentials: None, + azure_speech_subscription_key: None, + azure_speech_region: None, + cartesia_api_key: None, + openai_api_key: None, + azure_openai_api_key: None, + azure_openai_endpoint: None, + grok_api_key: None, + inworld_api_key: None, + gemini_api_key: None, + ultravox_api_key: None, + speechmatics_api_key: None, + yandex_api_key: None, + yandex_folder_id: None, + assemblyai_api_key: None, + hume_api_key: None, + groq_api_key: None, + ibm_watson_api_key: None, + ibm_watson_instance_id: None, + ibm_watson_region: None, + aws_access_key_id: None, + aws_secret_access_key: None, + aws_region: None, + gnani_token: None, + gnani_access_key: None, + gnani_certificate_path: None, + recording_s3_bucket: None, + recording_s3_region: None, + recording_s3_endpoint: None, + recording_s3_access_key: None, + recording_s3_secret_key: None, + recording_s3_prefix: None, + cache_path: None, + cache_ttl_seconds: Some(3600), + auth_service_url: None, + auth_signing_key_path: None, + auth_api_secrets: Vec::new(), + auth_timeout_seconds: 5, + auth_required: false, + sip: None, + cors_allowed_origins: None, + rate_limit_requests_per_second: 60, + rate_limit_burst_size: 10, + max_websocket_connections: None, + max_connections_per_ip: 100, + ws_processing_timeout_secs: 10, + realtime_processing_timeout_secs: 30, + sip_max_participants: 3, + realtime_endpoint_overrides: Default::default(), + aliases: Default::default(), + plugins: crate::config::PluginConfig::default(), + dag_timeouts: crate::config::DAGTimeoutsConfig::default(), + } +} + +/// An `AppState` in Bud mode over an in-memory control plane holding `keys` (FRD-023 tests). +/// +/// Uses `BudMode::for_plane`, which never connects to Redis and does NOT mark the process as in +/// Bud mode, so it cannot leak into tests running beside it. +pub(crate) async fn bud_state(keys: &[(&str, &str)]) -> Arc { + let store = Arc::new(bud_auth::MemoryStore::new()); + for (k, v) in keys { + store.set(k, v); + } + let plane = Arc::new(bud_auth::BudPlane::new( + store as Arc, + None, + )); + plane.boot().await.expect("plane boots"); + let mut state = AppState::new(minimal_config()).await; + Arc::get_mut(&mut state) + .expect("the state is not shared yet") + .bud_mode = Some(crate::auth::bud_mode::BudMode::for_plane(plane).expect("bud mode")); + state +} From dda42974b341b3f3bc7352486749a08348d7db59 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 03:13:20 +0530 Subject: [PATCH 03/17] fix(gateway): connection-slot acquire deadlocked with DEBUG logging on try_acquire_connection held its DashMap entry (the shard's WRITE lock) while the "Connection acquired" debug! evaluated ip_connection_count(), which takes the same shard's READ lock. The field is evaluated only when DEBUG is enabled, so any deployment run with debug logging froze the worker thread of the first WebSocket upgrade forever (FRD-022 connection slots). Found by FRD-023's relay tests, whose span capture enables DEBUG. The entry is dropped before logging, and the count comes from the fetch_add result. Regression test acquires and releases a slot under a DEBUG subscriber with a 10 s deadline. Co-Authored-By: Claude Opus 5.5 --- gateway/src/state/mod.rs | 65 +++++++++++++++++++++++++++++++++++++++- 1 file changed, 64 insertions(+), 1 deletion(-) diff --git a/gateway/src/state/mod.rs b/gateway/src/state/mod.rs index 367bae0d..69fbfe45 100644 --- a/gateway/src/state/mod.rs +++ b/gateway/src/state/mod.rs @@ -534,6 +534,7 @@ impl AppState { if max_per_ip != 0 && current_ip >= max_per_ip as usize { // Rollback both counters ip_entry.fetch_sub(1, Ordering::Relaxed); + drop(ip_entry); self.active_ws_connections.fetch_sub(1, Ordering::Relaxed); tracing::warn!( ip = %ip, @@ -544,10 +545,15 @@ impl AppState { return Err(ConnectionLimitError::PerIpLimitReached); } + // The entry holds its DashMap shard's WRITE lock. It must be dropped before anything reads + // the map again: `ip_connection_count` takes the same shard's read lock, and the `debug!` + // below evaluates it only when DEBUG is enabled — so with debug logging on, the first + // WebSocket connection deadlocked its worker thread (found by FRD-023's relay tests). + drop(ip_entry); tracing::debug!( ip = %ip, total_connections = self.active_ws_connections.load(Ordering::Relaxed), - ip_connections = self.ip_connection_count(&ip), + ip_connections = current_ip + 1, "Connection acquired" ); self.export_active_sessions(); @@ -918,3 +924,60 @@ mod tests { cleanup_core_runtime_env(); } } + +#[cfg(test)] +mod connection_slot_debug_tests { + //! A connection slot must be acquirable with DEBUG logging enabled (see `try_acquire_connection`). + use super::*; + + struct DebugOn; + impl tracing::Subscriber for DebugOn { + fn register_callsite( + &self, + _m: &'static tracing::Metadata<'static>, + ) -> tracing::subscriber::Interest { + tracing::subscriber::Interest::always() + } + fn enabled(&self, _m: &tracing::Metadata<'_>) -> bool { + true + } + fn new_span(&self, _a: &tracing::span::Attributes<'_>) -> tracing::span::Id { + tracing::span::Id::from_u64(1) + } + fn record(&self, _s: &tracing::span::Id, _v: &tracing::span::Record<'_>) {} + fn record_follows_from(&self, _s: &tracing::span::Id, _f: &tracing::span::Id) {} + fn event(&self, e: &tracing::Event<'_>) { + // Evaluate every field, as a real formatter does. + struct V; + impl tracing::field::Visit for V { + fn record_debug(&mut self, _f: &tracing::field::Field, v: &dyn std::fmt::Debug) { + let _ = format!("{v:?}"); + } + } + e.record(&mut V); + } + fn enter(&self, _s: &tracing::span::Id) {} + fn exit(&self, _s: &tracing::span::Id) {} + } + + #[tokio::test] + #[serial_test::serial] + async fn a_slot_is_acquired_and_released_with_debug_logging_on() { + let state = AppState::new(crate::test_support::minimal_config()).await; + let ip: IpAddr = "10.1.2.3".parse().unwrap(); + let (tx, rx) = std::sync::mpsc::channel(); + let worker = Arc::clone(&state); + std::thread::spawn(move || { + tracing::subscriber::with_default(DebugOn, || { + let acquired = worker.try_acquire_connection(ip).is_ok(); + worker.release_connection(ip); + let _ = tx.send(acquired); + }); + }); + let acquired = rx + .recv_timeout(std::time::Duration::from_secs(10)) + .expect("try_acquire_connection deadlocked with DEBUG logging enabled"); + assert!(acquired); + assert_eq!(state.ip_connection_count(&ip), 0); + } +} From 2a2ae5a86f5e7b302852496ffb38f44550b16f11 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 03:13:36 +0530 Subject: [PATCH 04/17] feat(gateway): /v1/realtime serves Bud deployments over OpenAI Realtime GA (FRD-023 RT2, RT3, RT5) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit wss://gateway/v1/realtime?model= now speaks OpenAI Realtime GA, so the OpenAI SDKs, the Agents SDKs, LiveKit and Pipecat connect by changing only the base URL and the key. WaaV's native protocol stays on /realtime. - Handshake (RT2.1): credentials from Authorization: Bearer, the api-key header, or the openai-insecure-api-key. subprotocol (the server selects `realtime`, never echoes the credential); ?token=, OpenAI-Beta and call_id refused; model required; no permessage-deflate. Pre-upgrade failures are HTTP in OpenAI's error envelope. - Resolution and admission (RT2.2): the deployment resolves through the caller's allowlist as on REST; FRD-022 admission is taken once and its concurrency slot held for the session; an open breaker refuses before connecting. - Upstream (RT2.4): URL, auth header and model only from voice_table — OpenAI wss://api.openai.com/v1/realtime (or api_base, SSRF-validated), Azure GA wss:///openai/v1/realtime with api-key and no api-version; transcription deployments connect with intent=transcription. Connect deadline 10 s. - Relay and policy (RT2.5): two-stage parse, audio frames forwarded byte-for-byte; session.model and tracing stripped; MCP tools, stored prompts, image input, instruction overrides, session type changes and off-list transcription models refused per the deployment's policy with event_not_allowed (session continues); max_output_tokens clamped; vendor rate_limits.updated dropped; session.model rewritten to the deployment name. The deployment's defaults are sent first and client frames held until the vendor applies them (5 s). A client that stops reading is closed 1011 client_too_slow after 5 s — audio is never dropped silently. - Lifecycle (RT2.6): pings both legs every 20 s (3 missed → 1011), idle and maximum-length limits with a 60 s warning, revalidation every 30 s and before each response.create (key revoked, JWT subject removed from the project, endpoint unpublished → session_revoked + 1008), drain → server_shutdown + 1012. JWT expiry alone does not end a started session. - Metering (RT3.1): a voice.turn per response.done, per input transcription and per 60 s duration segment under a minute/second price — each the ROOT of its own trace with a link to the session's voice.session span (VoiceTurnFact coalesces on TraceId). Per-modality token cost with cached tokens subtracted from their class; a component without a rate is named in bud.voice.unpriced_components, never priced at zero. New attributes are declared in voice_turn_span!/voice_session_span! and match budmetrics' contract byte for byte. - Client secrets (RT5.1): POST /v1/realtime/client_secrets mints a stateless ek_bud_ secret sealed with XChaCha20-Poly1305 (claims opaque to the holder), lifetime capped by a JWT parent's exp, parent = the full snapshot hash; connect revalidates the parent. Keys from WAAV_CLIENT_SECRET_KEYS (first seals, all open); a bad key fails startup; none → 501. - Metrics (RT2.7): waav_realtime_sessions_active/_total, relay latency, policy refusals, unknown client events. Tests: lib units per module and tests/openai_realtime_relay.rs (39 cases against a mock GA vendor: TC-HS, TC-UP, TC-EVT, TC-LIFE, TC-MET, TC-EK, TC-SEC-07). Co-Authored-By: Claude Opus 5.5 --- bud-auth/src/runtime.rs | 39 +- bud-auth/tests/realtime_contract.rs | 36 +- gateway/Cargo.lock | 92 +- gateway/Cargo.toml | 2 + gateway/src/auth/ephemeral.rs | 400 ++++ gateway/src/auth/mod.rs | 1 + gateway/src/core/mod.rs | 3 +- gateway/src/core/realtime_cost.rs | 479 ++++ gateway/src/core/voice_cost.rs | 1 + gateway/src/handlers/mod.rs | 1 + .../openai_realtime/client_secrets.rs | 200 ++ .../src/handlers/openai_realtime/handshake.rs | 437 ++++ .../src/handlers/openai_realtime/metering.rs | 336 +++ gateway/src/handlers/openai_realtime/mod.rs | 29 + .../src/handlers/openai_realtime/policy.rs | 797 +++++++ .../src/handlers/openai_realtime/session.rs | 1015 +++++++++ .../src/handlers/openai_realtime/upstream.rs | 380 ++++ gateway/src/main.rs | 9 + gateway/src/observability/voice_attrs.rs | 324 ++- gateway/src/routes/mod.rs | 1 + gateway/src/routes/openai_realtime.rs | 32 + gateway/src/routes/realtime.rs | 31 +- gateway/src/state/mod.rs | 8 + gateway/tests/openai_realtime_relay.rs | 1990 +++++++++++++++++ gateway/tests/voice_span_contract.json | 214 +- 25 files changed, 6753 insertions(+), 104 deletions(-) create mode 100644 gateway/src/auth/ephemeral.rs create mode 100644 gateway/src/core/realtime_cost.rs create mode 100644 gateway/src/handlers/openai_realtime/client_secrets.rs create mode 100644 gateway/src/handlers/openai_realtime/handshake.rs create mode 100644 gateway/src/handlers/openai_realtime/metering.rs create mode 100644 gateway/src/handlers/openai_realtime/mod.rs create mode 100644 gateway/src/handlers/openai_realtime/policy.rs create mode 100644 gateway/src/handlers/openai_realtime/session.rs create mode 100644 gateway/src/handlers/openai_realtime/upstream.rs create mode 100644 gateway/src/routes/openai_realtime.rs create mode 100644 gateway/tests/openai_realtime_relay.rs diff --git a/bud-auth/src/runtime.rs b/bud-auth/src/runtime.rs index 2ed3ae39..527ab08f 100644 --- a/bud-auth/src/runtime.rs +++ b/bud-auth/src/runtime.rs @@ -822,7 +822,10 @@ mod realtime_reach_tests { r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1","model_id":"m1"},"__metadata__":{"api_key_id":"ak1","user_id":"u1","api_key_project_id":"p1"}}"#.to_string() } - async fn plane_with(keys: &[(&str, String)], jwt: Option>) -> (Arc, BudPlane) { + async fn plane_with( + keys: &[(&str, String)], + jwt: Option>, + ) -> (Arc, BudPlane) { let store = Arc::new(MemoryStore::new()); for (k, v) in keys { store.set(k, v); @@ -856,7 +859,9 @@ mod realtime_reach_tests { let key = format!("api_key:{hashed}"); let (store, plane) = plane_with(&[(&key, key_blob())], None).await; - let entry = plane.hash_reaches(&hashed, "ep-rt", false).expect("reaches"); + let entry = plane + .hash_reaches(&hashed, "ep-rt", false) + .expect("reaches"); assert_eq!(entry.project_id.as_deref(), Some("p1")); assert!(plane.hash_reaches(&hashed, "ep-other", false).is_none()); @@ -873,7 +878,10 @@ mod realtime_reach_tests { let hashed = hash_api_key("bud_client_x"); let (_s, plane) = plane_with( &[ - (&format!("api_key:{hashed}"), r#"{"__metadata__":{"api_key_id":"ak"}}"#.to_string()), + ( + &format!("api_key:{hashed}"), + r#"{"__metadata__":{"api_key_id":"ak"}}"#.to_string(), + ), ( crate::hydrate::PUBLISHED_MODEL_INFO_KEY, r#"{"pub-rt":{"endpoint_id":"ep-pub","project_id":"p9"}}"#.to_string(), @@ -891,14 +899,23 @@ mod realtime_reach_tests { let jwt = verifier(); let (store, plane) = plane_with( &[ - ("user_projects:sub-1", r#"{"user_id":"u1","projects":["p1"]}"#.to_string()), - ("project_models:p1", r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1"}}"#.to_string()), + ( + "user_projects:sub-1", + r#"{"user_id":"u1","projects":["p1"]}"#.to_string(), + ), + ( + "project_models:p1", + r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1"}}"#.to_string(), + ), ], Some(Arc::clone(&jwt)), ) .await; assert!(plane.subject_reaches("sub-1", "ep-rt").await.is_some()); - assert!(jwt.cached_authz("sub-1").is_some(), "the resolution is cached"); + assert!( + jwt.cached_authz("sub-1").is_some(), + "the resolution is cached" + ); // The user is removed from the project. store.set("user_projects:sub-1", r#"{"user_id":"u1","projects":[]}"#); @@ -918,8 +935,14 @@ mod realtime_reach_tests { let jwt = verifier(); let (store, plane) = plane_with( &[ - ("user_projects:sub-1", r#"{"user_id":"u1","projects":["p1"]}"#.to_string()), - ("project_models:p1", r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1"}}"#.to_string()), + ( + "user_projects:sub-1", + r#"{"user_id":"u1","projects":["p1"]}"#.to_string(), + ), + ( + "project_models:p1", + r#"{"rt":{"endpoint_id":"ep-rt","project_id":"p1"}}"#.to_string(), + ), ], Some(Arc::clone(&jwt)), ) diff --git a/bud-auth/tests/realtime_contract.rs b/bud-auth/tests/realtime_contract.rs index 03019cfc..3276ac17 100644 --- a/bud-auth/tests/realtime_contract.rs +++ b/bud-auth/tests/realtime_contract.rs @@ -32,8 +32,14 @@ fn the_realtime_block_round_trips() { let d = rt.defaults.as_ref().expect("defaults"); assert_eq!(d.voice.as_deref(), Some("marin")); - assert_eq!(d.instructions.as_deref(), Some("You are a helpful assistant.")); - assert_eq!(d.output_modalities.as_deref(), Some(&["audio".to_string()][..])); + assert_eq!( + d.instructions.as_deref(), + Some("You are a helpful assistant.") + ); + assert_eq!( + d.output_modalities.as_deref(), + Some(&["audio".to_string()][..]) + ); assert_eq!( d.turn_detection, Some(serde_json::json!({"type": "semantic_vad", "eagerness": "auto"})) @@ -101,8 +107,14 @@ fn a_malformed_field_inside_realtime_drops_only_that_field() { let ep = parse(&blob); let rt = ep.config.realtime.expect("block kept"); assert_eq!(rt.session_type.as_deref(), Some("realtime")); - assert!(rt.defaults.is_none(), "the malformed defaults block is dropped"); - assert!(rt.policy.unwrap().allows_mcp_tools(), "the good sibling survives"); + assert!( + rt.defaults.is_none(), + "the malformed defaults block is dropped" + ); + assert!( + rt.policy.unwrap().allows_mcp_tools(), + "the good sibling survives" + ); } #[test] @@ -129,7 +141,11 @@ fn unknown_rate_keys_are_ignored_not_fatal() { .to_string(); let pricing = parse(&blob).pricing.expect("usable"); assert_eq!(pricing.rates.get("input_video"), None); - assert_eq!(pricing.rates.get("output_text"), Some(&24.0), "numeric strings are numbers"); + assert_eq!( + pricing.rates.get("output_text"), + Some(&24.0), + "numeric strings are numbers" + ); assert_eq!(pricing.rates.len(), 3); } @@ -165,10 +181,16 @@ fn absent_policy_fields_take_the_secure_defaults() { .to_string(); let policy = parse(&blob).config.realtime.expect("block").policy(); assert!(!policy.allows_mcp_tools(), "MCP tools reach the vendor org"); - assert!(!policy.allows_prompt_references(), "stored prompts belong to the vendor org"); + assert!( + !policy.allows_prompt_references(), + "stored prompts belong to the vendor org" + ); assert!(policy.allows_client_instructions()); assert!(policy.allows_image_input()); - assert!(policy.allows_transcription_model("anything"), "no allowlist = any model"); + assert!( + policy.allows_transcription_model("anything"), + "no allowlist = any model" + ); let none = bud_auth::RealtimePolicy::default(); assert!(!none.allows_mcp_tools() && !none.allows_prompt_references()); diff --git a/gateway/Cargo.lock b/gateway/Cargo.lock index 7d943b92..2da4e825 100644 --- a/gateway/Cargo.lock +++ b/gateway/Cargo.lock @@ -56,6 +56,16 @@ version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" +[[package]] +name = "aead" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" +dependencies = [ + "crypto-common 0.1.6", + "generic-array", +] + [[package]] name = "aes" version = "0.8.4" @@ -148,7 +158,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -159,7 +169,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" dependencies = [ "anstyle", "once_cell_polyfill", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -1173,6 +1183,17 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +[[package]] +name = "chacha20" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures 0.2.17", +] + [[package]] name = "chacha20" version = "0.10.2" @@ -1184,6 +1205,19 @@ dependencies = [ "rand_core 0.10.1", ] +[[package]] +name = "chacha20poly1305" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35" +dependencies = [ + "aead", + "chacha20 0.9.1", + "cipher", + "poly1305", + "zeroize", +] + [[package]] name = "chrono" version = "0.4.44" @@ -1233,6 +1267,7 @@ checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" dependencies = [ "crypto-common 0.1.6", "inout", + "zeroize", ] [[package]] @@ -1325,7 +1360,7 @@ version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "faf9468729b8cbcea668e36183cb69d317348c2e08e994829fb56ebfdfbaac34" dependencies = [ - "windows-sys 0.48.0", + "windows-sys 0.61.2", ] [[package]] @@ -1579,6 +1614,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1bfb12502f3fc46cca1bb51ac28df9d618d813cdc3d2f25b9fe775a34af26bb3" dependencies = [ "generic-array", + "rand_core 0.6.4", "typenum", ] @@ -2131,7 +2167,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -3189,7 +3225,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" dependencies = [ "hermit-abi", "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -4060,7 +4096,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -4265,6 +4301,12 @@ version = "11.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e" +[[package]] +name = "opaque-debug" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" + [[package]] name = "openssl" version = "0.10.80" @@ -4836,6 +4878,17 @@ dependencies = [ "miniz_oxide", ] +[[package]] +name = "poly1305" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" +dependencies = [ + "cpufeatures 0.2.17", + "opaque-debug", + "universal-hash", +] + [[package]] name = "portable-atomic" version = "1.13.1" @@ -4949,7 +5002,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "22505a5c94da8e3b7c2996394d1c933236c4d743e81a410bcca4e6989fc066a4" dependencies = [ "bytes", - "heck 0.4.1", + "heck 0.5.0", "itertools 0.12.1", "log", "multimap", @@ -5175,7 +5228,7 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207" dependencies = [ - "chacha20", + "chacha20 0.10.2", "getrandom 0.4.2", "rand_core 0.10.1", ] @@ -5691,7 +5744,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -5786,7 +5839,7 @@ dependencies = [ "security-framework 3.7.0", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -5970,7 +6023,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5b55fb86dfd3a2f5f76ea78310a88f96c4ea21a3031f8d212443d56123fd0521" dependencies = [ "libc", - "windows-sys 0.52.0", + "windows-sys 0.61.2", ] [[package]] @@ -6291,7 +6344,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" dependencies = [ "libc", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -6618,7 +6671,7 @@ dependencies = [ "getrandom 0.4.2", "once_cell", "rustix", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] @@ -7541,6 +7594,16 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" +[[package]] +name = "universal-hash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" +dependencies = [ + "crypto-common 0.1.6", + "subtle", +] + [[package]] name = "unsafe-libyaml" version = "0.2.11" @@ -7714,6 +7777,7 @@ dependencies = [ "bud-auth", "bytemuck", "bytes", + "chacha20poly1305", "clap", "criterion", "dashmap", @@ -8091,7 +8155,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.48.0", + "windows-sys 0.61.2", ] [[package]] diff --git a/gateway/Cargo.toml b/gateway/Cargo.toml index 761e58c4..990aafea 100644 --- a/gateway/Cargo.toml +++ b/gateway/Cargo.toml @@ -176,6 +176,8 @@ jsonwebtoken = { version = "10.2.0", features = ["rust_crypto"] } tokio-tungstenite = { version = "0.28.0", features = ["rustls-tls-webpki-roots"] } url = "2.5.0" base64 = "0.22.0" +# FRD-023 §5.8: sealed ek_bud_ client secrets (XChaCha20-Poly1305). +chacha20poly1305 = "0.10" http = "1.0" # TLS support diff --git a/gateway/src/auth/ephemeral.rs b/gateway/src/auth/ephemeral.rs new file mode 100644 index 00000000..c4a277c8 --- /dev/null +++ b/gateway/src/auth/ephemeral.rs @@ -0,0 +1,400 @@ +//! `ek_bud_…` client secrets (FRD-023 §5.8, D-7, S-4). +//! +//! A backend holding a Bud credential mints a short-lived secret and hands it to a browser, which +//! connects to `/v1/realtime` with it. WaaV stores NOTHING: the secret carries its own claims, +//! sealed with XChaCha20-Poly1305 so they are opaque to the holder (often a third party's end +//! user) and tamper-evident in one step. +//! +//! ```text +//! ek_bud_. +//! ``` +//! +//! * **XChaCha20** — its 192-bit random nonce removes the nonce-reuse ceiling AES-GCM's 96-bit +//! random nonce puts on a long-lived key. +//! * **The parent is the full snapshot hash** (`hash_api_key`), so the parent check at connect is +//! the same `HashMap::get` a direct connect makes; a truncated fingerprint could not be looked +//! up. +//! * **Alphabet** — `kid` and unpadded base64url are HTTP `tchar`s, so the token is legal inside +//! `Sec-WebSocket-Protocol`. +//! +//! Keys: `WAAV_CLIENT_SECRET_KEYS` = comma-separated `kid:base64(32 bytes)`; the FIRST seals, all +//! open. A key that is not 32 bytes, or a duplicate `kid`, fails startup. + +use base64::Engine; +use chacha20poly1305::aead::{Aead, AeadCore, KeyInit, OsRng, Payload}; +use chacha20poly1305::{XChaCha20Poly1305, XNonce}; +use serde::{Deserialize, Serialize}; + +pub const TOKEN_PREFIX: &str = "ek_bud_"; +const AAD_PREFIX: &str = "ek_bud:v1:"; +const NONCE_LEN: usize = 24; +pub const KEYS_ENV: &str = "WAAV_CLIENT_SECRET_KEYS"; + +/// Who minted the secret. Revalidated at connect and every 30 s of a live session (D-17). +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "k")] +pub enum Parent { + /// An API key, by its auth-snapshot key (`sha256("bud-" + key)` hex). `ck`: it was a + /// `bud_client_*` key, which also reaches the published overlay. + #[serde(rename = "api_key")] + ApiKey { + h: String, + #[serde(default)] + ck: bool, + }, + /// A Keycloak subject. + #[serde(rename = "jwt")] + Jwt { sub: String }, +} + +/// The sealed claims. +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct Claims { + pub v: u8, + pub iat: u64, + /// Already capped at the parent's expiry. + pub exp: u64, + /// The bound endpoint id. + pub ep: String, + /// The name it was minted for. + pub alias: String, + pub parent: Parent, + /// Attribution copied from the minting principal: project, user, API key id. + #[serde(default)] + pub pid: Option, + #[serde(default)] + pub uid: Option, + #[serde(default)] + pub akid: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum OpenError { + /// Not an `ek_bud_` token at all, or not in its shape. + Malformed, + /// No configured key has this `kid`. + UnknownKey, + /// The box did not open: any altered byte, or a different key's `kid`. + Tampered, + Expired, +} + +struct SealKey { + kid: String, + cipher: XChaCha20Poly1305, +} + +/// The configured sealing keys. The first seals; all open. +pub struct ClientSecretKeys { + keys: Vec, +} + +impl std::fmt::Debug for ClientSecretKeys { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ClientSecretKeys") + .field( + "kids", + &self.keys.iter().map(|k| &k.kid).collect::>(), + ) + .finish() + } +} + +fn valid_kid(kid: &str) -> bool { + (1..=16).contains(&kid.len()) + && kid + .bytes() + .all(|b| b.is_ascii_alphanumeric() || b == b'_' || b == b'-') +} + +impl ClientSecretKeys { + /// Parse `kid:base64(32 bytes),…`. Every problem is fatal and named (TC-EK-16). + pub fn parse(spec: &str) -> Result { + let mut keys: Vec = Vec::new(); + for part in spec.split(',').map(str::trim).filter(|p| !p.is_empty()) { + let (kid, b64) = part.split_once(':').ok_or_else(|| { + format!("{KEYS_ENV}: entry is not kid:base64 (got a value without ':')") + })?; + let kid = kid.trim(); + if !valid_kid(kid) { + return Err(format!( + "{KEYS_ENV}: kid '{kid}' must be 1-16 characters of [A-Za-z0-9_-]" + )); + } + if keys.iter().any(|k| k.kid == kid) { + return Err(format!("{KEYS_ENV}: duplicate kid '{kid}'")); + } + let raw = base64::engine::general_purpose::STANDARD + .decode(b64.trim()) + .map_err(|_| format!("{KEYS_ENV}: key '{kid}' is not valid base64"))?; + let key: [u8; 32] = raw.as_slice().try_into().map_err(|_| { + format!( + "{KEYS_ENV}: key '{kid}' is {} bytes; XChaCha20-Poly1305 needs exactly 32", + raw.len() + ) + })?; + keys.push(SealKey { + kid: kid.to_string(), + cipher: XChaCha20Poly1305::new(&key.into()), + }); + } + if keys.is_empty() { + return Err(format!("{KEYS_ENV} is set but names no key")); + } + Ok(Self { keys }) + } + + /// `Ok(None)` when unset (client secrets disabled: the route answers 501). + pub fn from_env() -> Result, String> { + match std::env::var(KEYS_ENV) { + Ok(v) if !v.trim().is_empty() => Self::parse(&v).map(Some), + _ => Ok(None), + } + } + + pub fn sealing_kid(&self) -> &str { + &self.keys[0].kid + } + + fn aad(kid: &str) -> Vec { + format!("{AAD_PREFIX}{kid}").into_bytes() + } + + /// Seal claims with the first key. + pub fn seal(&self, claims: &Claims) -> Result { + let key = &self.keys[0]; + let plaintext = serde_json::to_vec(claims).map_err(|e| e.to_string())?; + let nonce = XChaCha20Poly1305::generate_nonce(&mut OsRng); + let aad = Self::aad(&key.kid); + let sealed = key + .cipher + .encrypt( + &nonce, + Payload { + msg: &plaintext, + aad: &aad, + }, + ) + .map_err(|_| "sealing failed".to_string())?; + let mut body = Vec::with_capacity(NONCE_LEN + sealed.len()); + body.extend_from_slice(&nonce); + body.extend_from_slice(&sealed); + Ok(format!( + "{TOKEN_PREFIX}{}.{}", + key.kid, + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(body) + )) + } + + /// Open a token, checking its expiry against `now` (unix seconds). + pub fn open(&self, token: &str, now: u64) -> Result { + let rest = token + .strip_prefix(TOKEN_PREFIX) + .ok_or(OpenError::Malformed)?; + let (kid, body) = rest.split_once('.').ok_or(OpenError::Malformed)?; + if !valid_kid(kid) { + return Err(OpenError::Malformed); + } + let key = self + .keys + .iter() + .find(|k| k.kid == kid) + .ok_or(OpenError::UnknownKey)?; + let body = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(body) + .map_err(|_| OpenError::Malformed)?; + if body.len() <= NONCE_LEN { + return Err(OpenError::Malformed); + } + let (nonce, sealed) = body.split_at(NONCE_LEN); + let nonce: [u8; NONCE_LEN] = nonce.try_into().map_err(|_| OpenError::Malformed)?; + let aad = Self::aad(kid); + let plaintext = key + .cipher + .decrypt( + &XNonce::from(nonce), + Payload { + msg: sealed, + aad: &aad, + }, + ) + .map_err(|_| OpenError::Tampered)?; + let claims: Claims = serde_json::from_slice(&plaintext).map_err(|_| OpenError::Tampered)?; + if claims.v != 1 { + return Err(OpenError::Malformed); + } + if claims.exp <= now { + return Err(OpenError::Expired); + } + Ok(claims) + } +} + +pub fn now_epoch() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_secs()) + .unwrap_or(0) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn b64(bytes: &[u8]) -> String { + base64::engine::general_purpose::STANDARD.encode(bytes) + } + + fn keys(spec: &str) -> ClientSecretKeys { + ClientSecretKeys::parse(spec).unwrap() + } + + fn claims(exp: u64) -> Claims { + Claims { + v: 1, + iat: 1000, + exp, + ep: "0f5d9a3e-1111-4222-8333-444455556666".into(), + alias: "my-rt".into(), + parent: Parent::ApiKey { + h: bud_auth::hash_api_key("bud_parent_key"), + ck: false, + }, + pid: Some("proj-7c1e".into()), + uid: Some("user-9a2b".into()), + akid: Some("akid-44f0".into()), + } + } + + #[test] + fn a_sealed_secret_opens_to_its_claims() { + let k = keys(&format!("k1:{}", b64(&[7u8; 32]))); + let token = k.seal(&claims(5000)).unwrap(); + assert_eq!(k.open(&token, 2000).unwrap(), claims(5000)); + } + + /// TC-EK-06 🔒 — any altered byte, or another configured key's kid, fails. + #[test] + fn tc_ek_06_tampering_fails() { + let k = keys(&format!("k1:{},k2:{}", b64(&[7u8; 32]), b64(&[9u8; 32]))); + let token = k.seal(&claims(5000)).unwrap(); + let (head, body) = token.rsplit_once('.').unwrap(); + let mut bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(body) + .unwrap(); + let last = bytes.len() - 1; + bytes[last] ^= 0x01; + let flipped = format!( + "{head}.{}", + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(&bytes) + ); + assert_eq!(k.open(&flipped, 2000), Err(OpenError::Tampered)); + let other_kid = token.replacen("ek_bud_k1.", "ek_bud_k2.", 1); + assert_eq!(k.open(&other_kid, 2000), Err(OpenError::Tampered)); + } + + /// TC-EK-07 + #[test] + fn tc_ek_07_expired_secrets_are_refused() { + let k = keys(&format!("k1:{}", b64(&[7u8; 32]))); + let token = k.seal(&claims(5000)).unwrap(); + assert_eq!(k.open(&token, 5000), Err(OpenError::Expired)); + } + + /// TC-EK-10 — rotation: `k1` → `k2,k1` → `k2`. + #[test] + fn tc_ek_10_rotation() { + let (k1, k2) = (b64(&[1u8; 32]), b64(&[2u8; 32])); + let old = keys(&format!("k1:{k1}")).seal(&claims(5000)).unwrap(); + let both = keys(&format!("k2:{k2},k1:{k1}")); + assert!( + both.open(&old, 2000).is_ok(), + "an old secret opens during rotation" + ); + let new = both.seal(&claims(5000)).unwrap(); + assert!(new.starts_with("ek_bud_k2."), "the first key seals"); + let only_new = keys(&format!("k2:{k2}")); + assert!(only_new.open(&new, 2000).is_ok()); + assert_eq!(only_new.open(&old, 2000), Err(OpenError::UnknownKey)); + } + + /// TC-EK-13 🔒 — the claims are opaque to the holder. + #[test] + fn tc_ek_13_claims_are_opaque() { + let k = keys(&format!("k1:{}", b64(&[7u8; 32]))); + let c = claims(5000); + let token = k.seal(&c).unwrap(); + let body = token.rsplit_once('.').unwrap().1; + let decoded = base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(body) + .unwrap(); + let haystack = String::from_utf8_lossy(&decoded); + let Parent::ApiKey { h, .. } = &c.parent else { + unreachable!() + }; + for needle in [ + c.ep.as_str(), + c.alias.as_str(), + h.as_str(), + c.pid.as_deref().unwrap(), + c.uid.as_deref().unwrap(), + c.akid.as_deref().unwrap(), + ] { + assert!( + !haystack.contains(needle), + "{needle} is readable in the token" + ); + assert!(!token.contains(needle), "{needle} is readable in the token"); + } + } + + /// TC-EK-15 — legal inside `Sec-WebSocket-Protocol`. + #[test] + fn tc_ek_15_subprotocol_safe_alphabet() { + let k = keys(&format!("k-1_a:{}", b64(&[7u8; 32]))); + let token = k.seal(&claims(5000)).unwrap(); + let (head, body) = token.split_once('.').unwrap(); + assert!(head.starts_with("ek_bud_")); + let kid = &head["ek_bud_".len()..]; + assert!(valid_kid(kid)); + assert!( + body.bytes() + .all(|b| b.is_ascii_alphanumeric() || b == b'-' || b == b'_'), + "{body}" + ); + assert!(!token.contains('=')); + // ~410 characters with full UUID attribution (FRD-023 §5.8): comfortably inside a header. + assert!((200..=600).contains(&token.len()), "length {}", token.len()); + } + + /// TC-EK-16 — bad key configuration fails, naming the problem. + #[test] + fn tc_ek_16_bad_keys_fail_startup() { + let short = ClientSecretKeys::parse(&format!("k1:{}", b64(&[7u8; 16]))).unwrap_err(); + assert!(short.contains("16 bytes"), "{short}"); + let dup = ClientSecretKeys::parse(&format!("k1:{0},k1:{0}", b64(&[7u8; 32]))).unwrap_err(); + assert!(dup.contains("duplicate kid"), "{dup}"); + assert!(ClientSecretKeys::parse("k1").is_err()); + assert!(ClientSecretKeys::parse(&format!("bad.kid:{}", b64(&[7u8; 32]))).is_err()); + assert!(ClientSecretKeys::parse(" , ").is_err()); + } + + #[test] + fn junk_is_malformed() { + let k = keys(&format!("k1:{}", b64(&[7u8; 32]))); + for junk in [ + "", + "bud_key", + "ek_bud_", + "ek_bud_k1", + "ek_bud_k1.!!!", + "ek_bud_k1.AAAA", + ] { + assert!( + matches!(k.open(junk, 0), Err(OpenError::Malformed)), + "{junk}" + ); + } + } +} diff --git a/gateway/src/auth/mod.rs b/gateway/src/auth/mod.rs index ace346e0..91b0282d 100644 --- a/gateway/src/auth/mod.rs +++ b/gateway/src/auth/mod.rs @@ -2,6 +2,7 @@ pub mod api_secret; pub mod bud_mode; pub mod client; pub mod context; +pub mod ephemeral; pub mod jwt; // Re-export commonly used items diff --git a/gateway/src/core/mod.rs b/gateway/src/core/mod.rs index 4ee3b57e..af3bc743 100644 --- a/gateway/src/core/mod.rs +++ b/gateway/src/core/mod.rs @@ -36,6 +36,8 @@ pub mod pipeline; pub mod providers; pub mod readiness; pub mod realtime; +/// A voice call's cost by its deployment's published price (FRD-021 §6.4). +pub mod realtime_cost; pub mod resilience; pub mod silero_vad; pub mod smart_turn; @@ -53,7 +55,6 @@ pub mod vendor_error; /// catalog. Mirrors the [`emotion`] / [`lang`] mapper chassis; raw `voice_id` is the /// escape hatch, no-match → provider default + `config_warning` (never a 400). pub mod voice; -/// A voice call's cost by its deployment's published price (FRD-021 §6.4). pub mod voice_cost; /// Why a voice call failed, as the closed class vocabulary analytics count (FRD-021 §6.5). pub mod voice_error; diff --git a/gateway/src/core/realtime_cost.rs b/gateway/src/core/realtime_cost.rs new file mode 100644 index 00000000..33076bf5 --- /dev/null +++ b/gateway/src/core/realtime_cost.rs @@ -0,0 +1,479 @@ +//! What a realtime (speech-to-speech) record costs (FRD-023 §5.10, D-9, D-10). +//! +//! Realtime is the one place on the audio plane that bills in TOKENS, and at up to eight rates at +//! once: text, audio and image in, their cached forms, text and audio out. A single +//! `cost_per_unit` cannot express that — cached audio input is ~80× cheaper than uncached — so a +//! realtime price is a map of per-modality rates (`VoicePricing::rates`, CONTRACTS C1). +//! +//! Two rules are load-bearing: +//! +//! * **Cached tokens are a SUBSET of input tokens** (OpenAI's usage contract). Each modality's +//! uncached count is `input − cached`, priced at its own rate; pricing both in full would bill +//! the cached part twice. +//! * **A component with no rate is UNPRICED, never free.** It is left out of the cost and named in +//! `unpriced`, which the turn span records (`bud.voice.unpriced_components`). Pricing it at zero +//! would make a missing rate indistinguishable from a free one. +//! +//! Pure functions: no I/O, no spans. + +use bud_auth::VoicePricing; + +/// Token counts from one `response.done` (or their sum over a session). +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct RealtimeUsage { + pub input_text: u64, + pub input_audio: u64, + pub input_image: u64, + pub cached_text: u64, + pub cached_audio: u64, + pub cached_image: u64, + pub output_text: u64, + pub output_audio: u64, +} + +impl RealtimeUsage { + /// Read OpenAI GA `response.usage`: + /// `{input_tokens, output_tokens, input_token_details: {text_tokens, audio_tokens, + /// image_tokens, cached_tokens, cached_tokens_details: {text_tokens, audio_tokens, + /// image_tokens}}, output_token_details: {text_tokens, audio_tokens}}`. + /// + /// `None` when there is no usage object at all (a vendor that reports none). Missing counts + /// inside it read as zero, which is what they mean. + pub fn from_openai(usage: &serde_json::Value) -> Option { + let usage = usage.as_object()?; + let n = |v: Option<&serde_json::Value>| v.and_then(serde_json::Value::as_u64).unwrap_or(0); + let input = usage.get("input_token_details"); + let cached = input.and_then(|d| d.get("cached_tokens_details")); + let output = usage.get("output_token_details"); + let mut u = Self { + input_text: n(input.and_then(|d| d.get("text_tokens"))), + input_audio: n(input.and_then(|d| d.get("audio_tokens"))), + input_image: n(input.and_then(|d| d.get("image_tokens"))), + cached_text: n(cached.and_then(|d| d.get("text_tokens"))), + cached_audio: n(cached.and_then(|d| d.get("audio_tokens"))), + cached_image: n(cached.and_then(|d| d.get("image_tokens"))), + output_text: n(output.and_then(|d| d.get("text_tokens"))), + output_audio: n(output.and_then(|d| d.get("audio_tokens"))), + }; + // A vendor that reports only totals: attribute them to text rather than drop them, so + // the record still carries volume. An all-modality breakdown always wins when present. + if input.is_none() && output.is_none() { + u.input_text = n(usage.get("input_tokens")); + u.output_text = n(usage.get("output_tokens")); + } + // A cached count larger than its input is a vendor inconsistency; never let the + // uncached remainder go negative (it would REDUCE the bill). + u.cached_text = u.cached_text.min(u.input_text); + u.cached_audio = u.cached_audio.min(u.input_audio); + u.cached_image = u.cached_image.min(u.input_image); + Some(u) + } + + pub fn add(&mut self, other: &Self) { + self.input_text += other.input_text; + self.input_audio += other.input_audio; + self.input_image += other.input_image; + self.cached_text += other.cached_text; + self.cached_audio += other.cached_audio; + self.cached_image += other.cached_image; + self.output_text += other.output_text; + self.output_audio += other.output_audio; + } + + pub fn is_empty(&self) -> bool { + *self == Self::default() + } +} + +/// What an input transcription reported (`conversation.item.input_audio_transcription.completed` +/// `usage`): either seconds of audio or token counts, depending on the transcription model. +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum TranscriptionUsage { + Seconds(f64), + Tokens { + input_audio: u64, + input_text: u64, + output_text: u64, + }, +} + +impl TranscriptionUsage { + pub fn from_openai(usage: &serde_json::Value) -> Option { + let n = |v: Option<&serde_json::Value>| v.and_then(serde_json::Value::as_u64).unwrap_or(0); + match usage.get("type").and_then(|t| t.as_str()) { + Some("duration") => usage + .get("seconds") + .and_then(serde_json::Value::as_f64) + .filter(|s| s.is_finite() && *s >= 0.0) + .map(Self::Seconds), + Some("tokens") => { + let details = usage.get("input_token_details"); + let (audio, text) = match details { + Some(d) => (n(d.get("audio_tokens")), n(d.get("text_tokens"))), + // No breakdown: a transcription's input is audio. + None => (n(usage.get("input_tokens")), 0), + }; + Some(Self::Tokens { + input_audio: audio, + input_text: text, + output_text: n(usage.get("output_tokens")), + }) + } + _ => None, + } + } +} + +/// A computed price and what it could not price. +#[derive(Debug, Clone, PartialEq)] +pub struct RealtimeCost { + /// `None` when nothing was priceable (no price, or every present component unpriced). + pub cost: Option, + /// What `cost` was computed from: `token` | `minute` | `second`. `None` with `cost`. + pub unit: Option<&'static str>, + /// Components that were present and had no rate — named, never zero-priced. + pub unpriced: Vec<&'static str>, +} + +impl RealtimeCost { + fn none() -> Self { + Self { + cost: None, + unit: None, + unpriced: Vec::new(), + } + } +} + +fn rate(pricing: &VoicePricing, key: &str) -> Option { + pricing + .rates + .get(key) + .copied() + .filter(|r| r.is_finite() && *r >= 0.0) +} + +/// Accumulates `count × rate / per_units` over components, naming each unpriced one. +struct Tally<'p> { + pricing: &'p VoicePricing, + per_units: f64, + total: f64, + priced_any: bool, + unpriced: Vec<&'static str>, +} + +impl<'p> Tally<'p> { + fn new(pricing: &'p VoicePricing) -> Self { + Self { + pricing, + per_units: pricing.per_units as f64, + total: 0.0, + priced_any: false, + unpriced: Vec::new(), + } + } + + fn add(&mut self, count: u64, key: &'static str) { + if count == 0 { + return; + } + match rate(self.pricing, key) { + Some(r) => { + self.total += count as f64 * r / self.per_units; + self.priced_any = true; + } + None => self.unpriced.push(key), + } + } + + fn finish(self, nothing_present: bool) -> RealtimeCost { + let cost = if self.priced_any || nothing_present { + Some(self.total).filter(|c| c.is_finite()) + } else { + None + }; + RealtimeCost { + unit: cost.map(|_| "token"), + cost, + unpriced: self.unpriced, + } + } +} + +/// One `response.done`, under a token price (FRD-023 §5.10): +/// +/// ```text +/// cost = [ (in_text − cached_text)·r.input_text + cached_text·r.cached_input_text +/// + (in_audio − cached_audio)·r.input_audio + cached_audio·r.cached_input_audio +/// + (in_image − cached_image)·r.input_image + cached_image·r.cached_input_image +/// + out_text·r.output_text + out_audio·r.output_audio ] / per_units +/// ``` +/// +/// A minute or second price bills the session's DURATION instead ([`realtime_duration_cost`]), so +/// a response under one carries its tokens and no cost. +pub fn realtime_response_cost( + pricing: Option<&VoicePricing>, + usage: &RealtimeUsage, +) -> RealtimeCost { + let Some(pricing) = pricing else { + return RealtimeCost::none(); + }; + if pricing.unit != "token" || pricing.per_units == 0 { + return RealtimeCost::none(); + } + let mut t = Tally::new(pricing); + t.add(usage.input_text - usage.cached_text, "input_text"); + t.add(usage.cached_text, "cached_input_text"); + t.add(usage.input_audio - usage.cached_audio, "input_audio"); + t.add(usage.cached_audio, "cached_input_audio"); + t.add(usage.input_image - usage.cached_image, "input_image"); + t.add(usage.cached_image, "cached_input_image"); + t.add(usage.output_text, "output_text"); + t.add(usage.output_audio, "output_audio"); + t.finish(usage.is_empty()) +} + +/// One input transcription. Seconds are priced at `transcription_per_minute` (per MINUTE, not per +/// `per_units`); token usage at the three `transcription_*` token rates (per `per_units`). +/// Applies under every unit: a transcription is its own billed record. +pub fn realtime_transcription_cost( + pricing: Option<&VoicePricing>, + usage: &TranscriptionUsage, +) -> RealtimeCost { + let Some(pricing) = pricing else { + return RealtimeCost::none(); + }; + match *usage { + TranscriptionUsage::Seconds(secs) => match rate(pricing, "transcription_per_minute") { + Some(r) => { + let cost = secs / 60.0 * r; + RealtimeCost { + cost: cost.is_finite().then_some(cost), + unit: cost.is_finite().then_some("minute"), + unpriced: Vec::new(), + } + } + None if secs > 0.0 => RealtimeCost { + cost: None, + unit: None, + unpriced: vec!["transcription_per_minute"], + }, + None => RealtimeCost::none(), + }, + TranscriptionUsage::Tokens { + input_audio, + input_text, + output_text, + } => { + if pricing.per_units == 0 { + return RealtimeCost::none(); + } + let mut t = Tally::new(pricing); + t.add(input_audio, "transcription_input_audio"); + t.add(input_text, "transcription_input_text"); + t.add(output_text, "transcription_output_text"); + t.finish(input_audio == 0 && input_text == 0 && output_text == 0) + } + } +} + +/// A duration segment under a minute or second price (D-9: per 60 s, so a socket that drops at +/// minute 40 was still billed for 39 minutes). `None` under any other unit. +pub fn realtime_duration_cost(pricing: Option<&VoicePricing>, seconds: f64) -> RealtimeCost { + let Some(pricing) = pricing else { + return RealtimeCost::none(); + }; + if pricing.per_units == 0 || !seconds.is_finite() || seconds < 0.0 { + return RealtimeCost::none(); + } + let (units, unit) = match pricing.unit.as_str() { + "minute" => (seconds / 60.0, "minute"), + "second" => (seconds, "second"), + _ => return RealtimeCost::none(), + }; + let cost = units * pricing.cost_per_unit / pricing.per_units as f64; + RealtimeCost { + cost: cost.is_finite().then_some(cost), + unit: cost.is_finite().then_some(unit), + unpriced: Vec::new(), + } +} + +/// Whether this price bills the session's duration (so the relay emits 60 s segments). +pub fn bills_duration(pricing: Option<&VoicePricing>) -> bool { + pricing.is_some_and(|p| matches!(p.unit.as_str(), "minute" | "second")) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::BTreeMap; + + fn token_price(rates: &[(&str, f64)]) -> VoicePricing { + VoicePricing { + unit: "token".into(), + cost_per_unit: 0.0, + currency: Some("USD".into()), + per_units: 1_000_000, + rates: rates + .iter() + .map(|(k, v)| (k.to_string(), *v)) + .collect::>(), + } + } + + /// gpt-realtime-2.1 list prices (FRD §5.10). + fn gpt_realtime_21() -> VoicePricing { + token_price(&[ + ("input_text", 4.0), + ("input_audio", 32.0), + ("input_image", 5.0), + ("cached_input_text", 0.4), + ("cached_input_audio", 0.4), + ("cached_input_image", 0.5), + ("output_text", 24.0), + ("output_audio", 64.0), + ("transcription_per_minute", 0.003), + ]) + } + + /// OpenAI's documented `response.done` usage (the FRD's worked example). + fn documented_usage() -> serde_json::Value { + serde_json::json!({ + "total_tokens": 253, + "input_tokens": 132, + "output_tokens": 121, + "input_token_details": { + "text_tokens": 119, "audio_tokens": 13, "image_tokens": 0, + "cached_tokens": 64, + "cached_tokens_details": {"text_tokens": 64, "audio_tokens": 0, "image_tokens": 0} + }, + "output_token_details": {"text_tokens": 30, "audio_tokens": 91} + }) + } + + /// TC-MET-03 🔒 — the formula, to 1e-10. + #[test] + fn tc_met_03_the_worked_example_costs_0_0072056() { + let usage = RealtimeUsage::from_openai(&documented_usage()).unwrap(); + assert_eq!(usage.input_text, 119); + assert_eq!(usage.cached_text, 64); + let c = realtime_response_cost(Some(&gpt_realtime_21()), &usage); + let cost = c.cost.expect("priced"); + assert!((cost - 0.0072056).abs() < 1e-10, "cost {cost}"); + assert_eq!(c.unit, Some("token")); + assert!(c.unpriced.is_empty()); + } + + /// TC-MET-04 — cached audio at the cached rate, only the remainder at the full one. + #[test] + fn tc_met_04_cached_audio_is_subtracted_from_its_class() { + let usage = RealtimeUsage { + input_audio: 1000, + cached_audio: 800, + ..Default::default() + }; + let c = realtime_response_cost(Some(&gpt_realtime_21()), &usage); + let want = (200.0 * 32.0 + 800.0 * 0.4) / 1e6; + assert!((c.cost.unwrap() - want).abs() < 1e-12); + } + + /// TC-MET-05 — a present component without a rate is named, not zero-priced. + #[test] + fn tc_met_05_a_missing_rate_is_unpriced_not_free() { + let pricing = token_price(&[("input_audio", 32.0), ("output_audio", 64.0)]); + let usage = RealtimeUsage { + input_audio: 100, + input_image: 50, + output_audio: 10, + ..Default::default() + }; + let c = realtime_response_cost(Some(&pricing), &usage); + let want = (100.0 * 32.0 + 10.0 * 64.0) / 1e6; + assert!((c.cost.unwrap() - want).abs() < 1e-12, "images excluded"); + assert_eq!(c.unpriced, vec!["input_image"]); + } + + #[test] + fn every_present_component_unpriced_means_no_cost_at_all() { + let pricing = token_price(&[("output_audio", 64.0)]); + let usage = RealtimeUsage { + input_text: 10, + ..Default::default() + }; + let c = realtime_response_cost(Some(&pricing), &usage); + assert_eq!(c.cost, None, "an all-unpriced record must not read as free"); + assert_eq!(c.unpriced, vec!["input_text"]); + } + + /// TC-MET-06 — 12 s at $0.003/min. + #[test] + fn tc_met_06_duration_transcription_usage() { + let usage = TranscriptionUsage::from_openai( + &serde_json::json!({"type": "duration", "seconds": 12}), + ) + .unwrap(); + let c = realtime_transcription_cost(Some(&gpt_realtime_21()), &usage); + assert!((c.cost.unwrap() - 0.0006).abs() < 1e-12); + assert_eq!(c.unit, Some("minute")); + } + + #[test] + fn token_transcription_usage_uses_the_transcription_token_rates() { + let usage = TranscriptionUsage::from_openai(&serde_json::json!({ + "type": "tokens", "total_tokens": 30, "input_tokens": 20, "output_tokens": 10, + "input_token_details": {"text_tokens": 0, "audio_tokens": 20} + })) + .unwrap(); + let unpriced = realtime_transcription_cost(Some(&gpt_realtime_21()), &usage); + assert_eq!(unpriced.cost, None); + assert_eq!( + unpriced.unpriced, + vec!["transcription_input_audio", "transcription_output_text"] + ); + + let priced = token_price(&[ + ("transcription_input_audio", 3.0), + ("transcription_output_text", 5.0), + ]); + let c = realtime_transcription_cost(Some(&priced), &usage); + assert!((c.cost.unwrap() - (20.0 * 3.0 + 10.0 * 5.0) / 1e6).abs() < 1e-12); + } + + #[test] + fn a_minute_price_bills_duration_not_tokens() { + let pricing = VoicePricing { + unit: "minute".into(), + cost_per_unit: 0.06, + currency: None, + per_units: 1, + rates: BTreeMap::new(), + }; + assert!(bills_duration(Some(&pricing))); + let usage = RealtimeUsage::from_openai(&documented_usage()).unwrap(); + assert_eq!(realtime_response_cost(Some(&pricing), &usage).cost, None); + let seg = realtime_duration_cost(Some(&pricing), 60.0); + assert!((seg.cost.unwrap() - 0.06).abs() < 1e-12); + assert_eq!(seg.unit, Some("minute")); + let partial = realtime_duration_cost(Some(&pricing), 30.0); + assert!((partial.cost.unwrap() - 0.03).abs() < 1e-12); + } + + #[test] + fn a_cached_count_above_its_input_never_reduces_the_bill() { + let usage = RealtimeUsage::from_openai(&serde_json::json!({ + "input_token_details": {"text_tokens": 10, + "cached_tokens_details": {"text_tokens": 50}}, + "output_token_details": {} + })) + .unwrap(); + assert_eq!(usage.cached_text, 10); + } + + #[test] + fn no_price_prices_nothing() { + let usage = RealtimeUsage::from_openai(&documented_usage()).unwrap(); + assert_eq!(realtime_response_cost(None, &usage), RealtimeCost::none()); + assert!(!bills_duration(None)); + } +} diff --git a/gateway/src/core/voice_cost.rs b/gateway/src/core/voice_cost.rs index b072b0b3..4d51f63f 100644 --- a/gateway/src/core/voice_cost.rs +++ b/gateway/src/core/voice_cost.rs @@ -64,6 +64,7 @@ mod tests { cost_per_unit, currency: Some("USD".to_string()), per_units, + rates: Default::default(), } } diff --git a/gateway/src/handlers/mod.rs b/gateway/src/handlers/mod.rs index 57751fe0..23d3dbae 100644 --- a/gateway/src/handlers/mod.rs +++ b/gateway/src/handlers/mod.rs @@ -22,6 +22,7 @@ pub mod debug_profile; pub mod endpoint_settings; pub mod livekit; pub mod openai_audio; +pub mod openai_realtime; pub mod realtime; pub mod recording; pub mod sip; diff --git a/gateway/src/handlers/openai_realtime/client_secrets.rs b/gateway/src/handlers/openai_realtime/client_secrets.rs new file mode 100644 index 00000000..f4e10c7c --- /dev/null +++ b/gateway/src/handlers/openai_realtime/client_secrets.rs @@ -0,0 +1,200 @@ +//! `POST /v1/realtime/client_secrets` — mint an `ek_bud_…` (FRD-023 §5.8, FR-EK-1, D-7). +//! +//! OpenAI's request and response shapes, so an OpenAI client that mints through its SDK works +//! against the gateway unchanged. WaaV writes nothing anywhere: the secret is sealed +//! (`auth::ephemeral`) and carries its own claims. +//! +//! The lifetime is capped by the PARENT's: `exp = min(iat + seconds, parent expiry)`. A Keycloak +//! access token lives about five minutes, so without the cap a JWT could mint a two-hour bearer +//! credential. A clamp is not an error — `expires_at` reports the capped value. + +use std::sync::Arc; + +use axum::body::Bytes; +use axum::extract::State; +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; + +use crate::auth::ephemeral::{self, Claims, Parent}; +use crate::state::AppState; + +use super::handshake::{self, HandshakeError, REALTIME_CAPABILITY}; +use super::policy::{self, ClientOutcome, ClientRules}; +use super::session::{CallerCheck, authenticate}; + +pub const DEFAULT_TTL_SECS: u64 = 600; +pub const MIN_TTL_SECS: u64 = 10; +pub const MAX_TTL_SECS: u64 = 7200; +/// A JWT with less than this left cannot mint (TC-EK-12). +const MIN_PARENT_REMAINING_SECS: u64 = 10; + +fn bad_request( + code: &'static str, + message: impl Into, + param: &'static str, +) -> HandshakeError { + HandshakeError::new(StatusCode::BAD_REQUEST, code, message).param(param) +} + +/// Validate a mint request and seal the secret. Split from the handler for the tests. +pub async fn mint( + state: &AppState, + headers: &HeaderMap, + body: &[u8], + now: u64, +) -> Result { + let Some(keys) = state.realtime.client_secret_keys.as_ref() else { + return Err(HandshakeError::new( + StatusCode::NOT_IMPLEMENTED, + "client_secrets_not_configured", + "Client secrets are not enabled on this gateway.", + )); + }; + + let Some(credential) = handshake::extract_credential(headers)? else { + return Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "missing_api_key", + "Mint client secrets with a Bud API key or an access token in `Authorization: Bearer`.", + )); + }; + // No chaining: a client secret cannot mint another (TC-EK-04). + if credential.is_client_secret() { + return Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_api_key", + "A client secret cannot mint another client secret.", + )); + } + let caller = authenticate(state, &credential).await?; + + let request: serde_json::Value = if body.is_empty() { + serde_json::json!({}) + } else { + serde_json::from_slice(body).map_err(|e| { + bad_request( + "invalid_json", + format!("The body is not valid JSON: {e}"), + "body", + ) + })? + }; + + let expires_after = request.get("expires_after"); + if let Some(anchor) = expires_after.and_then(|e| e.get("anchor")) + && anchor.as_str() != Some("created_at") + { + return Err(bad_request( + "invalid_value", + "expires_after.anchor must be \"created_at\".", + "expires_after.anchor", + )); + } + let seconds = match expires_after.and_then(|e| e.get("seconds")) { + None => DEFAULT_TTL_SECS, + Some(v) => v + .as_u64() + .filter(|s| (MIN_TTL_SECS..=MAX_TTL_SECS).contains(s)) + .ok_or_else(|| { + bad_request( + "invalid_value", + format!("expires_after.seconds must be {MIN_TTL_SECS}..{MAX_TTL_SECS}."), + "expires_after.seconds", + ) + })?, + }; + + let session = request + .get("session") + .cloned() + .unwrap_or_else(|| serde_json::json!({})); + let model = session + .get("model") + .and_then(|m| m.as_str()) + .map(str::trim) + .filter(|m| !m.is_empty()) + .ok_or_else(|| { + bad_request( + "model_required", + "session.model is required.", + "session.model", + ) + })? + .to_string(); + + let resolved = state + .resolve_voice_endpoint(&model, REALTIME_CAPABILITY, Some(credential.expose())) + .ok_or_else(|| { + HandshakeError::new( + StatusCode::FORBIDDEN, + "model_not_allowed", + format!("Model '{model}' is not a realtime deployment this credential can reach."), + ) + .param("session.model") + })?; + + // Validate the rest of `session` against the deployment's policy (§5.5). Echoed, not bound + // (DEG-3). + let rules = ClientRules::from_settings(resolved.endpoint.config.realtime.as_ref()); + let probe = serde_json::json!({"type": "session.update", "session": session}).to_string(); + if let ClientOutcome::Refuse(r) = policy::client_event(&probe, &rules) { + return Err(HandshakeError::new( + StatusCode::BAD_REQUEST, + "event_not_allowed", + r.message, + )); + } + + // The parent's lifetime caps the secret's (TC-EK-12). + let mut exp = now + seconds; + if let Some(parent_exp) = caller.principal.expires_at { + if parent_exp < now + MIN_PARENT_REMAINING_SECS { + return Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "credential_expiring", + "The minting credential expires in under 10 seconds; refresh it first.", + )); + } + exp = exp.min(parent_exp); + } + + let parent = match &caller.check { + CallerCheck::ApiKey { hashed, client_key } => Parent::ApiKey { + h: hashed.clone(), + ck: *client_key, + }, + CallerCheck::Jwt { sub } => Parent::Jwt { sub: sub.clone() }, + }; + let claims = Claims { + v: 1, + iat: now, + exp, + ep: resolved.endpoint_id.clone(), + alias: model.clone(), + parent, + pid: caller.principal.project_id.clone(), + uid: caller.principal.user_id.clone(), + akid: caller.principal.api_key_id.clone(), + }; + let value = keys + .seal(&claims) + .map_err(|e| HandshakeError::new(StatusCode::INTERNAL_SERVER_ERROR, "internal_error", e))?; + metrics::counter!("waav_realtime_client_secrets_minted_total").increment(1); + Ok(serde_json::json!({ + "value": value, + "expires_at": exp, + "session": session, + })) +} + +/// The route handler. +pub async fn client_secrets_handler( + State(state): State>, + headers: HeaderMap, + body: Bytes, +) -> Response { + match mint(&state, &headers, &body, ephemeral::now_epoch()).await { + Ok(v) => (StatusCode::OK, axum::Json(v)).into_response(), + Err(e) => e.into_response(), + } +} diff --git a/gateway/src/handlers/openai_realtime/handshake.rs b/gateway/src/handlers/openai_realtime/handshake.rs new file mode 100644 index 00000000..4d998916 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/handshake.rs @@ -0,0 +1,437 @@ +//! The `/v1/realtime` handshake (FRD-023 §5.2, FR-RT-2…5, S-2, S-3). +//! +//! Everything here happens BEFORE the upgrade, so every refusal is an HTTP response in OpenAI's +//! error envelope — the shape an OpenAI SDK reads (`error.message`, `error.code`). +//! +//! The credential rules are the security-relevant part: +//! +//! * three sources — `Authorization: Bearer`, the `api-key` header (Azure-mode clients) and the +//! `openai-insecure-api-key.` subprotocol (browsers cannot set headers); +//! * `?token=` is REFUSED here (it stays valid on `/ws` and `/realtime` for existing WaaV +//! clients): query strings land in access logs; +//! * the credential never leaves this process: it is not echoed in the 101 (only `realtime` is +//! selected), and [`Credential`]'s `Debug` is redacted so it cannot reach a log line. + +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; + +/// The capability a `/v1/realtime` deployment serves (FRD D-4). +pub const REALTIME_CAPABILITY: &str = "realtime_session"; +/// The subprotocol the server selects in the 101. The Node `ws` package fails a handshake in +/// which the client offered protocols and the server selected none. +pub const SUBPROTOCOL: &str = "realtime"; +/// The credential-bearing subprotocol prefix OpenAI clients send from a browser. +pub const KEY_SUBPROTOCOL_PREFIX: &str = "openai-insecure-api-key."; +/// The beta shape, removed by OpenAI on 2026-05-12. +const BETA_SUBPROTOCOL: &str = "openai-beta.realtime-v1"; +/// Parse budget per message (S-8). +pub const MAX_MESSAGE_BYTES: usize = 10 * 1024 * 1024; + +/// Where the credential came from. Recorded (never the value) so a log can say how a caller +/// authenticated. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CredentialSource { + Bearer, + ApiKeyHeader, + Subprotocol, +} + +impl CredentialSource { + pub fn as_str(self) -> &'static str { + match self { + Self::Bearer => "authorization", + Self::ApiKeyHeader => "api-key", + Self::Subprotocol => "subprotocol", + } + } +} + +/// A caller's credential. `Debug` never prints the value (S-2, TC-SEC-07). +#[derive(Clone)] +pub struct Credential { + value: String, + pub source: CredentialSource, +} + +impl Credential { + pub fn new(value: impl Into, source: CredentialSource) -> Self { + Self { + value: value.into(), + source, + } + } + + /// The raw credential. Callers pass it to the auth plane and nowhere else. + pub fn expose(&self) -> &str { + &self.value + } + + /// An `ek_bud_` client secret (§5.8) rather than a Bud key or a JWT. + pub fn is_client_secret(&self) -> bool { + self.value.starts_with(crate::auth::ephemeral::TOKEN_PREFIX) + } +} + +impl std::fmt::Debug for Credential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Credential") + .field("value", &"[redacted]") + .field("source", &self.source) + .finish() + } +} + +/// A handshake that passed the pre-auth rules. +#[derive(Debug)] +pub struct Handshake { + /// The deployment the caller named (`?model=`): an alias or an endpoint UUID. + pub model: String, + pub credential: Credential, +} + +/// A refusal before the upgrade, rendered in OpenAI's error envelope. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct HandshakeError { + pub status: StatusCode, + pub kind: &'static str, + pub code: &'static str, + pub message: String, + pub param: Option<&'static str>, + pub retry_after: Option, +} + +impl HandshakeError { + pub fn new(status: StatusCode, code: &'static str, message: impl Into) -> Self { + let kind = match status { + StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => "invalid_request_error", + StatusCode::TOO_MANY_REQUESTS => "rate_limit_error", + s if s.is_server_error() => "api_error", + _ => "invalid_request_error", + }; + Self { + status, + kind, + code, + message: message.into(), + param: None, + retry_after: None, + } + } + + pub fn param(mut self, param: &'static str) -> Self { + self.param = Some(param); + self + } + + pub fn retry_after(mut self, secs: u64) -> Self { + self.retry_after = Some(secs.max(1)); + self + } +} + +impl IntoResponse for HandshakeError { + fn into_response(self) -> Response { + let body = serde_json::json!({ + "error": { + "message": self.message, + "type": self.kind, + "param": self.param, + "code": self.code, + } + }); + let mut resp = (self.status, axum::Json(body)).into_response(); + if let Some(secs) = self.retry_after + && let Ok(v) = axum::http::HeaderValue::from_str(&secs.to_string()) + { + resp.headers_mut() + .insert(axum::http::header::RETRY_AFTER, v); + } + resp + } +} + +/// Every value of every `Sec-WebSocket-Protocol` header, comma-split and trimmed. +pub fn offered_subprotocols(headers: &HeaderMap) -> Vec { + headers + .get_all(axum::http::header::SEC_WEBSOCKET_PROTOCOL) + .iter() + .filter_map(|v| v.to_str().ok()) + .flat_map(|v| v.split(',')) + .map(|p| p.trim().to_string()) + .filter(|p| !p.is_empty()) + .collect() +} + +/// The credential, from the first source present: `Authorization: Bearer` > `api-key` > +/// subprotocol. A malformed `Authorization` header is a refusal, not a fall-through: a caller who +/// sent one meant it. +pub fn extract_credential(headers: &HeaderMap) -> Result, HandshakeError> { + if let Some(value) = headers.get(axum::http::header::AUTHORIZATION) { + let text = value.to_str().map_err(|_| invalid_credential())?; + let token = text + .strip_prefix("Bearer ") + .or_else(|| text.strip_prefix("bearer ")) + .map(str::trim) + .filter(|t| !t.is_empty()) + .ok_or_else(invalid_credential)?; + return Ok(Some(Credential::new(token, CredentialSource::Bearer))); + } + if let Some(value) = headers.get("api-key") { + let token = value + .to_str() + .map(str::trim) + .ok() + .filter(|t| !t.is_empty()) + .ok_or_else(invalid_credential)?; + return Ok(Some(Credential::new(token, CredentialSource::ApiKeyHeader))); + } + for proto in offered_subprotocols(headers) { + if let Some(token) = proto.strip_prefix(KEY_SUBPROTOCOL_PREFIX) + && !token.is_empty() + { + return Ok(Some(Credential::new(token, CredentialSource::Subprotocol))); + } + } + Ok(None) +} + +fn invalid_credential() -> HandshakeError { + HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_api_key", + "The credential could not be read. Send `Authorization: Bearer `, an `api-key` \ + header, or the `openai-insecure-api-key.` subprotocol.", + ) +} + +/// Apply the §5.2 handshake rules, in this order: the beta shape, the query, the credential. +pub fn parse(query: Option<&str>, headers: &HeaderMap) -> Result { + let beta_subprotocol = offered_subprotocols(headers) + .iter() + .any(|p| p.eq_ignore_ascii_case(BETA_SUBPROTOCOL)); + if headers.contains_key("openai-beta") || beta_subprotocol { + return Err(HandshakeError::new( + StatusCode::BAD_REQUEST, + "beta_api_shape_disabled", + "The Realtime beta interface is not supported. Remove the `OpenAI-Beta` header \ + (or the `openai-beta.realtime-v1` subprotocol) and use the GA interface.", + )); + } + + let mut model: Option = None; + for (key, value) in url::form_urlencoded::parse(query.unwrap_or("").as_bytes()) { + match key.as_ref() { + "token" => { + return Err(HandshakeError::new( + StatusCode::BAD_REQUEST, + "use_subprotocol", + "Credentials are not accepted in the query string on /v1/realtime (query \ + strings are logged). Use `Authorization: Bearer`, the `api-key` header, or \ + the `openai-insecure-api-key.` subprotocol.", + ) + .param("token")); + } + "call_id" => { + return Err(HandshakeError::new( + StatusCode::BAD_REQUEST, + "unsupported_parameter", + "`call_id` (sideband control of a WebRTC or SIP call) is not supported by \ + this gateway.", + ) + .param("call_id")); + } + "model" => model = Some(value.trim().to_string()).filter(|m| !m.is_empty()), + // `intent=transcription` is accepted and ignored: the deployment decides the session + // type. Everything else is ignored and never forwarded upstream. + _ => {} + } + } + let Some(model) = model else { + return Err(HandshakeError::new( + StatusCode::BAD_REQUEST, + "model_required", + "`model` is required: connect to /v1/realtime?model=.", + ) + .param("model")); + }; + + let Some(credential) = extract_credential(headers)? else { + return Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "missing_api_key", + "You didn't provide a credential. Send `Authorization: Bearer `, an `api-key` \ + header, or the `openai-insecure-api-key.` subprotocol.", + )); + }; + + Ok(Handshake { model, credential }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn headers(pairs: &[(&str, &str)]) -> HeaderMap { + let mut h = HeaderMap::new(); + for (k, v) in pairs { + h.append( + axum::http::HeaderName::from_bytes(k.as_bytes()).unwrap(), + v.parse().unwrap(), + ); + } + h + } + + #[test] + fn tc_hs_01_bearer() { + let hs = parse( + Some("model=rt"), + &headers(&[("authorization", "Bearer bud_k")]), + ) + .unwrap(); + assert_eq!(hs.model, "rt"); + assert_eq!(hs.credential.expose(), "bud_k"); + assert_eq!(hs.credential.source, CredentialSource::Bearer); + } + + #[test] + fn tc_hs_02_api_key_header() { + let hs = parse(Some("model=rt"), &headers(&[("api-key", "bud_k")])).unwrap(); + assert_eq!(hs.credential.source, CredentialSource::ApiKeyHeader); + } + + #[test] + fn tc_hs_03_subprotocol_credential() { + let hs = parse( + Some("model=rt"), + &headers(&[( + "sec-websocket-protocol", + "realtime, openai-insecure-api-key.bud_k", + )]), + ) + .unwrap(); + assert_eq!(hs.credential.expose(), "bud_k"); + assert_eq!(hs.credential.source, CredentialSource::Subprotocol); + } + + #[test] + fn tc_hs_04_extra_subprotocols_are_tolerated() { + let hs = parse( + Some("model=rt"), + &headers(&[( + "sec-websocket-protocol", + "realtime, openai-agents-sdk.v0.18, openai-project.p, openai-insecure-api-key.k", + )]), + ) + .unwrap(); + assert_eq!(hs.credential.expose(), "k"); + } + + #[test] + fn tc_hs_06_model_is_required() { + let err = parse(None, &headers(&[("authorization", "Bearer k")])).unwrap_err(); + assert_eq!( + (err.status, err.code), + (StatusCode::BAD_REQUEST, "model_required") + ); + let err = parse(Some("model="), &headers(&[("authorization", "Bearer k")])).unwrap_err(); + assert_eq!(err.code, "model_required"); + } + + #[test] + fn tc_hs_07_query_token_is_refused() { + let err = parse(Some("token=bud_k&model=rt"), &HeaderMap::new()).unwrap_err(); + assert_eq!( + (err.status, err.code), + (StatusCode::BAD_REQUEST, "use_subprotocol") + ); + assert!(!err.message.contains("bud_k")); + } + + #[test] + fn tc_hs_08_the_beta_shape_is_refused() { + let err = parse( + Some("model=rt"), + &headers(&[ + ("authorization", "Bearer k"), + ("openai-beta", "realtime=v1"), + ]), + ) + .unwrap_err(); + assert_eq!(err.code, "beta_api_shape_disabled"); + let err = parse( + Some("model=rt"), + &headers(&[( + "sec-websocket-protocol", + "realtime, openai-beta.realtime-v1, openai-insecure-api-key.k", + )]), + ) + .unwrap_err(); + assert_eq!(err.code, "beta_api_shape_disabled"); + } + + #[test] + fn tc_hs_09_call_id_is_refused() { + let err = parse( + Some("model=rt&call_id=rtc_x"), + &headers(&[("authorization", "Bearer k")]), + ) + .unwrap_err(); + assert_eq!(err.code, "unsupported_parameter"); + } + + #[test] + fn tc_hs_10_intent_is_ignored() { + let hs = parse( + Some("model=rt&intent=transcription"), + &headers(&[("authorization", "Bearer k")]), + ) + .unwrap(); + assert_eq!(hs.model, "rt"); + } + + #[test] + fn no_credential_is_a_401() { + let err = parse(Some("model=rt"), &HeaderMap::new()).unwrap_err(); + assert_eq!( + (err.status, err.code), + (StatusCode::UNAUTHORIZED, "missing_api_key") + ); + } + + #[test] + fn a_malformed_authorization_header_is_refused_not_skipped() { + let err = parse( + Some("model=rt"), + &headers(&[("authorization", "Basic abc"), ("api-key", "k")]), + ) + .unwrap_err(); + assert_eq!(err.status, StatusCode::UNAUTHORIZED); + } + + /// TC-SEC-07 (unit half): the credential cannot reach a log through `Debug`. + #[test] + fn tc_sec_07_the_credential_debug_is_redacted() { + let hs = parse( + Some("model=rt"), + &headers(&[("authorization", "Bearer bud_secret_value")]), + ) + .unwrap(); + let printed = format!("{hs:?}"); + assert!(!printed.contains("bud_secret_value"), "{printed}"); + assert!(printed.contains("[redacted]")); + } + + #[test] + fn the_error_envelope_is_openai_shaped() { + let resp = HandshakeError::new( + StatusCode::TOO_MANY_REQUESTS, + "rate_limit_exceeded", + "slow down", + ) + .retry_after(3) + .into_response(); + assert_eq!(resp.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!(resp.headers()[axum::http::header::RETRY_AFTER], "3"); + } +} diff --git a/gateway/src/handlers/openai_realtime/metering.rs b/gateway/src/handlers/openai_realtime/metering.rs new file mode 100644 index 00000000..18059b7a --- /dev/null +++ b/gateway/src/handlers/openai_realtime/metering.rs @@ -0,0 +1,336 @@ +//! Metering a relayed session (FRD-023 §5.10, D-9, D-19, FR-MET-1…3). +//! +//! * **One `voice.turn` per billed record** — each `response.done`, each input-transcription +//! completion, and each 60 s duration segment under a minute/second price. Metering at the event +//! rather than at close is the point: a socket that drops at minute 40 has still been billed for +//! everything before it (and Kong shipped exactly this bug — realtime tokens that never reached +//! metrics). +//! * **Each turn is the ROOT of its own trace**, linked to the session's trace. `VoiceTurnFact` +//! coalesces rows on `(date, TraceId)`; turns sharing the session's trace would collapse into +//! one row and lose every response but one. +//! * **One `voice.session` span per session**, its own root, opened at start (so turns can link to +//! it) and ended at close with the totals. A pod crash loses it; the turns survive (DEG-4). + +use bud_auth::VoicePricing; +use opentelemetry::trace::TraceContextExt; +use serde_json::Value; +use tracing::Span; +use tracing_opentelemetry::OpenTelemetrySpanExt; + +use crate::core::realtime_cost::{ + RealtimeCost, RealtimeUsage, TranscriptionUsage, realtime_duration_cost, + realtime_response_cost, realtime_transcription_cost, +}; +use crate::observability::voice_attrs::{realtime as rt, session as sess, turn}; + +use super::handshake::REALTIME_CAPABILITY; + +/// Who a session is for, recorded on every span it emits (CONTRACTS C2). +#[derive(Debug, Clone, Default)] +pub struct Attribution { + pub project_id: Option, + pub endpoint_id: String, + pub model_id: Option, + pub api_key_id: Option, + pub user_id: Option, + pub api_key_project_id: Option, + /// The `?model=` the client connected with. + pub endpoint_name: String, + pub vendor: String, + /// The vendor model (`voice_table.model`). + pub model: Option, + pub session_type: String, +} + +fn record_text(span: &Span, key: &'static str, value: Option<&str>) { + // An empty string is not NULL; it would make the column look populated. + if let Some(v) = value.filter(|v| !v.is_empty()) { + span.record(key, v); + } +} + +fn record_attribution(span: &Span, a: &Attribution, session_id: &str) { + record_text(span, turn::PROJECT_ID, a.project_id.as_deref()); + record_text(span, turn::ENDPOINT_ID, Some(&a.endpoint_id)); + record_text(span, turn::MODEL_ID, a.model_id.as_deref()); + record_text(span, turn::API_KEY_ID, a.api_key_id.as_deref()); + record_text(span, turn::USER_ID, a.user_id.as_deref()); + record_text( + span, + turn::API_KEY_PROJECT_ID, + a.api_key_project_id.as_deref(), + ); + record_text(span, turn::ENDPOINT_NAME, Some(&a.endpoint_name)); + record_text(span, turn::SESSION_ID, Some(session_id)); + record_text(span, rt::VENDOR, Some(&a.vendor)); + record_text(span, rt::MODEL, a.model.as_deref()); +} + +fn record_usage(span: &Span, u: &RealtimeUsage) { + span.record(rt::INPUT_TEXT_TOKENS, u.input_text); + span.record(rt::INPUT_AUDIO_TOKENS, u.input_audio); + span.record(rt::INPUT_IMAGE_TOKENS, u.input_image); + span.record(rt::CACHED_TEXT_TOKENS, u.cached_text); + span.record(rt::CACHED_AUDIO_TOKENS, u.cached_audio); + span.record(rt::CACHED_IMAGE_TOKENS, u.cached_image); + span.record(rt::OUTPUT_TEXT_TOKENS, u.output_text); + span.record(rt::OUTPUT_AUDIO_TOKENS, u.output_audio); +} + +fn record_cost(span: &Span, c: &RealtimeCost) { + if let (Some(cost), Some(unit)) = (c.cost, c.unit) { + span.record(turn::COST, cost); + span.record(turn::PRICING_UNIT, unit); + } + if !c.unpriced.is_empty() { + span.record(rt::UNPRICED_COMPONENTS, c.unpriced.join(",").as_str()); + } +} + +/// The transcripts a `response.done` carries (`response.output[].content[].transcript|text`). +fn response_transcript(event: &Value) -> Option { + let mut parts = Vec::new(); + for item in event + .get("response") + .and_then(|r| r.get("output")) + .and_then(Value::as_array) + .into_iter() + .flatten() + { + for c in item + .get("content") + .and_then(Value::as_array) + .into_iter() + .flatten() + { + if let Some(t) = c + .get("transcript") + .or_else(|| c.get("text")) + .and_then(Value::as_str) + { + parts.push(t.to_string()); + } + } + } + (!parts.is_empty()).then(|| parts.join(" ")) +} + +/// The meter for one session. +pub struct SessionMeter { + session_span: Span, + session_id: String, + attribution: Attribution, + pricing: Option, + capture: bool, + turn_index: u64, + totals: RealtimeUsage, + billed_seconds: f64, + cost_total: f64, + priced_any: bool, + pricing_unit: Option<&'static str>, + vendor_session_id: Option, + started: std::time::Instant, +} + +impl SessionMeter { + /// Open the session's `voice.session` span (a root: its own trace). + pub fn start( + session_id: String, + attribution: Attribution, + pricing: Option, + ) -> Self { + let session_span = + crate::voice_session_span!(capability = REALTIME_CAPABILITY, transport = "websocket"); + record_attribution(&session_span, &attribution, &session_id); + record_text( + &session_span, + rt::SESSION_TYPE, + Some(&attribution.session_type), + ); + Self { + session_span, + session_id, + attribution, + pricing, + capture: crate::observability::trace_redact::capture_content(), + turn_index: 0, + totals: RealtimeUsage::default(), + billed_seconds: 0.0, + cost_total: 0.0, + priced_any: false, + pricing_unit: None, + vendor_session_id: None, + started: std::time::Instant::now(), + } + } + + pub fn session_id(&self) -> &str { + &self.session_id + } + + pub fn turns(&self) -> u64 { + self.turn_index + } + + pub fn set_vendor_session_id(&mut self, id: Option) { + if id.is_some() { + self.vendor_session_id = id; + } + } + + /// A root `voice.turn`, linked to the session trace, with attribution recorded. + fn open_turn(&mut self, component: &'static str) -> Span { + let span = crate::voice_turn_span!( + parent: None, + capability = REALTIME_CAPABILITY, + transport = "websocket" + ); + let link = self.session_span.context().span().span_context().clone(); + if link.is_valid() { + span.add_link(link); + } + record_attribution(&span, &self.attribution, &self.session_id); + record_text( + &span, + rt::VENDOR_SESSION_ID, + self.vendor_session_id.as_deref(), + ); + span.record(turn::TURN_INDEX, self.turn_index); + span.record(rt::COMPONENT, component); + self.turn_index += 1; + span + } + + fn account(&mut self, c: &RealtimeCost) { + if let (Some(cost), Some(unit)) = (c.cost, c.unit) { + self.cost_total += cost; + self.priced_any = true; + self.pricing_unit.get_or_insert(unit); + } + } + + /// One `response.done` (§5.5 tap). + pub fn response_done(&mut self, event: &Value) { + let response = event.get("response"); + let usage = response + .and_then(|r| r.get("usage")) + .and_then(RealtimeUsage::from_openai) + .unwrap_or_default(); + let cost = realtime_response_cost(self.pricing.as_ref(), &usage); + let span = self.open_turn("response"); + record_text( + &span, + rt::RESPONSE_ID, + response.and_then(|r| r.get("id")).and_then(Value::as_str), + ); + let status = response + .and_then(|r| r.get("status")) + .and_then(Value::as_str); + record_text(&span, rt::RESPONSE_STATUS, status); + if status == Some("failed") { + span.record("otel.status_code", "ERROR"); + span.record(turn::ERROR_TYPE, "vendor_error"); + } + record_usage(&span, &usage); + record_cost(&span, &cost); + if self.capture + && let Some(t) = response_transcript(event) + { + span.record( + turn::TRANSCRIPT, + crate::observability::trace_redact::sanitize_body(&t).as_str(), + ); + } + self.totals.add(&usage); + self.account(&cost); + } + + /// One `conversation.item.input_audio_transcription.completed`. + pub fn transcription_completed(&mut self, event: &Value) { + let usage = event.get("usage").and_then(TranscriptionUsage::from_openai); + let span = self.open_turn("input_transcription"); + let cost = match usage { + Some(u) => { + match u { + TranscriptionUsage::Seconds(s) => { + span.record(rt::BILLED_SECONDS, s); + self.billed_seconds += s; + } + TranscriptionUsage::Tokens { + input_audio, + input_text, + output_text, + } => { + let t = RealtimeUsage { + input_audio, + input_text, + output_text, + ..Default::default() + }; + record_usage(&span, &t); + } + } + realtime_transcription_cost(self.pricing.as_ref(), &u) + } + None => RealtimeCost { + cost: None, + unit: None, + unpriced: Vec::new(), + }, + }; + record_cost(&span, &cost); + if self.capture + && let Some(t) = event.get("transcript").and_then(Value::as_str) + { + span.record( + turn::TRANSCRIPT, + crate::observability::trace_redact::sanitize_body(t).as_str(), + ); + } + self.account(&cost); + } + + /// One duration segment under a minute/second price (D-9). + pub fn duration_segment(&mut self, seconds: f64) { + if seconds <= 0.0 { + return; + } + let cost = realtime_duration_cost(self.pricing.as_ref(), seconds); + let span = self.open_turn("duration_segment"); + span.record(rt::BILLED_SECONDS, seconds); + record_cost(&span, &cost); + self.billed_seconds += seconds; + self.account(&cost); + } + + /// Close the session span with its totals (FR-MET-3). + pub fn finish(self, end_reason: &str, close_code: u16) { + let span = &self.session_span; + record_text( + span, + rt::VENDOR_SESSION_ID, + self.vendor_session_id.as_deref(), + ); + span.record( + sess::DURATION_MS, + self.started.elapsed().as_secs_f64() * 1000.0, + ); + span.record(sess::TURNS, self.turn_index); + span.record(sess::END_REASON, end_reason); + span.record(sess::CLOSE_CODE, u64::from(close_code)); + record_usage(span, &self.totals); + if self.billed_seconds > 0.0 { + span.record(rt::BILLED_SECONDS, self.billed_seconds); + } + if self.priced_any { + span.record(turn::COST, self.cost_total); + if let Some(unit) = self.pricing_unit { + span.record(turn::PRICING_UNIT, unit); + } + } + if matches!(end_reason, "upstream_error" | "error") { + span.record("otel.status_code", "ERROR"); + } + // Dropping `self` ends the span. + } +} diff --git a/gateway/src/handlers/openai_realtime/mod.rs b/gateway/src/handlers/openai_realtime/mod.rs new file mode 100644 index 00000000..9993283c --- /dev/null +++ b/gateway/src/handlers/openai_realtime/mod.rs @@ -0,0 +1,29 @@ +//! `/v1/realtime`: OpenAI Realtime GA on Bud deployments (FRD-023, D-1, D-2). +//! +//! `wss://gateway/v1/realtime?model=` speaks the OpenAI Realtime GA protocol, so the +//! unmodified OpenAI SDKs, the Agents SDKs, LiveKit and Pipecat work against it by changing only +//! the base URL and the key. For vendors that already speak GA (OpenAI, Azure OpenAI) frames are +//! RELAYED after policy — full fidelity, no per-event translation. WaaV's native protocol stays +//! on `/realtime`. +//! +//! * [`handshake`] — credential sources and the pre-upgrade rules (§5.2). +//! * [`upstream`] — the vendor URL, auth header and SSRF check (§5.3). +//! * [`policy`] — what crosses the relay, per event (§5.5, §5.6). +//! * [`metering`] — a `voice.turn` per billed record, a `voice.session` per session (§5.10). +//! * [`session`] — the session engine: admission, relay, timers, revalidation, teardown (§5.4). +//! * [`client_secrets`] — `POST /v1/realtime/client_secrets`, the `ek_bud_` mint (§5.8). + +pub mod client_secrets; +pub mod handshake; +pub mod metering; +pub mod policy; +pub mod session; +pub mod upstream; + +pub use client_secrets::client_secrets_handler; +pub use session::{RealtimeRuntime, Timings, realtime_ws_handler}; + +/// The public path (FR-RT-1). The ingress already routes the whole `/v1/realtime` prefix here. +pub const OPENAI_REALTIME_PATH: &str = "/v1/realtime"; +/// The client-secret mint path (§5.8). +pub const CLIENT_SECRETS_PATH: &str = "/v1/realtime/client_secrets"; diff --git a/gateway/src/handlers/openai_realtime/policy.rs b/gateway/src/handlers/openai_realtime/policy.rs new file mode 100644 index 00000000..fb8d77fc --- /dev/null +++ b/gateway/src/handlers/openai_realtime/policy.rs @@ -0,0 +1,797 @@ +//! What crosses the relay, per event (FRD-023 §5.5, §5.6, D-11, D-12, S-5, S-6). +//! +//! Pure functions over the event text. The session loop calls [`client_event`] for every client +//! frame and [`vendor_event`] for every vendor frame; nothing here performs I/O. +//! +//! **Two-stage parse (R-2).** Every frame's envelope (`type`, `event_id`) is read; the body only +//! for the event types a rule inspects. Audio-bearing events — `input_audio_buffer.append` from the +//! client, `response.output_audio.delta` from the vendor — are forwarded as the original text, +//! never re-serialised. +//! +//! **Refuse loudly, strip only the routine.** A refused client event is not forwarded and the +//! client gets an `error` naming the field; the session continues. Only `model` and `tracing` are +//! removed silently, because the SDKs send them on every `session.update`. + +use std::borrow::Cow; + +use bud_auth::{RealtimePolicy, RealtimeSettings}; +use serde::Deserialize; +use serde_json::{Map, Value}; + +/// The first stage: just enough to route. +#[derive(Deserialize)] +struct Envelope<'a> { + #[serde(rename = "type", borrow)] + kind: Cow<'a, str>, + #[serde(default, borrow)] + event_id: Option>, +} + +/// Client events forwarded on the envelope alone. +const PASS_THROUGH: &[&str] = &[ + "input_audio_buffer.append", + "input_audio_buffer.commit", + "input_audio_buffer.clear", + "conversation.item.retrieve", + "conversation.item.truncate", + "conversation.item.delete", + "response.cancel", + "output_audio_buffer.clear", +]; + +/// Client events the rules inspect. +const INSPECTED: &[&str] = &[ + "session.update", + "response.create", + "conversation.item.create", +]; + +/// A deployment's client-facing rules. +#[derive(Debug, Clone, Default)] +pub struct ClientRules { + pub policy: RealtimePolicy, + /// `realtime` or `transcription`: a client may not change it. + pub session_type: String, + /// The deployment's output cap, when it sets one. + pub max_output_tokens: Option, +} + +impl ClientRules { + pub fn from_settings(settings: Option<&RealtimeSettings>) -> Self { + let settings = settings.cloned().unwrap_or_default(); + Self { + policy: settings.policy(), + session_type: settings + .session_type + .clone() + .unwrap_or_else(|| "realtime".to_string()), + max_output_tokens: settings.defaults.as_ref().and_then(|d| d.max_output_tokens), + } + } +} + +/// A policy refusal: the field, and what to tell the client. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Refusal { + pub param: String, + pub message: String, + pub event_id: Option, +} + +/// What to do with one client frame. +#[derive(Debug, PartialEq, Eq)] +pub enum ClientOutcome<'a> { + /// Forward this text. `response_create` asks the loop to revalidate first (D-17). + Forward { + text: Cow<'a, str>, + kind: String, + response_create: bool, + /// The type is not one this gateway knows (forwarded; counted). + unknown: bool, + }, + Refuse(Refusal), + /// Not a JSON event at all. + Invalid(String), +} + +fn refusal( + param: &str, + message: impl Into, + event_id: Option<&str>, +) -> ClientOutcome<'static> { + ClientOutcome::Refuse(Refusal { + param: param.to_string(), + message: message.into(), + event_id: event_id.map(str::to_string), + }) +} + +fn has_mcp_tool(tools: Option<&Value>) -> bool { + tools.and_then(Value::as_array).is_some_and(|t| { + t.iter() + .any(|tool| tool.get("type").and_then(Value::as_str) == Some("mcp")) + }) +} + +fn has_image(content: Option<&Value>) -> bool { + content.and_then(Value::as_array).is_some_and(|parts| { + parts + .iter() + .any(|p| p.get("type").and_then(Value::as_str) == Some("input_image")) + }) +} + +/// Clamp a `max_output_tokens` (a number or `"inf"`) to the deployment's cap. +fn clamp_output_tokens(obj: &mut Map, cap: Option) -> bool { + let Some(cap) = cap else { return false }; + let over = match obj.get("max_output_tokens") { + None => false, + Some(Value::String(s)) => s == "inf", + Some(v) => v.as_u64().is_none_or(|n| n > u64::from(cap)), + }; + if over { + obj.insert("max_output_tokens".into(), Value::from(cap)); + } + over +} + +/// Rules shared by a session and a per-response override (`instructions`, `tools`, `prompt`). +fn check_overridable( + obj: &Map, + prefix: &str, + rules: &ClientRules, + event_id: Option<&str>, +) -> Option> { + if obj.get("instructions").is_some() && !rules.policy.allows_client_instructions() { + return Some(refusal( + &format!("{prefix}.instructions"), + "This deployment fixes its instructions; `instructions` may not be set by the client.", + event_id, + )); + } + if has_mcp_tool(obj.get("tools")) && !rules.policy.allows_mcp_tools() { + return Some(refusal( + &format!("{prefix}.tools.mcp"), + "MCP tools are not enabled on this deployment (they run in the vendor's organization).", + event_id, + )); + } + if obj.get("prompt").is_some_and(|p| !p.is_null()) && !rules.policy.allows_prompt_references() { + return Some(refusal( + &format!("{prefix}.prompt"), + "Stored prompt references are not enabled on this deployment.", + event_id, + )); + } + None +} + +fn apply_session_update( + event: &mut Value, + rules: &ClientRules, + event_id: Option<&str>, +) -> Option> { + let Some(session) = event.get_mut("session").and_then(Value::as_object_mut) else { + return None; + }; + // Routine: SDKs send both on every update. `model` is the deployment's (D-3); `tracing` + // would open a trace in the vendor org every project on the credential shares (D-11). + session.remove("model"); + session.remove("tracing"); + + if let Some(kind) = session.get("type").and_then(Value::as_str) + && kind != rules.session_type + { + return Some(refusal( + "session.type", + format!( + "This deployment serves `{}` sessions; `session.type` cannot be changed.", + rules.session_type + ), + event_id, + )); + } + if let Some(refused) = check_overridable(session, "session", rules, event_id) { + return Some(refused); + } + if let Some(model) = session + .get("audio") + .and_then(|a| a.get("input")) + .and_then(|i| i.get("transcription")) + .and_then(|t| t.get("model")) + .and_then(Value::as_str) + && !rules.policy.allows_transcription_model(model) + { + return Some(refusal( + "session.audio.input.transcription.model", + format!("Transcription model '{model}' is not enabled on this deployment."), + event_id, + )); + } + clamp_output_tokens(session, rules.max_output_tokens); + None +} + +fn apply_response_create( + event: &mut Value, + rules: &ClientRules, + event_id: Option<&str>, +) -> Option> { + let Some(response) = event.get_mut("response").and_then(Value::as_object_mut) else { + return None; + }; + response.remove("model"); + if let Some(refused) = check_overridable(response, "response", rules, event_id) { + return Some(refused); + } + if !rules.policy.allows_image_input() + && response + .get("input") + .and_then(Value::as_array) + .is_some_and(|items| items.iter().any(|i| has_image(i.get("content")))) + { + return Some(refusal( + "response.input.content.input_image", + "Image input is not enabled on this deployment.", + event_id, + )); + } + clamp_output_tokens(response, rules.max_output_tokens); + None +} + +/// Apply the client → vendor rules to one frame (FRD §5.5). +pub fn client_event<'a>(raw: &'a str, rules: &ClientRules) -> ClientOutcome<'a> { + let env: Envelope = match serde_json::from_str(raw) { + Ok(e) => e, + Err(e) => return ClientOutcome::Invalid(format!("not a JSON event with a `type`: {e}")), + }; + let kind = env.kind.to_string(); + + if PASS_THROUGH.contains(&kind.as_str()) { + return ClientOutcome::Forward { + text: Cow::Borrowed(raw), + kind, + response_create: false, + unknown: false, + }; + } + if !INSPECTED.contains(&kind.as_str()) { + // Forward-compatible with vendor additions (R-6); counted by the caller. + return ClientOutcome::Forward { + text: Cow::Borrowed(raw), + kind, + response_create: false, + unknown: true, + }; + } + + let event_id = env.event_id.map(|e| e.to_string()); + let mut event: Value = match serde_json::from_str(raw) { + Ok(v) => v, + Err(e) => return ClientOutcome::Invalid(e.to_string()), + }; + let refused = match kind.as_str() { + "session.update" => apply_session_update(&mut event, rules, event_id.as_deref()), + "response.create" => apply_response_create(&mut event, rules, event_id.as_deref()), + "conversation.item.create" => { + let image = has_image(event.get("item").and_then(|i| i.get("content"))); + (image && !rules.policy.allows_image_input()).then(|| { + refusal( + "item.content.input_image", + "Image input is not enabled on this deployment.", + event_id.as_deref(), + ) + }) + } + _ => None, + }; + if let Some(refused) = refused { + return refused; + } + ClientOutcome::Forward { + text: Cow::Owned(event.to_string()), + response_create: kind == "response.create", + kind, + unknown: false, + } +} + +/// What a vendor frame means to the session, beyond forwarding it. +#[derive(Debug, PartialEq)] +pub enum Tap { + None, + SessionCreated { vendor_session_id: Option }, + SessionUpdated { event_id: Option }, + ResponseDone(Value), + TranscriptionCompleted(Value), + Error(Value), +} + +/// What to do with one vendor frame. +#[derive(Debug, PartialEq)] +pub enum VendorOutcome<'a> { + Forward(Cow<'a, str>, Tap), + /// Never reaches the client (D-12). + Drop, + Invalid, +} + +/// Rewrite `session.model` to the name the client connected with (§5.5 taps). +fn rewrite_session_model(event: &mut Value, deployment: &str) { + if let Some(session) = event.get_mut("session").and_then(Value::as_object_mut) + && session.contains_key("model") + { + session.insert("model".into(), Value::from(deployment)); + } +} + +/// Apply the vendor → client taps to one frame (FRD §5.5). +pub fn vendor_event<'a>(raw: &'a str, deployment: &str) -> VendorOutcome<'a> { + let env: Envelope = match serde_json::from_str(raw) { + Ok(e) => e, + Err(_) => return VendorOutcome::Invalid, + }; + match env.kind.as_ref() { + // It reports the VENDOR ORG's remaining budget across every tenant on the credential. + "rate_limits.updated" => VendorOutcome::Drop, + "session.created" | "session.updated" => { + let created = env.kind == "session.created"; + let event_id = env.event_id.map(|e| e.to_string()); + let Ok(mut event) = serde_json::from_str::(raw) else { + return VendorOutcome::Invalid; + }; + let vendor_session_id = event + .get("session") + .and_then(|s| s.get("id")) + .and_then(Value::as_str) + .map(str::to_string); + rewrite_session_model(&mut event, deployment); + let tap = if created { + Tap::SessionCreated { vendor_session_id } + } else { + Tap::SessionUpdated { event_id } + }; + VendorOutcome::Forward(Cow::Owned(event.to_string()), tap) + } + "response.done" => match serde_json::from_str::(raw) { + Ok(v) => VendorOutcome::Forward(Cow::Borrowed(raw), Tap::ResponseDone(v)), + Err(_) => VendorOutcome::Invalid, + }, + "conversation.item.input_audio_transcription.completed" => { + match serde_json::from_str::(raw) { + Ok(v) => VendorOutcome::Forward(Cow::Borrowed(raw), Tap::TranscriptionCompleted(v)), + Err(_) => VendorOutcome::Invalid, + } + } + "error" => match serde_json::from_str::(raw) { + Ok(v) => VendorOutcome::Forward(Cow::Borrowed(raw), Tap::Error(v)), + Err(_) => VendorOutcome::Invalid, + }, + _ => VendorOutcome::Forward(Cow::Borrowed(raw), Tap::None), + } +} + +/// The `session.update` WaaV sends first, built from the deployment's defaults (FRD §5.6). +/// +/// `None` when a realtime deployment has no defaults — the vendor's apply. A transcription +/// deployment always gets one: its session type and model are set here. +pub fn defaults_update( + settings: Option<&RealtimeSettings>, + vendor_model: Option<&str>, + event_id: &str, +) -> Option { + let transcription = settings.is_some_and(RealtimeSettings::is_transcription); + let defaults = settings + .and_then(|s| s.defaults.clone()) + .unwrap_or_default(); + + let mut session = Map::new(); + session.insert( + "type".into(), + Value::from(if transcription { + "transcription" + } else { + "realtime" + }), + ); + let mut input = Map::new(); + let mut output = Map::new(); + + let mut transcription_cfg = Map::new(); + if let Some(t) = &defaults.input_transcription { + if let Some(m) = &t.model { + transcription_cfg.insert("model".into(), Value::from(m.as_str())); + } + if let Some(l) = &t.language { + transcription_cfg.insert("language".into(), Value::from(l.as_str())); + } + if let Some(p) = &t.prompt { + transcription_cfg.insert("prompt".into(), Value::from(p.as_str())); + } + } + if transcription && let Some(model) = vendor_model { + // A transcription deployment's model IS its transcription model. + transcription_cfg.insert("model".into(), Value::from(model)); + } + if !transcription_cfg.is_empty() { + input.insert("transcription".into(), Value::Object(transcription_cfg)); + } + if let Some(td) = &defaults.turn_detection { + input.insert("turn_detection".into(), td.clone()); + } + if let Some(nr) = &defaults.noise_reduction { + input.insert("noise_reduction".into(), serde_json::json!({ "type": nr })); + } + if !transcription { + if let Some(v) = &defaults.voice { + output.insert("voice".into(), Value::from(v.as_str())); + } + if let Some(s) = defaults.speed { + output.insert("speed".into(), Value::from(s)); + } + if let Some(i) = &defaults.instructions { + session.insert("instructions".into(), Value::from(i.as_str())); + } + if let Some(m) = &defaults.output_modalities { + session.insert("output_modalities".into(), Value::from(m.clone())); + } + if let Some(m) = defaults.max_output_tokens { + session.insert("max_output_tokens".into(), Value::from(m)); + } + } + + let mut audio = Map::new(); + if !input.is_empty() { + audio.insert("input".into(), Value::Object(input)); + } + if !output.is_empty() { + audio.insert("output".into(), Value::Object(output)); + } + if !audio.is_empty() { + session.insert("audio".into(), Value::Object(audio)); + } + if !transcription && session.len() == 1 { + return None; // only `type`: nothing to apply + } + Some( + serde_json::json!({ + "type": "session.update", + "event_id": event_id, + "session": Value::Object(session), + }) + .to_string(), + ) +} + +/// An OpenAI `error` event from the gateway. +pub fn error_event( + event_id: &str, + kind: &str, + code: &str, + message: &str, + param: Option<&str>, + client_event_id: Option<&str>, +) -> String { + let mut error = serde_json::json!({"type": kind, "code": code, "message": message}); + if let Some(p) = param { + error["param"] = Value::from(p); + } + if let Some(e) = client_event_id { + error["event_id"] = Value::from(e); + } + serde_json::json!({"type": "error", "event_id": event_id, "error": error}).to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + + fn rules() -> ClientRules { + ClientRules { + policy: RealtimePolicy { + input_transcription_models: Some(vec!["gpt-4o-mini-transcribe".into()]), + ..Default::default() + }, + session_type: "realtime".into(), + max_output_tokens: Some(4096), + } + } + + fn forwarded(outcome: ClientOutcome<'_>) -> Value { + match outcome { + ClientOutcome::Forward { text, .. } => serde_json::from_str(&text).unwrap(), + other => panic!("expected forward, got {other:?}"), + } + } + + fn refused(outcome: ClientOutcome<'_>) -> Refusal { + match outcome { + ClientOutcome::Refuse(r) => r, + other => panic!("expected refusal, got {other:?}"), + } + } + + /// TC-EVT-01 + #[test] + fn tc_evt_01_model_is_stripped_silently() { + let out = forwarded(client_event( + r#"{"type":"session.update","session":{"type":"realtime","model":"gpt-4o-realtime-preview","voice":"x"}}"#, + &rules(), + )); + assert!(out["session"].get("model").is_none()); + assert_eq!(out["session"]["voice"], "x"); + } + + /// TC-EVT-02 🔒 / TC-EVT-03 + #[test] + fn tc_evt_02_mcp_tools_are_refused_unless_enabled() { + let raw = r#"{"type":"session.update","event_id":"c1","session":{"tools":[{"type":"mcp","server_url":"https://x"}]}}"#; + let r = refused(client_event(raw, &rules())); + assert_eq!(r.param, "session.tools.mcp"); + assert_eq!(r.event_id.as_deref(), Some("c1")); + + let mut open = rules(); + open.policy.allow_mcp_tools = Some(true); + forwarded(client_event(raw, &open)); + } + + /// TC-EVT-04 🔒 + #[test] + fn tc_evt_04_stored_prompts_are_refused_by_default() { + let r = refused(client_event( + r#"{"type":"session.update","session":{"prompt":{"id":"pmpt_x"}}}"#, + &rules(), + )); + assert_eq!(r.param, "session.prompt"); + } + + /// TC-EVT-05 + #[test] + fn tc_evt_05_tracing_is_stripped() { + let out = forwarded(client_event( + r#"{"type":"session.update","session":{"tracing":"auto"}}"#, + &rules(), + )); + assert!(out["session"].get("tracing").is_none()); + } + + /// TC-EVT-06 🔒 + #[test] + fn tc_evt_06_transcription_model_allowlist() { + let r = refused(client_event( + r#"{"type":"session.update","session":{"audio":{"input":{"transcription":{"model":"gpt-4o-transcribe"}}}}}"#, + &rules(), + )); + assert_eq!(r.param, "session.audio.input.transcription.model"); + forwarded(client_event( + r#"{"type":"session.update","session":{"audio":{"input":{"transcription":{"model":"gpt-4o-mini-transcribe"}}}}}"#, + &rules(), + )); + } + + /// TC-EVT-07 + #[test] + fn tc_evt_07_session_type_is_locked() { + let r = refused(client_event( + r#"{"type":"session.update","session":{"type":"transcription"}}"#, + &rules(), + )); + assert_eq!(r.param, "session.type"); + } + + /// TC-EVT-08 + #[test] + fn tc_evt_08_instructions_lock_applies_to_session_and_response() { + let mut locked = rules(); + locked.policy.allow_client_instructions = Some(false); + let r = refused(client_event( + r#"{"type":"session.update","session":{"instructions":"be rude"}}"#, + &locked, + )); + assert_eq!(r.param, "session.instructions"); + let r = refused(client_event( + r#"{"type":"response.create","response":{"instructions":"be rude"}}"#, + &locked, + )); + assert_eq!(r.param, "response.instructions"); + } + + /// TC-EVT-09 + #[test] + fn tc_evt_09_output_cap_is_clamped() { + let out = forwarded(client_event( + r#"{"type":"session.update","session":{"max_output_tokens":"inf"}}"#, + &rules(), + )); + assert_eq!(out["session"]["max_output_tokens"], 4096); + let out = forwarded(client_event( + r#"{"type":"response.create","response":{"max_output_tokens":100000}}"#, + &rules(), + )); + assert_eq!(out["response"]["max_output_tokens"], 4096); + let out = forwarded(client_event( + r#"{"type":"response.create","response":{"max_output_tokens":10}}"#, + &rules(), + )); + assert_eq!(out["response"]["max_output_tokens"], 10); + } + + /// TC-EVT-10 + #[test] + fn tc_evt_10_image_input_refused_when_disabled() { + let mut no_images = rules(); + no_images.policy.allow_image_input = Some(false); + let raw = r#"{"type":"conversation.item.create","item":{"type":"message","role":"user","content":[{"type":"input_image","image_url":"data:x"}]}}"#; + assert_eq!( + refused(client_event(raw, &no_images)).param, + "item.content.input_image" + ); + forwarded(client_event(raw, &rules())); + } + + /// TC-EVT-11 + #[test] + fn tc_evt_11_unknown_events_are_forwarded_and_flagged() { + match client_event(r#"{"type":"future.event","x":1}"#, &rules()) { + ClientOutcome::Forward { unknown, text, .. } => { + assert!(unknown); + assert_eq!(text, r#"{"type":"future.event","x":1}"#); + } + other => panic!("{other:?}"), + } + } + + #[test] + fn audio_frames_are_forwarded_byte_for_byte() { + let raw = r#"{"type":"input_audio_buffer.append","audio":"AAAA//8="}"#; + match client_event(raw, &rules()) { + ClientOutcome::Forward { text, .. } => { + assert!(matches!(text, Cow::Borrowed(t) if t == raw)) + } + other => panic!("{other:?}"), + } + } + + #[test] + fn response_create_asks_for_revalidation() { + match client_event(r#"{"type":"response.create"}"#, &rules()) { + ClientOutcome::Forward { + response_create, .. + } => assert!(response_create), + other => panic!("{other:?}"), + } + } + + #[test] + fn junk_is_invalid_not_forwarded() { + assert!(matches!( + client_event("not json", &rules()), + ClientOutcome::Invalid(_) + )); + assert!(matches!( + client_event(r#"{"no":"type"}"#, &rules()), + ClientOutcome::Invalid(_) + )); + } + + /// TC-EVT-12 🔒 + #[test] + fn tc_evt_12_vendor_rate_limits_are_dropped() { + assert_eq!( + vendor_event(r#"{"type":"rate_limits.updated","rate_limits":[]}"#, "rt"), + VendorOutcome::Drop + ); + } + + /// TC-EVT-13 + #[test] + fn tc_evt_13_session_model_is_rewritten_to_the_deployment() { + match vendor_event( + r#"{"type":"session.created","session":{"id":"sess_1","model":"gpt-realtime-2.1"}}"#, + "my-rt", + ) { + VendorOutcome::Forward(text, Tap::SessionCreated { vendor_session_id }) => { + let v: Value = serde_json::from_str(&text).unwrap(); + assert_eq!(v["session"]["model"], "my-rt"); + assert_eq!(vendor_session_id.as_deref(), Some("sess_1")); + } + other => panic!("{other:?}"), + } + } + + #[test] + fn response_done_is_tapped_and_forwarded_untouched() { + let raw = r#"{"type":"response.done","response":{"id":"r1","usage":{}}}"#; + match vendor_event(raw, "rt") { + VendorOutcome::Forward(text, Tap::ResponseDone(v)) => { + assert_eq!(text, raw); + assert_eq!(v["response"]["id"], "r1"); + } + other => panic!("{other:?}"), + } + } + + #[test] + fn audio_deltas_pass_through_untouched() { + let raw = r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#; + assert_eq!( + vendor_event(raw, "rt"), + VendorOutcome::Forward(Cow::Borrowed(raw), Tap::None) + ); + } + + /// TC-EVT-14 (builder half): the defaults update carries the deployment's voice and settings. + #[test] + fn tc_evt_14_defaults_update_shape() { + let settings: RealtimeSettings = serde_json::from_value(serde_json::json!({ + "session_type": "realtime", + "defaults": { + "voice": "marin", "instructions": "hi", "output_modalities": ["audio"], + "turn_detection": {"type": "semantic_vad", "eagerness": "auto"}, + "input_transcription": {"model": "gpt-4o-mini-transcribe", "language": "en"}, + "noise_reduction": "near_field", "max_output_tokens": 4096, "speed": 1.0 + } + })) + .unwrap(); + let update: Value = serde_json::from_str( + &defaults_update(Some(&settings), Some("gpt-realtime-2.1"), "evt_bud_d1").unwrap(), + ) + .unwrap(); + assert_eq!(update["type"], "session.update"); + assert_eq!(update["event_id"], "evt_bud_d1"); + let s = &update["session"]; + assert_eq!(s["type"], "realtime"); + assert_eq!(s["instructions"], "hi"); + assert_eq!(s["max_output_tokens"], 4096); + assert_eq!(s["audio"]["output"]["voice"], "marin"); + assert_eq!(s["audio"]["output"]["speed"], 1.0); + assert_eq!( + s["audio"]["input"]["turn_detection"]["type"], + "semantic_vad" + ); + assert_eq!( + s["audio"]["input"]["transcription"]["model"], + "gpt-4o-mini-transcribe" + ); + assert_eq!(s["audio"]["input"]["noise_reduction"]["type"], "near_field"); + assert!(s.get("model").is_none()); + } + + #[test] + fn a_realtime_deployment_without_defaults_sends_no_update() { + assert!(defaults_update(None, Some("m"), "e").is_none()); + } + + #[test] + fn a_transcription_deployment_always_sets_its_type_and_model() { + let settings: RealtimeSettings = + serde_json::from_value(serde_json::json!({"session_type": "transcription"})).unwrap(); + let update: Value = serde_json::from_str( + &defaults_update(Some(&settings), Some("gpt-4o-transcribe"), "e").unwrap(), + ) + .unwrap(); + assert_eq!(update["session"]["type"], "transcription"); + assert_eq!( + update["session"]["audio"]["input"]["transcription"]["model"], + "gpt-4o-transcribe" + ); + } + + #[test] + fn the_error_event_is_openai_shaped() { + let e: Value = serde_json::from_str(&error_event( + "evt_bud_1", + "invalid_request_error", + "event_not_allowed", + "no", + Some("session.prompt"), + Some("c9"), + )) + .unwrap(); + assert_eq!(e["type"], "error"); + assert_eq!(e["error"]["code"], "event_not_allowed"); + assert_eq!(e["error"]["param"], "session.prompt"); + assert_eq!(e["error"]["event_id"], "c9"); + } +} diff --git a/gateway/src/handlers/openai_realtime/session.rs b/gateway/src/handlers/openai_realtime/session.rs new file mode 100644 index 00000000..0ee71f17 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/session.rs @@ -0,0 +1,1015 @@ +//! One relayed session: upgrade → authenticate → resolve → admit → connect → relay → close +//! (FRD-023 §5.4, §5.5, §5.6; D-8, D-13, D-14, D-17). +//! +//! Everything that can be refused is refused BEFORE the upgrade, as HTTP with OpenAI's error +//! envelope. After the upgrade a failure is an `error` event followed by a close with the D-14 +//! code, and every exit path runs the same teardown: the session span, the metrics, and — by +//! dropping — the deployment's concurrency slot and the connection slot. + +use std::collections::VecDeque; +use std::sync::Arc; +use std::time::Duration; + +use axum::extract::ws::{CloseFrame, Message, WebSocket, WebSocketUpgrade}; +use axum::extract::{Extension, State}; +use axum::http::{HeaderMap, StatusCode, Uri}; +use axum::response::{IntoResponse, Response}; +use futures_util::{SinkExt, StreamExt}; +use tokio::sync::mpsc; +use tokio::time::Instant; +use tokio_tungstenite::tungstenite::Message as UpMessage; +use tracing::{debug, info, warn}; + +use bud_auth::{AliasMetadata, Principal, PrincipalKind, RealtimeSettings, VoiceEndpoint}; + +use crate::auth::ephemeral::{self, OpenError, Parent}; +use crate::core::deployment_policy::{Admission, Rejection, vendor_key}; +use crate::middleware::connection_limit::ConnectionSlot; +use crate::state::AppState; + +use super::handshake::{ + self, Credential, HandshakeError, MAX_MESSAGE_BYTES, REALTIME_CAPABILITY, SUBPROTOCOL, +}; +use super::metering::{Attribution, SessionMeter}; +use super::policy::{self, ClientOutcome, ClientRules, Tap, VendorOutcome}; +use super::upstream::{self, UpstreamRequest}; + +/// Session timings (D-13). From the environment, with the FRD's defaults; tests shorten them. +#[derive(Debug, Clone)] +pub struct Timings { + pub ping: Duration, + pub max_missed_pongs: u32, + pub revalidate: Duration, + pub connect: Duration, + /// How long client frames are held for the vendor to apply the deployment defaults. + pub hold: Duration, + /// How long a full client-bound queue is tolerated before `client_too_slow`. + pub slow_client: Duration, + /// The ceiling on any deployment's maximum session length. + pub max_session: Duration, + pub default_idle: Duration, + pub warn_before: Duration, + pub segment: Duration, + pub upstream_send: Duration, +} + +impl Default for Timings { + fn default() -> Self { + Self { + ping: Duration::from_secs(20), + max_missed_pongs: 3, + revalidate: Duration::from_secs(30), + connect: Duration::from_secs(10), + hold: Duration::from_secs(5), + slow_client: Duration::from_secs(5), + max_session: Duration::from_secs(3600), + default_idle: Duration::from_secs(300), + warn_before: Duration::from_secs(60), + segment: Duration::from_secs(60), + upstream_send: Duration::from_secs(10), + } + } +} + +impl Timings { + pub fn from_env() -> Self { + Self::from_lookup(|k| std::env::var(k).ok()) + } + + pub fn from_lookup(get: impl Fn(&str) -> Option) -> Self { + let d = Self::default(); + let secs = |k: &str, dflt: Duration| { + get(k) + .and_then(|v| v.trim().parse::().ok()) + .filter(|v| *v > 0) + .map(Duration::from_secs) + .unwrap_or(dflt) + }; + Self { + ping: secs("WAAV_REALTIME_PING_SECS", d.ping), + revalidate: secs("WAAV_REALTIME_REVALIDATE_SECS", d.revalidate), + max_session: secs("WAAV_REALTIME_MAX_SESSION_SECS", d.max_session), + default_idle: secs("WAAV_REALTIME_IDLE_SECS", d.default_idle), + ..d + } + } +} + +/// Process-wide realtime state, built once at startup. +#[derive(Debug, Default)] +pub struct RealtimeRuntime { + pub timings: Timings, + /// `None`: client secrets are not configured (the mint route answers 501, `ek_bud_` refused). + pub client_secret_keys: Option, +} + +impl RealtimeRuntime { + /// Fails on a bad `WAAV_CLIENT_SECRET_KEYS` (TC-EK-16): a pod must not start with keys it + /// cannot use. + pub fn from_env() -> Result { + Ok(Self { + timings: Timings::from_env(), + client_secret_keys: ephemeral::ClientSecretKeys::from_env()?, + }) + } +} + +/// How a caller is re-checked during the session (D-17). +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CallerCheck { + ApiKey { hashed: String, client_key: bool }, + Jwt { sub: String }, +} + +/// An authenticated caller. +#[derive(Debug, Clone)] +pub struct Caller { + pub principal: Principal, + pub check: CallerCheck, + /// Minted by a client secret (the principal is the parent's). + pub via_client_secret: bool, +} + +/// Does the caller still reach the endpoint, and is it still a realtime deployment? +pub async fn still_allowed( + state: &AppState, + check: &CallerCheck, + endpoint_id: &str, +) -> Option { + let plane = state.bud_mode.as_ref()?.plane(); + let serving = plane + .voice_endpoint(endpoint_id) + .is_some_and(|e| e.serves(REALTIME_CAPABILITY)); + if !serving { + return None; + } + match check { + CallerCheck::ApiKey { hashed, client_key } => { + plane.hash_reaches(hashed, endpoint_id, *client_key) + } + CallerCheck::Jwt { sub } => plane.subject_reaches(sub, endpoint_id).await, + } +} + +/// The per-response quota hook (NG-8, §5.4): always `Allow` until FRD-021 Q-1 decides it. +pub fn admit_response(_caller: &Caller, _endpoint_id: &str) -> bool { + true +} + +fn not_found(model: &str) -> HandshakeError { + HandshakeError::new( + StatusCode::NOT_FOUND, + "model_not_found", + format!("Model '{model}' not found or does not support {REALTIME_CAPABILITY}."), + ) + .param("model") +} + +/// Authenticate a Bud key or JWT. +pub async fn authenticate( + state: &AppState, + credential: &Credential, +) -> Result { + let Some(bud) = state.bud_mode.as_ref() else { + return Err(bud_mode_required()); + }; + let raw = credential.expose(); + match bud.plane().authenticate(raw).await { + Ok(principal) => { + let check = match principal.via { + PrincipalKind::ApiKey => CallerCheck::ApiKey { + hashed: bud_auth::hash_api_key(raw), + client_key: raw.starts_with("bud_client"), + }, + PrincipalKind::Jwt => CallerCheck::Jwt { + sub: principal.user_id.clone().unwrap_or_default(), + }, + }; + Ok(Caller { + principal, + check, + via_client_secret: false, + }) + } + Err(bud_auth::AuthFailure::NotReady) => Err(HandshakeError::new( + StatusCode::SERVICE_UNAVAILABLE, + "not_ready", + "The gateway is still loading its control plane; retry shortly.", + ) + .retry_after(1)), + Err(bud_auth::AuthFailure::Throttled) => Err(HandshakeError::new( + StatusCode::TOO_MANY_REQUESTS, + "rate_limit_exceeded", + "Too many authentication failures; retry later.", + ) + .retry_after(1)), + Err(bud_auth::AuthFailure::JwtRejected) => Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_api_key", + "The bearer token (a JWT) was rejected: check that it has not expired and that its \ + client is allowed on this gateway.", + )), + Err(_) => Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_api_key", + "Invalid API key.", + )), + } +} + +fn bud_mode_required() -> HandshakeError { + HandshakeError::new( + StatusCode::NOT_FOUND, + "bud_mode_required", + "/v1/realtime serves Bud deployments and this gateway has no Bud control plane; \ + WaaV's native realtime protocol is at /realtime.", + ) +} + +/// Open and check an `ek_bud_` secret against `?model` and its parent (§5.8 "Validation"). +async fn authenticate_client_secret( + state: &AppState, + credential: &Credential, + model: &str, +) -> Result<(Caller, String, Option), HandshakeError> { + if state.bud_mode.is_none() { + return Err(bud_mode_required()); + } + let Some(keys) = state.realtime.client_secret_keys.as_ref() else { + return Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_client_secret", + "Client secrets are not enabled on this gateway.", + )); + }; + let claims = keys + .open(credential.expose(), ephemeral::now_epoch()) + .map_err(|e| match e { + OpenError::Expired => HandshakeError::new( + StatusCode::UNAUTHORIZED, + "client_secret_expired", + "The client secret has expired; mint a new one.", + ), + _ => HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_client_secret", + "The client secret is not valid.", + ), + })?; + if model != claims.alias && model != claims.ep { + return Err(HandshakeError::new( + StatusCode::FORBIDDEN, + "model_mismatch", + "This client secret was minted for a different deployment.", + ) + .param("model")); + } + let (check, via) = match &claims.parent { + Parent::ApiKey { h, ck } => ( + CallerCheck::ApiKey { + hashed: h.clone(), + client_key: *ck, + }, + PrincipalKind::ApiKey, + ), + Parent::Jwt { sub } => (CallerCheck::Jwt { sub: sub.clone() }, PrincipalKind::Jwt), + }; + let alias = still_allowed(state, &check, &claims.ep).await; + if alias.is_none() { + return Err(HandshakeError::new( + StatusCode::UNAUTHORIZED, + "invalid_client_secret", + "The credential that minted this client secret no longer reaches the deployment.", + )); + } + let caller = Caller { + principal: Principal { + project_id: claims.pid.clone(), + api_key_id: claims.akid.clone(), + user_id: claims.uid.clone(), + via, + expires_at: None, + }, + check, + via_client_secret: true, + }; + Ok((caller, claims.ep, alias)) +} + +/// Everything a session needs, decided before the upgrade. +pub struct Prepared { + pub caller: Caller, + pub endpoint_id: String, + pub endpoint_name: String, + pub endpoint: VoiceEndpoint, + pub alias: Option, + pub settings: Option, + pub rules: ClientRules, + pub upstream: UpstreamRequest, + pub vkey: String, + /// Held for the session: the deployment's concurrency slot (D-8). + pub admission: Admission, +} + +impl std::fmt::Debug for Prepared { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Prepared") + .field("endpoint_id", &self.endpoint_id) + .field("endpoint_name", &self.endpoint_name) + .field("vendor", &self.endpoint.vendor) + .finish() + } +} + +fn rejection_error(rej: Rejection) -> HandshakeError { + let retry = rej.retry_after().as_secs().max(1); + match rej { + Rejection::Rate(_) => HandshakeError::new( + StatusCode::TOO_MANY_REQUESTS, + "rate_limit_exceeded", + "This deployment's rate limit was reached; retry after the indicated delay.", + ), + Rejection::Concurrency(_) => HandshakeError::new( + StatusCode::TOO_MANY_REQUESTS, + "concurrency_limit_exceeded", + "This deployment's concurrent-session limit was reached; retry after the indicated delay.", + ), + } + .retry_after(retry) +} + +/// Authenticate, resolve, admit and build the vendor request (FRD §5.3 steps 1-5). +pub async fn prepare( + state: &AppState, + query: Option<&str>, + headers: &HeaderMap, +) -> Result { + let hs = handshake::parse(query, headers)?; + debug!(model = %hs.model, source = hs.credential.source.as_str(), "realtime handshake parsed"); + if state.bud_mode.is_none() { + return Err(bud_mode_required()); + } + + let (caller, endpoint_id, alias) = if hs.credential.is_client_secret() { + authenticate_client_secret(state, &hs.credential, &hs.model).await? + } else { + let caller = authenticate(state, &hs.credential).await?; + let resolved = state + .resolve_voice_endpoint(&hs.model, REALTIME_CAPABILITY, Some(hs.credential.expose())) + .ok_or_else(|| not_found(&hs.model))?; + (caller, resolved.endpoint_id, resolved.alias) + }; + + debug!(endpoint_id = %endpoint_id, "realtime caller authenticated and endpoint resolved"); + let plane = state + .bud_mode + .as_ref() + .map(|b| b.plane()) + .ok_or_else(bud_mode_required)?; + let endpoint = plane + .voice_endpoint(&endpoint_id) + .filter(|e| e.serves(REALTIME_CAPABILITY)) + .ok_or_else(|| not_found(&hs.model))?; + let settings = endpoint.config.realtime.clone(); + let transcription = settings + .as_ref() + .is_some_and(RealtimeSettings::is_transcription); + + // Build before admitting: a deployment the relay cannot serve must not take a slot. + let upstream_req = upstream::build(&endpoint, transcription).map_err(|e| match e { + upstream::UpstreamError::UnsupportedVendor(_) => HandshakeError::new( + StatusCode::NOT_IMPLEMENTED, + "unsupported_vendor", + e.to_string(), + ), + other => HandshakeError::new(StatusCode::BAD_GATEWAY, "upstream_error", other.to_string()), + })?; + upstream::validate(&upstream_req).await.map_err(|e| { + HandshakeError::new(StatusCode::BAD_GATEWAY, "upstream_error", e.to_string()) + })?; + + debug!(endpoint_id = %endpoint_id, "realtime upstream request validated"); + let vkey = vendor_key(&endpoint.vendor, endpoint.api_base.as_deref()); + if let Some(p) = &state.policies + && let Err(open) = p.breakers().check(&endpoint_id, &vkey) + { + return Err(HandshakeError::new( + StatusCode::SERVICE_UNAVAILABLE, + "circuit_open", + "This deployment's vendor is failing; its circuit breaker is open.", + ) + .retry_after(open.retry_in.as_secs())); + } + + let admission = state + .admit_deployment(&endpoint_id) + .await + .map_err(rejection_error)?; + debug!(endpoint_id = %endpoint_id, "realtime session admitted"); + + Ok(Prepared { + rules: ClientRules::from_settings(settings.as_ref()), + caller, + endpoint_name: hs.model, + endpoint_id, + endpoint, + alias, + settings, + upstream: upstream_req, + vkey, + admission, + }) +} + +/// `GET /v1/realtime` (FR-RT-1). +pub async fn realtime_ws_handler( + State(state): State>, + uri: Uri, + headers: HeaderMap, + slot: Option>, + ws: Result, +) -> Response { + let prepared = match prepare(&state, uri.query(), &headers).await { + Ok(p) => p, + Err(e) => { + metrics::counter!("waav_realtime_refusals_total", "code" => e.code).increment(1); + return e.into_response(); + } + }; + let Ok(ws) = ws else { + return HandshakeError::new( + StatusCode::BAD_REQUEST, + "websocket_required", + "/v1/realtime is a WebSocket endpoint; send an Upgrade request.", + ) + .into_response(); + }; + let slot = slot.map(|Extension(s)| s); + ws.protocols([SUBPROTOCOL]) + .max_message_size(MAX_MESSAGE_BYTES) + .max_frame_size(MAX_MESSAGE_BYTES) + .on_upgrade(move |socket| run(state, prepared, socket, slot)) +} + +/// How a session ended. +#[derive(Debug, Clone)] +pub struct End { + pub reason: &'static str, + pub close_code: u16, + /// The `error` event sent before the close: (code, message). + pub error: Option<(&'static str, String)>, +} + +impl End { + fn new(reason: &'static str, close_code: u16) -> Self { + Self { + reason, + close_code, + error: None, + } + } + + fn with_error(mut self, code: &'static str, message: impl Into) -> Self { + self.error = Some((code, message.into())); + self + } +} + +enum Outbound { + Frame(Message), + Close { + error: Option, + code: u16, + reason: String, + }, +} + +fn next_event_id() -> String { + format!("evt_bud_{}", uuid::Uuid::new_v4().simple()) +} + +fn gateway_error( + code: &str, + message: &str, + param: Option<&str>, + client_event_id: Option<&str>, +) -> String { + let kind = match code { + "upstream_error" | "server_shutdown" | "client_too_slow" => "server_error", + _ => "invalid_request_error", + }; + policy::error_event( + &next_event_id(), + kind, + code, + message, + param, + client_event_id, + ) +} + +/// The client-bound writer: a bounded queue drained by its own task, so a slow client applies +/// backpressure without the relay dropping a frame (FR-EVT-4). +fn spawn_writer( + mut sink: futures_util::stream::SplitSink, + capacity: usize, +) -> (mpsc::Sender, tokio::task::JoinHandle<()>) { + let (tx, mut rx) = mpsc::channel::(capacity); + let task = tokio::spawn(async move { + while let Some(out) = rx.recv().await { + match out { + Outbound::Frame(m) => { + if sink.send(m).await.is_err() { + return; + } + } + Outbound::Close { + error, + code, + reason, + } => { + if let Some(e) = error { + let _ = sink.send(Message::Text(e.into())).await; + } + let _ = sink + .send(Message::Close(Some(CloseFrame { + code, + reason: reason.into(), + }))) + .await; + let _ = sink.close().await; + return; + } + } + } + }); + (tx, task) +} + +/// The client-bound queue: 2 s of 24 kHz audio at 20 ms frames, plus headroom for events. +const CLIENT_QUEUE: usize = 256; +/// Client frames held while the vendor applies the defaults. +const MAX_HELD: usize = 4096; + +struct Relay<'a> { + state: &'a AppState, + p: &'a Prepared, + meter: SessionMeter, + client_tx: mpsc::Sender, + up_tx: futures_util::stream::SplitSink, + timings: Timings, + /// Client frames wait here until the vendor has applied the deployment defaults (§5.6). + held: VecDeque<(String, bool)>, + ready: bool, + /// The `event_id` of the defaults update in flight, and its deadline. + awaiting_defaults: Option, + hold_deadline: Option, + last_activity: Instant, + client_missed: u32, + vendor_missed: u32, +} + +impl Relay<'_> { + async fn to_client(&self, text: String) -> Result<(), End> { + let started = std::time::Instant::now(); + match tokio::time::timeout( + self.timings.slow_client, + self.client_tx + .send(Outbound::Frame(Message::Text(text.into()))), + ) + .await + { + Ok(Ok(())) => { + metrics::histogram!("waav_realtime_relay_latency_seconds", "direction" => "vendor_to_client") + .record(started.elapsed().as_secs_f64()); + Ok(()) + } + Ok(Err(_)) => Err(End::new("client_close", 1006)), + Err(_) => Err(End::new("client_too_slow", 1011).with_error( + "client_too_slow", + "The client did not read the session's output fast enough; audio is never dropped \ + silently, so the session is closed.", + )), + } + } + + async fn to_vendor(&mut self, text: String) -> Result<(), End> { + let started = std::time::Instant::now(); + match tokio::time::timeout( + self.timings.upstream_send, + self.up_tx.send(UpMessage::Text(text.into())), + ) + .await + { + Ok(Ok(())) => { + metrics::histogram!("waav_realtime_relay_latency_seconds", "direction" => "client_to_vendor") + .record(started.elapsed().as_secs_f64()); + Ok(()) + } + _ => Err(End::new("upstream_error", 1011) + .with_error("upstream_error", "The connection to the vendor failed.")), + } + } + + async fn refuse( + &self, + code: &str, + message: &str, + param: Option<&str>, + event_id: Option<&str>, + ) -> Result<(), End> { + self.to_client(gateway_error(code, message, param, event_id)) + .await + } + + async fn revalidate(&self) -> Result<(), End> { + if still_allowed(self.state, &self.p.caller.check, &self.p.endpoint_id) + .await + .is_some() + { + return Ok(()); + } + info!(endpoint_id = %self.p.endpoint_id, "realtime session revoked"); + Err(End::new("revoked", 1008).with_error( + "session_revoked", + "The credential or the deployment was revoked; the session is closed.", + )) + } + + /// Forward one policy-approved client frame, revalidating a `response.create` first. + async fn forward_client(&mut self, text: String, response_create: bool) -> Result<(), End> { + if response_create { + self.revalidate().await?; + if !admit_response(&self.p.caller, &self.p.endpoint_id) { + return self + .refuse( + "quota_exceeded", + "The project's spend quota is exhausted.", + None, + None, + ) + .await; + } + } + self.to_vendor(text).await + } + + async fn release_held(&mut self) -> Result<(), End> { + self.ready = true; + self.awaiting_defaults = None; + self.hold_deadline = None; + while let Some((text, response_create)) = self.held.pop_front() { + self.forward_client(text, response_create).await?; + } + Ok(()) + } + + async fn on_client_text(&mut self, raw: &str) -> Result<(), End> { + self.last_activity = Instant::now(); + match policy::client_event(raw, &self.p.rules) { + ClientOutcome::Invalid(why) => { + self.refuse( + "invalid_event", + &format!("The event could not be read: {why}"), + None, + None, + ) + .await + } + ClientOutcome::Refuse(r) => { + metrics::counter!("waav_realtime_policy_refusals_total", "field" => r.param.clone()) + .increment(1); + self.refuse( + "event_not_allowed", + &r.message, + Some(&r.param), + r.event_id.as_deref(), + ) + .await + } + ClientOutcome::Forward { + text, + response_create, + unknown, + .. + } => { + if unknown { + metrics::counter!("waav_realtime_unknown_client_events_total").increment(1); + } + if !self.ready { + if self.held.len() >= MAX_HELD { + return Err(End::new("upstream_error", 1011).with_error( + "upstream_error", + "The vendor session did not become ready.", + )); + } + self.held.push_back((text.into_owned(), response_create)); + return Ok(()); + } + self.forward_client(text.into_owned(), response_create) + .await + } + } + } + + async fn on_vendor_text(&mut self, raw: &str) -> Result<(), End> { + self.last_activity = Instant::now(); + match policy::vendor_event(raw, &self.p.endpoint_name) { + VendorOutcome::Drop => Ok(()), + VendorOutcome::Invalid => self.to_client(raw.to_string()).await, + VendorOutcome::Forward(text, tap) => { + let text = text.into_owned(); + match tap { + Tap::SessionCreated { vendor_session_id } => { + self.meter.set_vendor_session_id(vendor_session_id); + self.to_client(text).await?; + let event_id = + format!("evt_bud_defaults_{}", uuid::Uuid::new_v4().simple()); + match policy::defaults_update( + self.p.settings.as_ref(), + self.p.endpoint.model.as_deref(), + &event_id, + ) { + Some(update) => { + self.to_vendor(update).await?; + self.awaiting_defaults = Some(event_id); + self.hold_deadline = Some(Instant::now() + self.timings.hold); + Ok(()) + } + None => self.release_held().await, + } + } + Tap::SessionUpdated { .. } => { + self.to_client(text).await?; + if self.awaiting_defaults.is_some() { + self.release_held().await?; + } + Ok(()) + } + Tap::ResponseDone(event) => { + self.to_client(text).await?; + self.meter.response_done(&event); + Ok(()) + } + Tap::TranscriptionCompleted(event) => { + self.to_client(text).await?; + self.meter.transcription_completed(&event); + Ok(()) + } + Tap::Error(event) => { + let about_defaults = self.awaiting_defaults.as_deref().is_some_and(|id| { + event + .get("error") + .and_then(|e| e.get("event_id")) + .and_then(|v| v.as_str()) + == Some(id) + }); + self.to_client(text).await?; + if about_defaults { + // A deployment whose defaults the vendor refuses still serves the + // session on the vendor's defaults; the client has seen the error. + warn!(endpoint_id = %self.p.endpoint_id, "vendor refused the deployment defaults"); + self.release_held().await?; + } + Ok(()) + } + Tap::None => self.to_client(text).await, + } + } + } + } +} + +async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option) { + let _slot = slot; + let timings = state.realtime.timings.clone(); + let session_id = format!("sess_bud_{}", uuid::Uuid::new_v4().simple()); + let vendor = p.endpoint.vendor.clone(); + let attribution = Attribution { + project_id: p + .alias + .as_ref() + .and_then(|a| a.project_id.clone()) + .or_else(|| p.caller.principal.project_id.clone()), + endpoint_id: p.endpoint_id.clone(), + model_id: p.alias.as_ref().and_then(|a| a.model_id.clone()), + api_key_id: p.caller.principal.api_key_id.clone(), + user_id: p.caller.principal.user_id.clone(), + api_key_project_id: p.caller.principal.project_id.clone(), + endpoint_name: p.endpoint_name.clone(), + vendor: vendor.clone(), + model: p.endpoint.model.clone(), + session_type: p + .settings + .as_ref() + .and_then(|s| s.session_type.clone()) + .unwrap_or_else(|| "realtime".into()), + }; + let meter = SessionMeter::start(session_id.clone(), attribution, p.endpoint.pricing.clone()); + metrics::gauge!("waav_realtime_sessions_active", "vendor" => vendor.clone()).increment(1.0); + info!(session_id = %session_id, endpoint_id = %p.endpoint_id, vendor = %vendor, "realtime session opened"); + + debug!(session_id = %session_id, "realtime upgrade complete; connecting upstream"); + let (client_sink, mut client_rx) = socket.split(); + let (client_tx, writer) = spawn_writer(client_sink, CLIENT_QUEUE); + + let upstream_socket = match upstream::connect(&p.upstream, timings.connect).await { + Ok(s) => { + if let Some(pol) = &state.policies { + pol.breakers().record_success(&p.endpoint_id, &p.vkey); + } + s + } + Err(e) => { + warn!(session_id = %session_id, error = %e, "realtime vendor connect failed"); + if let Some(pol) = &state.policies { + let verdict = + crate::core::deployment_policy::classify_message(&e.to_string(), None); + pol.breakers() + .record_failure(&p.endpoint_id, &p.vkey, &verdict); + } + let end = End::new("upstream_error", 1011).with_error("upstream_error", e.to_string()); + finish(meter, end, &client_tx, writer, None, &vendor).await; + return; + } + }; + debug!(session_id = %session_id, "realtime upstream connected"); + let (up_tx, mut up_rx) = upstream_socket.split(); + + let now = Instant::now(); + let limits = p + .settings + .as_ref() + .and_then(|s| s.limits.clone()) + .unwrap_or_default(); + let max_len = limits + .max_session_seconds + .map(Duration::from_secs) + .map_or(timings.max_session, |d| d.min(timings.max_session)); + let idle = limits + .idle_timeout_seconds + .map(Duration::from_secs) + .unwrap_or(timings.default_idle); + let max_at = now + max_len; + let warn_at = max_len.checked_sub(timings.warn_before).map(|d| now + d); + let bills_duration = crate::core::realtime_cost::bills_duration(p.endpoint.pricing.as_ref()); + + let mut relay = Relay { + state: &state, + p: &p, + meter, + client_tx: client_tx.clone(), + up_tx, + timings: timings.clone(), + held: VecDeque::new(), + ready: false, + awaiting_defaults: None, + // The vendor must say `session.created` within the hold window as well. + hold_deadline: Some(now + timings.connect), + last_activity: now, + client_missed: 0, + vendor_missed: 0, + }; + let mut ping = tokio::time::interval_at(now + timings.ping, timings.ping); + let mut revalidate = tokio::time::interval_at(now + timings.revalidate, timings.revalidate); + let mut segment = tokio::time::interval_at(now + timings.segment, timings.segment); + let mut last_segment = now; + let mut warned = warn_at.is_none(); + + let end: End = loop { + let idle_at = relay.last_activity + idle; + let hold_at = relay.hold_deadline; + let step: Result<(), End> = tokio::select! { + _ = state.shutdown.cancelled() => Err(End::new("drain", 1012) + .with_error("server_shutdown", "The server is restarting; reconnect.")), + msg = client_rx.next() => match msg { + None | Some(Err(_)) => Err(End::new("client_close", 1006)), + Some(Ok(Message::Close(frame))) => Err(End::new("client_close", frame.map_or(1005, |f| f.code))), + Some(Ok(Message::Pong(_))) => { relay.client_missed = 0; Ok(()) } + Some(Ok(Message::Ping(_))) => Ok(()), + Some(Ok(Message::Binary(_))) => relay.refuse( + "invalid_event", "Binary frames are not part of the Realtime protocol; send JSON events.", None, None, + ).await, + Some(Ok(Message::Text(t))) => relay.on_client_text(t.as_str()).await, + }, + msg = up_rx.next() => match msg { + None | Some(Err(_)) => Err(End::new("upstream_error", 1011) + .with_error("upstream_error", "The connection to the vendor was lost.")), + Some(Ok(UpMessage::Close(frame))) => { + let detail = frame.map(|f| format!(" (code {}: {})", u16::from(f.code), f.reason)).unwrap_or_default(); + Err(End::new("upstream_error", 1011) + .with_error("upstream_error", format!("The vendor closed the session{detail}."))) + } + Some(Ok(UpMessage::Pong(_))) => { relay.vendor_missed = 0; Ok(()) } + Some(Ok(UpMessage::Ping(_))) => { let _ = relay.up_tx.flush().await; Ok(()) } + Some(Ok(UpMessage::Text(t))) => relay.on_vendor_text(t.as_str()).await, + Some(Ok(_)) => Ok(()), + }, + _ = ping.tick() => { + if relay.client_missed >= timings.max_missed_pongs { + Err(End::new("client_timeout", 1011) + .with_error("client_timeout", "The client stopped answering pings.")) + } else if relay.vendor_missed >= timings.max_missed_pongs { + Err(End::new("upstream_error", 1011) + .with_error("upstream_error", "The vendor stopped answering pings.")) + } else { + relay.client_missed += 1; + relay.vendor_missed += 1; + let _ = relay.client_tx.try_send(Outbound::Frame(Message::Ping(Vec::new().into()))); + match tokio::time::timeout(timings.upstream_send, relay.up_tx.send(UpMessage::Ping(Vec::new().into()))).await { + Ok(Ok(())) => Ok(()), + _ => Err(End::new("upstream_error", 1011) + .with_error("upstream_error", "The connection to the vendor failed.")), + } + } + } + _ = revalidate.tick() => relay.revalidate().await, + _ = tokio::time::sleep_until(idle_at) => Err(End::new("idle", 1000) + .with_error("session_expired", format!("The session was idle for {} s.", idle.as_secs()))), + _ = tokio::time::sleep_until(warn_at.unwrap_or(max_at)), if !warned => { + warned = true; + relay.refuse("session_expiring", + &format!("The session reaches its maximum length in {} s.", timings.warn_before.as_secs()), + None, None).await + } + _ = tokio::time::sleep_until(max_at) => Err(End::new("max_duration", 1000) + .with_error("session_expired", format!("The session reached its maximum length of {} s.", max_len.as_secs()))), + _ = tokio::time::sleep_until(hold_at.unwrap_or(max_at)), if hold_at.is_some() => Err(End::new("upstream_error", 1011) + .with_error("upstream_error", "The vendor did not start or configure the session in time.")), + _ = segment.tick(), if bills_duration => { + relay.meter.duration_segment(timings.segment.as_secs_f64()); + last_segment = Instant::now(); + Ok(()) + } + }; + if let Err(end) = step { + break end; + } + }; + + if bills_duration { + // The final partial segment: a 150 s session bills 60 + 60 + 30 (TC-XL-07). + relay + .meter + .duration_segment(last_segment.elapsed().as_secs_f64()); + } + let Relay { + meter, mut up_tx, .. + } = relay; + let _ = tokio::time::timeout( + Duration::from_secs(2), + up_tx.send(UpMessage::Close(Some( + tokio_tungstenite::tungstenite::protocol::CloseFrame { + code: tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode::Normal, + reason: "".into(), + }, + ))), + ) + .await; + finish(meter, end, &client_tx, writer, Some(&session_id), &vendor).await; + drop(p); +} + +/// The common teardown: the error event and close, the session record, the metrics. +async fn finish( + meter: SessionMeter, + end: End, + client_tx: &mpsc::Sender, + writer: tokio::task::JoinHandle<()>, + session_id: Option<&str>, + vendor: &str, +) { + let error = end + .error + .as_ref() + .map(|(code, message)| gateway_error(code, message, None, None)); + // Behind any backlog, so every frame queued before the close still arrives in order. + let _ = tokio::time::timeout( + Duration::from_secs(30), + client_tx.send(Outbound::Close { + error, + code: end.close_code.clamp(1000, 4999), + reason: end.reason.to_string(), + }), + ) + .await; + let abort = writer.abort_handle(); + if tokio::time::timeout(Duration::from_secs(30), writer) + .await + .is_err() + { + abort.abort(); + } + info!( + session_id = session_id.unwrap_or(meter.session_id()), + end_reason = end.reason, + close_code = end.close_code, + turns = meter.turns(), + "realtime session closed" + ); + metrics::counter!("waav_realtime_sessions_total", "vendor" => vendor.to_string(), "end_reason" => end.reason) + .increment(1); + metrics::gauge!("waav_realtime_sessions_active", "vendor" => vendor.to_string()).decrement(1.0); + meter.finish(end.reason, end.close_code); + debug!("realtime session torn down"); +} diff --git a/gateway/src/handlers/openai_realtime/upstream.rs b/gateway/src/handlers/openai_realtime/upstream.rs new file mode 100644 index 00000000..ff52a84d --- /dev/null +++ b/gateway/src/handlers/openai_realtime/upstream.rs @@ -0,0 +1,380 @@ +//! The vendor leg of a relayed session (FRD-023 §5.3, FR-CRED-2/3, S-7). +//! +//! The URL, the auth header and the model come from the deployment's `voice_table` entry and from +//! nowhere else: not the client, not WaaV's environment (D-5). A client-controlled host on this +//! path is exactly the bug that made Helicone remove realtime (S-7), so the only host that is not +//! a vendor constant — an `api_base` — is SSRF-validated before it is dialled. + +use std::time::Duration; + +use bud_auth::VoiceEndpoint; +use tokio::net::TcpStream; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream}; + +/// The relay-capable vendors (D-2). Everything else needs the translate engine (RT7). +pub const RELAY_VENDORS: &[&str] = &["openai", "azure_openai"]; + +const OPENAI_DEFAULT_BASE: &str = "https://api.openai.com/v1"; + +pub type UpstreamSocket = WebSocketStream>; + +/// Why the vendor leg could not be built or opened. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum UpstreamError { + /// The deployment's vendor is not served by the relay (yet). + UnsupportedVendor(String), + /// The deployment carries no usable credential. + MissingCredential, + /// The deployment names no vendor model. + MissingModel, + /// An Azure deployment without its resource endpoint. + MissingApiBase, + /// The `api_base` is not a URL, or failed SSRF validation. + InvalidApiBase(String), + /// Connecting or upgrading failed. + Connect(String), + /// No answer within the connect deadline. + Timeout, +} + +impl std::fmt::Display for UpstreamError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::UnsupportedVendor(v) => write!( + f, + "realtime sessions are not yet served for vendor '{v}' (relay vendors: openai, azure_openai)" + ), + Self::MissingCredential => write!(f, "the deployment has no usable vendor credential"), + Self::MissingModel => write!(f, "the deployment names no vendor model"), + Self::MissingApiBase => write!( + f, + "an Azure OpenAI deployment needs its resource endpoint (api_base)" + ), + Self::InvalidApiBase(why) => write!(f, "the deployment's api_base was refused: {why}"), + Self::Connect(why) => write!(f, "could not connect to the vendor: {why}"), + Self::Timeout => write!(f, "the vendor did not answer within the connect deadline"), + } + } +} + +/// A built vendor request. `Debug` never prints header values: one of them is the credential. +#[derive(Clone)] +pub struct UpstreamRequest { + pub url: String, + pub headers: Vec<(&'static str, String)>, + /// The host an SSRF check must clear, when the URL is not a vendor constant. + pub needs_ssrf_check: bool, +} + +impl std::fmt::Debug for UpstreamRequest { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("UpstreamRequest") + .field("url", &self.url) + .field( + "headers", + &self.headers.iter().map(|(k, _)| *k).collect::>(), + ) + .finish() + } +} + +/// `https://…` → `wss://…`, `http://…` → `ws://…`; `ws(s)` kept. Trailing slashes trimmed. +fn to_ws_base(base: &str) -> Result { + let base = base.trim().trim_end_matches('/'); + let converted = if let Some(rest) = base.strip_prefix("https://") { + format!("wss://{rest}") + } else if let Some(rest) = base.strip_prefix("http://") { + format!("ws://{rest}") + } else if base.starts_with("wss://") || base.starts_with("ws://") { + base.to_string() + } else { + return Err(UpstreamError::InvalidApiBase( + "api_base must be an http(s) or ws(s) URL".into(), + )); + }; + url::Url::parse(&converted).map_err(|e| UpstreamError::InvalidApiBase(e.to_string()))?; + Ok(converted) +} + +fn query_value(v: &str) -> String { + url::form_urlencoded::byte_serialize(v.as_bytes()).collect() +} + +/// Build the vendor request for a deployment (FRD §5.3 table). +/// +/// A transcription deployment connects with `?intent=transcription`; its model is set by the +/// defaults `session.update` (`audio.input.transcription.model`). +pub fn build( + endpoint: &VoiceEndpoint, + transcription: bool, +) -> Result { + let vendor = endpoint.vendor.trim().to_ascii_lowercase(); + let credential = endpoint + .credential + .as_deref() + .map(str::trim) + .filter(|c| !c.is_empty()) + .ok_or(UpstreamError::MissingCredential)? + .to_string(); + let model = endpoint + .model + .as_deref() + .map(str::trim) + .filter(|m| !m.is_empty()); + let query = |model: Option<&str>| -> Result { + if transcription { + Ok("intent=transcription".to_string()) + } else { + Ok(format!( + "model={}", + query_value(model.ok_or(UpstreamError::MissingModel)?) + )) + } + }; + + match vendor.as_str() { + "openai" => { + let api_base = endpoint + .api_base + .as_deref() + .filter(|b| !b.trim().is_empty()); + let base = to_ws_base(api_base.unwrap_or(OPENAI_DEFAULT_BASE))?; + Ok(UpstreamRequest { + url: format!("{base}/realtime?{}", query(model)?), + headers: vec![("authorization", format!("Bearer {credential}"))], + needs_ssrf_check: api_base.is_some(), + }) + } + "azure_openai" => { + let api_base = endpoint + .api_base + .as_deref() + .filter(|b| !b.trim().is_empty()) + .ok_or(UpstreamError::MissingApiBase)?; + let parsed = url::Url::parse(api_base.trim()) + .map_err(|e| UpstreamError::InvalidApiBase(e.to_string()))?; + let host = parsed + .host_str() + .ok_or_else(|| UpstreamError::InvalidApiBase("api_base has no host".into()))?; + let authority = match parsed.port() { + Some(port) => format!("{host}:{port}"), + None => host.to_string(), + }; + let scheme = if parsed.scheme() == "http" { + "ws" + } else { + "wss" + }; + // GA: no `api-version` — the v1 URL answers 401 with one (FRD §5.3). The Azure + // deployment name is the `model`. + Ok(UpstreamRequest { + url: format!( + "{scheme}://{authority}/openai/v1/realtime?{}", + query(model)? + ), + headers: vec![("api-key", credential)], + needs_ssrf_check: true, + }) + } + other => Err(UpstreamError::UnsupportedVendor(other.to_string())), + } +} + +/// SSRF-validate the request's host (blocking DNS, so off the async workers). +pub async fn validate(req: &UpstreamRequest) -> Result<(), UpstreamError> { + if !req.needs_ssrf_check { + return Ok(()); + } + let url = req.url.clone(); + tokio::task::spawn_blocking(move || { + crate::core::net::validate_url_for_ssrf(&url, &["ws", "wss"]) + }) + .await + .map_err(|e| UpstreamError::InvalidApiBase(format!("validation task failed: {e}")))? + .map_err(UpstreamError::InvalidApiBase) +} + +/// Open the vendor socket within `deadline`. No extensions are offered (no permessage-deflate: +/// frames are parsed, S-8). +pub async fn connect( + req: &UpstreamRequest, + deadline: Duration, +) -> Result { + let mut request = req + .url + .as_str() + .into_client_request() + .map_err(|e| UpstreamError::Connect(e.to_string()))?; + for (name, value) in &req.headers { + let value = value + .parse() + .map_err(|_| UpstreamError::Connect(format!("header {name} is not a valid value")))?; + request.headers_mut().insert(*name, value); + } + let config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default() + .max_message_size(Some(super::handshake::MAX_MESSAGE_BYTES)) + .max_frame_size(Some(super::handshake::MAX_MESSAGE_BYTES)); + match tokio::time::timeout( + deadline, + tokio_tungstenite::connect_async_with_config(request, Some(config), true), + ) + .await + { + Ok(Ok((socket, _response))) => Ok(socket), + // The error text can carry the URL; never the headers. + Ok(Err(e)) => Err(UpstreamError::Connect(e.to_string())), + Err(_) => Err(UpstreamError::Timeout), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn endpoint(vendor: &str, api_base: Option<&str>, model: Option<&str>) -> VoiceEndpoint { + let mut entry = serde_json::json!({"vendor": vendor, "endpoints": ["realtime_session"]}); + if let Some(b) = api_base { + entry["api_base"] = b.into(); + } + if let Some(m) = model { + entry["model"] = m.into(); + } + let blob = serde_json::json!({ "ep": entry }).to_string(); + let mut ep = bud_auth::credentials::parse_voice_blob( + &blob, + &bud_auth::CredentialDecryptor::disabled(), + ) + .unwrap() + .remove("ep") + .unwrap(); + ep.credential = Some("sk-vendor".into()); + ep + } + + /// TC-UP-01 — OpenAI: GA URL, Bearer, no beta header. + #[test] + fn tc_up_01_openai_url_and_auth() { + let req = build(&endpoint("openai", None, Some("gpt-realtime-2.1")), false).unwrap(); + assert_eq!( + req.url, + "wss://api.openai.com/v1/realtime?model=gpt-realtime-2.1" + ); + assert_eq!( + req.headers, + vec![("authorization", "Bearer sk-vendor".to_string())] + ); + assert!( + !req.needs_ssrf_check, + "the vendor constant needs no SSRF check" + ); + assert!( + !req.headers + .iter() + .any(|(k, _)| k.eq_ignore_ascii_case("openai-beta")) + ); + } + + /// TC-UP-02 🔒 — Azure GA URL: `/openai/v1/realtime`, `api-key`, and NO `api-version`. + #[test] + fn tc_up_02_azure_ga_url_has_no_api_version() { + let req = build( + &endpoint( + "azure_openai", + Some("https://r.openai.azure.com/"), + Some("rt-prod"), + ), + false, + ) + .unwrap(); + assert_eq!( + req.url, + "wss://r.openai.azure.com/openai/v1/realtime?model=rt-prod" + ); + assert!(!req.url.contains("api-version")); + assert_eq!(req.headers, vec![("api-key", "sk-vendor".to_string())]); + assert!(req.needs_ssrf_check); + } + + /// TC-UP-03 — an `api_base` has its scheme converted and keeps its path. + #[test] + fn tc_up_03_api_base_scheme_conversion() { + let req = build( + &endpoint( + "openai", + Some("https://proxy.example/v1"), + Some("gpt-realtime-2.1"), + ), + false, + ) + .unwrap(); + assert_eq!( + req.url, + "wss://proxy.example/v1/realtime?model=gpt-realtime-2.1" + ); + assert!(req.needs_ssrf_check); + } + + /// TC-UP-04 🔒 — a private `api_base` is refused before any connect. + #[tokio::test] + async fn tc_up_04_ssrf_refuses_private_hosts() { + for base in [ + "http://10.0.0.5", + "http://169.254.169.254/v1", + "http://localhost:9000", + ] { + let req = build(&endpoint("openai", Some(base), Some("m")), false).unwrap(); + let err = validate(&req).await.expect_err(base); + assert!( + matches!(err, UpstreamError::InvalidApiBase(_)), + "{base}: {err:?}" + ); + } + } + + /// TC-UP-05 — the vendor model comes from `voice_table`, never the client. + #[test] + fn tc_up_05_model_is_the_deployments() { + let req = build(&endpoint("openai", None, Some("gpt-realtime-2.1")), false).unwrap(); + assert!(req.url.ends_with("model=gpt-realtime-2.1")); + } + + #[test] + fn a_transcription_deployment_connects_with_the_transcription_intent() { + let req = build(&endpoint("openai", None, Some("gpt-4o-transcribe")), true).unwrap(); + assert_eq!( + req.url, + "wss://api.openai.com/v1/realtime?intent=transcription" + ); + } + + #[test] + fn missing_pieces_are_named() { + let mut ep = endpoint("openai", None, Some("m")); + ep.credential = None; + assert_eq!( + build(&ep, false).unwrap_err(), + UpstreamError::MissingCredential + ); + assert_eq!( + build(&endpoint("openai", None, None), false).unwrap_err(), + UpstreamError::MissingModel + ); + assert_eq!( + build(&endpoint("azure_openai", None, Some("d")), false).unwrap_err(), + UpstreamError::MissingApiBase + ); + assert!(matches!( + build(&endpoint("gemini", None, Some("m")), false).unwrap_err(), + UpstreamError::UnsupportedVendor(_) + )); + } + + /// TC-SEC-07 (unit half): the vendor credential cannot reach a log through `Debug`. + #[test] + fn the_request_debug_never_prints_the_credential() { + let req = build(&endpoint("openai", None, Some("m")), false).unwrap(); + let printed = format!("{req:?}"); + assert!(!printed.contains("sk-vendor"), "{printed}"); + } +} diff --git a/gateway/src/main.rs b/gateway/src/main.rs index b353191f..85b522c9 100644 --- a/gateway/src/main.rs +++ b/gateway/src/main.rs @@ -296,6 +296,14 @@ async fn main() -> anyhow::Result<()> { connection_limit_middleware, )); + // FRD-023: `/v1/realtime` (OpenAI Realtime GA on Bud deployments) and the client-secret mint. + // No auth_middleware: the handlers authenticate themselves (subprotocol and `api-key` + // credentials, `ek_bud_` secrets, `?token=` refused). The connection limit still applies, so + // every session holds a slot (FRD-022). + let openai_realtime_routes = routes::openai_realtime::create_openai_realtime_router().layer( + middleware::from_fn_with_state(app_state.clone(), connection_limit_middleware), + ); + // Live latency-profile debug surface (`WAAV_DEBUG_PROFILE=1`): JSON snapshot // + per-turn SSE. Auth-gated exactly like the protected API; without the env // flag the routes are not mounted at all (double lock, never public). @@ -475,6 +483,7 @@ async fn main() -> anyhow::Result<()> { .merge(protected_routes) .merge(ws_routes) .merge(realtime_routes) + .merge(openai_realtime_routes) .merge(debug_profile_routes) .with_state(app_state.clone()) .layer(tower::util::option_layer(governor_layer)); diff --git a/gateway/src/observability/voice_attrs.rs b/gateway/src/observability/voice_attrs.rs index df7de66c..25ab8a6e 100644 --- a/gateway/src/observability/voice_attrs.rs +++ b/gateway/src/observability/voice_attrs.rs @@ -97,6 +97,49 @@ pub mod resilience { pub const RATE_LIMIT_OUTCOME: &str = "bud.rate_limit.outcome"; } +/// A realtime (speech-to-speech) billed record (FRD-023 §5.10, CONTRACTS C2). +/// +/// Realtime is the one place on the audio plane billed in TOKENS — at up to eight per-modality +/// rates — so these are the only token attributes in the vocabulary. +pub mod realtime { + /// `response` | `input_transcription` | `duration_segment`. + pub const COMPONENT: &str = "bud.voice.rt.component"; + /// The vendor serving the session (`voice_table.vendor`). + pub const VENDOR: &str = "bud.voice.rt.vendor"; + /// The vendor's model (`voice_table.model`; the Azure deployment name for Azure). + pub const MODEL: &str = "bud.voice.rt.model"; + pub const RESPONSE_ID: &str = "bud.voice.rt.response_id"; + /// `completed` | `cancelled` | `incomplete` | `failed`. + pub const RESPONSE_STATUS: &str = "bud.voice.rt.response_status"; + pub const INPUT_TEXT_TOKENS: &str = "bud.voice.rt.input_text_tokens"; + pub const INPUT_AUDIO_TOKENS: &str = "bud.voice.rt.input_audio_tokens"; + pub const INPUT_IMAGE_TOKENS: &str = "bud.voice.rt.input_image_tokens"; + /// Cached tokens are a SUBSET of their input class, not additional to it. + pub const CACHED_TEXT_TOKENS: &str = "bud.voice.rt.cached_text_tokens"; + pub const CACHED_AUDIO_TOKENS: &str = "bud.voice.rt.cached_audio_tokens"; + pub const CACHED_IMAGE_TOKENS: &str = "bud.voice.rt.cached_image_tokens"; + pub const OUTPUT_TEXT_TOKENS: &str = "bud.voice.rt.output_text_tokens"; + pub const OUTPUT_AUDIO_TOKENS: &str = "bud.voice.rt.output_audio_tokens"; + /// Seconds billed: an input transcription's audio, or a duration segment. + pub const BILLED_SECONDS: &str = "bud.voice.billed_seconds"; + /// Components present with no rate — named, never priced at zero. + pub const UNPRICED_COMPONENTS: &str = "bud.voice.unpriced_components"; + /// The vendor's own session id (`session.created.session.id`). + pub const VENDOR_SESSION_ID: &str = "bud.voice.vendor_session_id"; + /// `realtime` | `transcription` (the session span). + pub const SESSION_TYPE: &str = "bud.voice.rt.session_type"; +} + +/// Attributes of the `voice.session` span, one per realtime session (FRD-023 §5.10). +pub mod session { + pub const DURATION_MS: &str = "bud.voice.session.duration_ms"; + pub const TURNS: &str = "bud.voice.session.turns"; + /// `client_close` | `idle` | `max_duration` | `revoked` | `drain` | `upstream_error` | + /// `rate_limited` | `client_too_slow` | `client_timeout`. + pub const END_REASON: &str = "bud.voice.session.end_reason"; + pub const CLOSE_CODE: &str = "bud.voice.session.close_code"; +} + /// Attributes carried by a CLIENT span covering one leg of a turn. /// /// Per-leg rather than a single `provider`/`duration` pair, because a turn routinely spans two @@ -176,6 +219,72 @@ pub const ALL: &[&str] = &[ resilience::FALLBACK_FROM, resilience::RETRY_COUNT, resilience::RATE_LIMIT_OUTCOME, + realtime::COMPONENT, + realtime::VENDOR, + realtime::MODEL, + realtime::RESPONSE_ID, + realtime::RESPONSE_STATUS, + realtime::INPUT_TEXT_TOKENS, + realtime::INPUT_AUDIO_TOKENS, + realtime::INPUT_IMAGE_TOKENS, + realtime::CACHED_TEXT_TOKENS, + realtime::CACHED_AUDIO_TOKENS, + realtime::CACHED_IMAGE_TOKENS, + realtime::OUTPUT_TEXT_TOKENS, + realtime::OUTPUT_AUDIO_TOKENS, + realtime::BILLED_SECONDS, + realtime::UNPRICED_COMPONENTS, + realtime::VENDOR_SESSION_ID, + realtime::SESSION_TYPE, + session::DURATION_MS, + session::TURNS, + session::END_REASON, + session::CLOSE_CODE, +]; + +/// Attributes only the `voice.session` span carries; every other attribute in [`ALL`] is a +/// `voice.turn` attribute and declared by [`voice_turn_span!`]. +pub const SESSION_ONLY: &[&str] = &[ + realtime::SESSION_TYPE, + session::DURATION_MS, + session::TURNS, + session::END_REASON, + session::CLOSE_CODE, +]; + +/// The attributes a `voice.session` span declares (FRD-023): attribution, the realtime shape, the +/// session's own fields and the totals. Kept beside [`ALL`] so the session macro and its test +/// read one list. +pub const SESSION: &[&str] = &[ + turn::PROJECT_ID, + turn::ENDPOINT_ID, + turn::MODEL_ID, + turn::API_KEY_ID, + turn::USER_ID, + turn::API_KEY_PROJECT_ID, + turn::ENDPOINT_NAME, + turn::CAPABILITY, + turn::TRANSPORT, + turn::SESSION_ID, + turn::COST, + turn::PRICING_UNIT, + realtime::VENDOR, + realtime::MODEL, + realtime::SESSION_TYPE, + realtime::VENDOR_SESSION_ID, + realtime::INPUT_TEXT_TOKENS, + realtime::INPUT_AUDIO_TOKENS, + realtime::INPUT_IMAGE_TOKENS, + realtime::CACHED_TEXT_TOKENS, + realtime::CACHED_AUDIO_TOKENS, + realtime::CACHED_IMAGE_TOKENS, + realtime::OUTPUT_TEXT_TOKENS, + realtime::OUTPUT_AUDIO_TOKENS, + realtime::BILLED_SECONDS, + session::DURATION_MS, + session::TURNS, + session::END_REASON, + session::CLOSE_CODE, ]; /// Open a `voice.turn` span that declares EVERY attribute in [`ALL`] up front. @@ -198,9 +307,100 @@ pub const ALL: &[&str] = &[ /// extra field can be written in any form tracing accepts, `{ CONST } = value` included. #[macro_export] macro_rules! voice_turn_span { + // FRD-023 D-19: a realtime billed record is the ROOT of its own trace. + (parent: $parent:expr, capability = $capability:expr, transport = $transport:expr $(, $($extra:tt)*)?) => { + $crate::__voice_turn_span_fields!((parent: $parent,) $capability, $transport $(, $($extra)*)?) + }; (capability = $capability:expr, transport = $transport:expr $(, $($extra:tt)*)?) => { + $crate::__voice_turn_span_fields!(() $capability, $transport $(, $($extra)*)?) + }; +} + +/// The field list behind [`voice_turn_span!`], in one place for both of its forms. +#[doc(hidden)] +#[macro_export] +macro_rules! __voice_turn_span_fields { + (($($parent:tt)*) $capability:expr, $transport:expr $(, $($extra:tt)*)?) => { ::tracing::info_span!( + $($parent)* "voice.turn", + { $crate::observability::voice_attrs::turn::CAPABILITY } = $capability, + { $crate::observability::voice_attrs::turn::TRANSPORT } = $transport, + { $crate::observability::voice_attrs::turn::PROJECT_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::ENDPOINT_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::MODEL_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::API_KEY_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::USER_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::SESSION_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::TURN_INDEX } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::CHARACTERS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::AUDIO_SECONDS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::COST } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::RESPONSE_LATENCY_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::BARGE_IN } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::TURN_DETECTOR } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::LANGUAGE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::TRANSCRIPT } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::SYNTHESIS_INPUT } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::STT_VENDOR } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::STT_DURATION_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::STT_TTFB_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::TTS_VENDOR } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::TTS_DURATION_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::TTS_TTFB_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::LLM_MODEL } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::LLM_DURATION_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::ENDPOINT_NAME } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::API_KEY_PROJECT_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::PRICING_UNIT } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::ERROR_TYPE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::VENDOR_STATUS_CODE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::OUTPUT_AUDIO_SECONDS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::DETECTED_LANGUAGE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::AUDIO_FORMAT } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::SAMPLE_RATE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::INPUT_AUDIO_BYTES } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::VENDOR_REQUEST_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::STT_MODEL } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::TTS_MODEL } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::STT_CONFIDENCE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::TTS_VOICE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::leg::STT_NOISE_SUPPRESSION } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::resilience::SERVED_ENDPOINT_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::resilience::FALLBACK_FROM } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::resilience::RETRY_COUNT } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::resilience::RATE_LIMIT_OUTCOME } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::COMPONENT } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::VENDOR } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::MODEL } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::RESPONSE_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::RESPONSE_STATUS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::INPUT_TEXT_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::INPUT_AUDIO_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::INPUT_IMAGE_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::CACHED_TEXT_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::CACHED_AUDIO_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::CACHED_IMAGE_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::OUTPUT_TEXT_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::OUTPUT_AUDIO_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::BILLED_SECONDS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::UNPRICED_COMPONENTS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::VENDOR_SESSION_ID } = ::tracing::field::Empty, + otel.status_code = ::tracing::field::Empty, + otel.status_message = ::tracing::field::Empty + $(, $($extra)*)? + ) + }; +} + +/// Open a `voice.session` span (FRD-023 §5.10): the ROOT of its own trace, declaring every +/// attribute in [`SESSION`] up front (the declare-before-record rule of [`voice_turn_span!`]). +#[macro_export] +macro_rules! voice_session_span { + (capability = $capability:expr, transport = $transport:expr) => { + ::tracing::info_span!( + parent: None, + "voice.session", { $crate::observability::voice_attrs::turn::CAPABILITY } = $capability, { $crate::observability::voice_attrs::turn::TRANSPORT } = $transport, { $crate::observability::voice_attrs::turn::PROJECT_ID } = ::tracing::field::Empty, @@ -208,48 +408,30 @@ macro_rules! voice_turn_span { { $crate::observability::voice_attrs::turn::MODEL_ID } = ::tracing::field::Empty, { $crate::observability::voice_attrs::turn::API_KEY_ID } = ::tracing::field::Empty, { $crate::observability::voice_attrs::turn::USER_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::API_KEY_PROJECT_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::turn::ENDPOINT_NAME } = ::tracing::field::Empty, { $crate::observability::voice_attrs::turn::SESSION_ID } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::TURN_INDEX } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::CHARACTERS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::AUDIO_SECONDS } = ::tracing::field::Empty, { $crate::observability::voice_attrs::turn::COST } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::RESPONSE_LATENCY_MS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::BARGE_IN } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::TURN_DETECTOR } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::LANGUAGE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::TRANSCRIPT } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::SYNTHESIS_INPUT } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::STT_VENDOR } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::STT_DURATION_MS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::STT_TTFB_MS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::TTS_VENDOR } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::TTS_DURATION_MS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::TTS_TTFB_MS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::LLM_MODEL } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::LLM_DURATION_MS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::ENDPOINT_NAME } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::API_KEY_PROJECT_ID } = ::tracing::field::Empty, { $crate::observability::voice_attrs::turn::PRICING_UNIT } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::ERROR_TYPE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::VENDOR_STATUS_CODE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::OUTPUT_AUDIO_SECONDS } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::DETECTED_LANGUAGE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::AUDIO_FORMAT } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::SAMPLE_RATE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::INPUT_AUDIO_BYTES } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::turn::VENDOR_REQUEST_ID } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::STT_MODEL } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::TTS_MODEL } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::STT_CONFIDENCE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::TTS_VOICE } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::leg::STT_NOISE_SUPPRESSION } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::resilience::SERVED_ENDPOINT_ID } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::resilience::FALLBACK_FROM } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::resilience::RETRY_COUNT } = ::tracing::field::Empty, - { $crate::observability::voice_attrs::resilience::RATE_LIMIT_OUTCOME } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::VENDOR } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::MODEL } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::SESSION_TYPE } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::VENDOR_SESSION_ID } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::INPUT_TEXT_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::INPUT_AUDIO_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::INPUT_IMAGE_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::CACHED_TEXT_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::CACHED_AUDIO_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::CACHED_IMAGE_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::OUTPUT_TEXT_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::OUTPUT_AUDIO_TOKENS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::realtime::BILLED_SECONDS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::session::DURATION_MS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::session::TURNS } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::session::END_REASON } = ::tracing::field::Empty, + { $crate::observability::voice_attrs::session::CLOSE_CODE } = ::tracing::field::Empty, otel.status_code = ::tracing::field::Empty, otel.status_message = ::tracing::field::Empty - $(, $($extra)*)? ) }; } @@ -358,12 +540,15 @@ mod tests { tracing::subscriber::with_default(capture::Sub(seen.clone()), || { let span = crate::voice_turn_span!(capability = "conversation", transport = "websocket"); - for name in ALL { + for name in ALL.iter().filter(|a| !SESSION_ONLY.contains(a)) { span.record(*name, "x"); } }); let seen = seen.lock().unwrap(); - let missing: Vec<_> = ALL.iter().filter(|a| !seen.contains(**a)).collect(); + let missing: Vec<_> = ALL + .iter() + .filter(|a| !SESSION_ONLY.contains(a) && !seen.contains(**a)) + .collect(); assert!( missing.is_empty(), "voice_turn_span! does not declare {missing:?}; recording them is a silent no-op, \ @@ -429,13 +614,66 @@ mod tests { #[test] fn the_billing_dimensions_are_the_vendors_own_units() { - // A token count on a voice turn would be all zeroes: TTS bills per character, STT per - // second of audio. + // A token count on an HTTP voice turn would be all zeroes: TTS bills per character, STT + // per second of audio. Realtime (FRD-023) is the one capability billed in tokens, so its + // `bud.voice.rt.*_tokens` are the only token attributes. assert!(ALL.contains(&turn::CHARACTERS)); assert!(ALL.contains(&turn::AUDIO_SECONDS)); + let tokens: Vec<_> = ALL.iter().filter(|a| a.contains("token")).collect(); + assert_eq!(tokens.len(), 8, "{tokens:?}"); + assert!( + tokens.iter().all(|a| a.starts_with("bud.voice.rt.")), + "token counts belong to realtime only: {tokens:?}" + ); + } + + /// TC-MET-10 🔒: every realtime attribute is declared by BOTH span macros that carry it, so a + /// record on it is never a silent no-op. + #[test] + fn tc_met_10_the_realtime_turn_root_declares_every_attribute_in_all() { + let seen = std::sync::Arc::new(std::sync::Mutex::new(std::collections::HashSet::new())); + tracing::subscriber::with_default(capture::Sub(seen.clone()), || { + let span = crate::voice_turn_span!( + parent: None, + capability = "realtime_session", + transport = "websocket" + ); + for name in ALL.iter().filter(|a| !SESSION_ONLY.contains(a)) { + span.record(*name, "x"); + } + }); + let seen = seen.lock().unwrap(); + let missing: Vec<_> = ALL + .iter() + .filter(|a| !SESSION_ONLY.contains(a) && !seen.contains(**a)) + .collect(); + assert!( + missing.is_empty(), + "the root form does not declare {missing:?}" + ); + } + + #[test] + fn tc_met_10_the_session_span_declares_every_session_attribute() { + let seen = std::sync::Arc::new(std::sync::Mutex::new(std::collections::HashSet::new())); + tracing::subscriber::with_default(capture::Sub(seen.clone()), || { + let span = crate::voice_session_span!( + capability = "realtime_session", + transport = "websocket" + ); + for name in SESSION { + span.record(*name, "x"); + } + }); + let seen = seen.lock().unwrap(); + let missing: Vec<_> = SESSION.iter().filter(|a| !seen.contains(**a)).collect(); + assert!( + missing.is_empty(), + "voice_session_span! does not declare {missing:?}" + ); assert!( - !ALL.iter().any(|a| a.contains("token")), - "token counts are meaningless for a voice turn" + SESSION.iter().all(|a| ALL.contains(a)), + "a session attribute is missing from ALL" ); } } diff --git a/gateway/src/routes/mod.rs b/gateway/src/routes/mod.rs index 48033b3b..f221d91b 100644 --- a/gateway/src/routes/mod.rs +++ b/gateway/src/routes/mod.rs @@ -1,4 +1,5 @@ pub mod api; +pub mod openai_realtime; pub mod realtime; pub mod webhooks; pub mod ws; diff --git a/gateway/src/routes/openai_realtime.rs b/gateway/src/routes/openai_realtime.rs new file mode 100644 index 00000000..5635846d --- /dev/null +++ b/gateway/src/routes/openai_realtime.rs @@ -0,0 +1,32 @@ +//! `/v1/realtime` and `/v1/realtime/client_secrets` (FRD-023). +//! +//! Mounted WITHOUT `auth_middleware`: these handlers authenticate themselves, because they need +//! credential sources the middleware does not read (the `openai-insecure-api-key.` subprotocol and +//! the `api-key` header), refuse one it does accept (`?token=`), and validate `ek_bud_` client +//! secrets. `main.rs` still layers `connection_limit_middleware` (FRD-022) outside them, so every +//! session holds a connection slot. + +use std::sync::Arc; + +use axum::Router; +use axum::routing::{get, post}; + +use crate::handlers::openai_realtime::{ + CLIENT_SECRETS_PATH, OPENAI_REALTIME_PATH, client_secrets_handler, realtime_ws_handler, +}; +use crate::state::AppState; + +pub fn create_openai_realtime_router() -> Router> { + Router::new() + .route(OPENAI_REALTIME_PATH, get(realtime_ws_handler)) + .route(CLIENT_SECRETS_PATH, post(client_secrets_handler)) +} + +#[cfg(test)] +mod tests { + #[test] + fn the_paths_are_the_openai_ones() { + assert_eq!(super::OPENAI_REALTIME_PATH, "/v1/realtime"); + assert_eq!(super::CLIENT_SECRETS_PATH, "/v1/realtime/client_secrets"); + } +} diff --git a/gateway/src/routes/realtime.rs b/gateway/src/routes/realtime.rs index 03ddab9b..fb08c511 100644 --- a/gateway/src/routes/realtime.rs +++ b/gateway/src/routes/realtime.rs @@ -44,21 +44,16 @@ use std::sync::Arc; /// // Client sends audio as binary frames /// // Server sends back transcripts and audio /// ``` -/// Every path the realtime handler is served at. +/// Every path the NATIVE realtime handler is served at. /// /// A CONSTANT the router iterates, rather than a list of `.route()` calls, so the test below -/// asserts on the same data the router uses. A test that greps this file for `.route(...)` -/// would be checking its own transcription of the truth; this one cannot drift from it. +/// asserts on the same data the router uses. /// -/// * `/realtime` — WaaV's native path, kept for standalone deployments already using it. -/// * `/v1/realtime` — the OpenAI-compatible path, and the one Bud's ingress routes here. -/// -/// Both are needed because a Kubernetes Ingress does NOT rewrite paths: a rule sending -/// `/v1/realtime` to this service delivers it verbatim, so serving only `/realtime` left the -/// ingress pointing at a route that did not exist. The failure is invisible from outside — -/// the auth middleware answers 401 before routing, so a missing route and a missing -/// credential are indistinguishable. -pub const REALTIME_PATHS: &[&str] = &["/realtime", "/v1/realtime"]; +/// Only `/realtime`: FRD-023 D-1 made `/v1/realtime` the OpenAI Realtime GA endpoint for Bud +/// deployments (`routes::openai_realtime`). WaaV's own SDKs connect to `/realtime`, and the +/// native protocol could never have served a Bud deployment on `/v1/realtime` (it had no way to +/// name one, nor a credential to use). +pub const REALTIME_PATHS: &[&str] = &["/realtime"]; pub fn create_realtime_router() -> Router> { let mut router = Router::new(); @@ -73,11 +68,13 @@ mod route_tests { use super::REALTIME_PATHS; #[test] - fn the_openai_compatible_path_is_served() { - assert!( - REALTIME_PATHS.contains(&"/v1/realtime"), - "Bud's ingress routes /v1/realtime here and does not rewrite the path; \ - dropping it makes the ingress point at nothing, which reads as a 401" + fn the_openai_compatible_path_is_not_the_native_handler() { + // FRD-023 D-1: `/v1/realtime` speaks OpenAI GA (routes::openai_realtime). Serving it with + // the native handler too would register the path twice — axum panics at construction. + assert!(!REALTIME_PATHS.contains(&"/v1/realtime")); + assert_eq!( + crate::handlers::openai_realtime::OPENAI_REALTIME_PATH, + "/v1/realtime" ); } diff --git a/gateway/src/state/mod.rs b/gateway/src/state/mod.rs index 69fbfe45..2b06ca5f 100644 --- a/gateway/src/state/mod.rs +++ b/gateway/src/state/mod.rs @@ -72,6 +72,9 @@ pub struct AppState { /// `GET /transcribe/batch/{job_id}` can return them. In-process (single-node); a multi-node /// deployment would back this with a shared store. pub batch_jobs: Arc>, + + /// FRD-023: `/v1/realtime` session timings and the `ek_bud_` sealing keys. + pub realtime: Arc, } impl AppState { @@ -441,6 +444,10 @@ impl AppState { None }; + // FRD-023: a bad WAAV_CLIENT_SECRET_KEYS must stop the pod (TC-EK-16), not leave a gateway + // whose client secrets silently never open. + let realtime = Arc::new(crate::handlers::openai_realtime::RealtimeRuntime::from_env()?); + Ok(Arc::new(Self { // Installed after construction by main(), once the control-plane connection is up: // AppState::try_new runs before the Redis URL is known. @@ -457,6 +464,7 @@ impl AppState { connections_per_ip: Arc::new(DashMap::new()), shutdown: CancellationToken::new(), batch_jobs: Arc::new(DashMap::new()), + realtime, })) } diff --git a/gateway/tests/openai_realtime_relay.rs b/gateway/tests/openai_realtime_relay.rs new file mode 100644 index 00000000..f6b7bc93 --- /dev/null +++ b/gateway/tests/openai_realtime_relay.rs @@ -0,0 +1,1990 @@ +//! FRD-023 `/v1/realtime` end to end, in process: a Bud-mode gateway over an in-memory control +//! plane, a mock OpenAI-GA vendor, and real WebSocket clients (TC-HS, TC-UP, TC-EVT, TC-LIFE, +//! TC-MET, TC-EK, TC-SEC-07). +//! +//! Each test uses its own deployment names and ids, so tests run in parallel against their own +//! gateway and vendor; spans are captured per test on the test's own (current-thread) runtime. + +use std::collections::HashMap; +use std::net::SocketAddr; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use futures_util::{SinkExt, StreamExt}; +use opentelemetry::trace::{SpanId, TracerProvider as _}; +use opentelemetry_sdk::error::OTelSdkResult; +use opentelemetry_sdk::trace::{SdkTracerProvider, SpanData, SpanExporter}; +use serde_json::{Value as Json, json}; +use tokio::net::TcpListener; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tracing_subscriber::layer::SubscriberExt; + +use waav_gateway::config::{DAGTimeoutsConfig, PluginConfig, ServerConfig}; +use waav_gateway::handlers::openai_realtime::{RealtimeRuntime, Timings}; +use waav_gateway::state::AppState; + +// ============================================================================================= +// Fixtures +// ============================================================================================= + +const KEY: &str = "bud_realtime_relay_test_key"; +const OTHER_KEY: &str = "bud_realtime_relay_other_project_key"; +const PROJECT: &str = "5b0c7e1d-0000-4000-8000-00000000aa01"; +const OTHER_PROJECT: &str = "5b0c7e1d-0000-4000-8000-00000000aa02"; +const USER: &str = "5b0c7e1d-0000-4000-8000-00000000bb01"; +const API_KEY_ID: &str = "5b0c7e1d-0000-4000-8000-00000000cc01"; +const MODEL_ID: &str = "5b0c7e1d-0000-4000-8000-00000000dd01"; + +/// bud-auth's fixture ciphertext; the plaintext is its `PLAIN`. +const TEST_CREDENTIAL: &str = include_str!("../../bud-auth/tests/fixtures/test_cred_encrypted.hex"); +const VENDOR_KEY: &str = "dg_vendor_key_abc123"; + +const JWKS: &str = include_str!("../../bud-auth/tests/fixtures/test_jwks.json"); +const JWT_ISSUER: &str = "https://auth.test/realms/bud"; + +fn test_pem() -> String { + std::fs::read_to_string(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../bud-auth/tests/fixtures/test_cred_private.pem" + )) + .expect("bud-auth's fixture key (git-ignored *.pem) must be present locally") +} + +fn jwt_private_pem() -> String { + std::fs::read_to_string(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../bud-auth/tests/fixtures/test_rsa_private.pem" + )) + .expect("bud-auth's JWT fixture key must be present locally") +} + +fn now() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs() +} + +fn jwt(sub: &str, exp_in: i64) -> String { + let mut header = jsonwebtoken::Header::new(jsonwebtoken::Algorithm::RS256); + header.kid = Some("test-key-1".into()); + let key = jsonwebtoken::EncodingKey::from_rsa_pem(jwt_private_pem().as_bytes()).unwrap(); + let now = now() as i64; + jsonwebtoken::encode( + &header, + &json!({"iss": JWT_ISSUER, "sub": sub, "azp": "bud-playground", "exp": now + exp_in, "iat": now}), + &key, + ) + .unwrap() +} + +/// The loopback escape hatch, so the mock vendor on 127.0.0.1 passes SSRF validation. Set once, +/// before any test builds a gateway; the SSRF refusal itself is covered by lib tests (TC-UP-04). +fn allow_loopback() { + static ONCE: std::sync::Once = std::sync::Once::new(); + ONCE.call_once(|| unsafe { std::env::set_var("WAAV_ALLOW_LOOPBACK_ENDPOINTS", "1") }); +} + +// ============================================================================================= +// Span capture (per test, thread-local) +// ============================================================================================= + +#[derive(Debug, Clone, Default)] +struct Exported(Arc>>); + +impl SpanExporter for Exported { + fn export( + &self, + batch: Vec, + ) -> impl std::future::Future + Send { + self.0.lock().unwrap().extend(batch); + std::future::ready(Ok(())) + } +} + +struct Capture { + exported: Exported, + _guard: tracing::subscriber::DefaultGuard, + _provider: SdkTracerProvider, + logs: Arc>>, +} + +#[derive(Clone)] +struct LogSink(Arc>>); +impl std::io::Write for LogSink { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + self.0.lock().unwrap().extend_from_slice(buf); + Ok(buf.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +impl Capture { + fn install() -> Self { + let exported = Exported::default(); + let provider = SdkTracerProvider::builder() + .with_simple_exporter(exported.clone()) + .build(); + let logs = Arc::new(Mutex::new(Vec::new())); + let sink = LogSink(logs.clone()); + let subscriber = tracing_subscriber::registry() + .with(tracing_subscriber::filter::LevelFilter::DEBUG) + .with(tracing_opentelemetry::layer().with_tracer(provider.tracer("realtime-relay"))) + .with( + tracing_subscriber::fmt::layer() + .with_writer(move || sink.clone()) + .with_ansi(false), + ); + let guard = tracing::subscriber::set_default(subscriber); + Self { + exported, + _guard: guard, + _provider: provider, + logs, + } + } + + fn spans(&self) -> Vec { + self.exported.0.lock().unwrap().clone() + } + + fn logs(&self) -> String { + String::from_utf8_lossy(&self.logs.lock().unwrap()).into_owned() + } + + async fn wait_for(&self, name: &str, n: usize) -> Vec { + for _ in 0..300 { + let found: Vec = self + .spans() + .into_iter() + .filter(|s| s.name == name) + .collect(); + if found.len() >= n { + return found; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!( + "expected {n} `{name}` spans, got {:?}", + self.spans() + .iter() + .map(|s| s.name.to_string()) + .collect::>() + ); + } +} + +fn attr(span: &SpanData, key: &str) -> Option { + span.attributes + .iter() + .find(|kv| kv.key.as_str() == key) + .map(|kv| kv.value.clone()) +} + +fn text(span: &SpanData, key: &str) -> Option { + attr(span, key).map(|v| v.as_str().into_owned()) +} + +fn number(span: &SpanData, key: &str) -> Option { + match attr(span, key)? { + opentelemetry::Value::I64(i) => Some(i as f64), + opentelemetry::Value::F64(f) => Some(f), + opentelemetry::Value::String(s) => s.as_str().parse().ok(), + _ => None, + } +} + +// ============================================================================================= +// The mock vendor (OpenAI Realtime GA) +// ============================================================================================= + +#[derive(Clone)] +struct Behaviour { + /// Complete the WebSocket handshake at all (TC-UP-07: a vendor that never answers). + accept: bool, + /// Answer a `session.update` with `session.updated` (TC-EVT-15 turns it off). + answer_updates: bool, + /// Send `rate_limits.updated` after each response (TC-EVT-12). + rate_limits: bool, + /// Audio deltas per response. + audio_deltas: usize, + /// Size of each audio delta's base64 payload. + delta_bytes: usize, + /// Stop reading (and therefore ponging) after this many client frames (TC-LIFE-03). + go_silent_after: Option, + /// Answer `input_audio_buffer.commit` with a completed transcription carrying this usage. + transcription_usage: Option, + usage: Json, +} + +impl Default for Behaviour { + fn default() -> Self { + Self { + accept: true, + answer_updates: true, + rate_limits: false, + audio_deltas: 2, + delta_bytes: 16, + go_silent_after: None, + transcription_usage: None, + usage: documented_usage(), + } + } +} + +fn documented_usage() -> Json { + json!({ + "total_tokens": 253, "input_tokens": 132, "output_tokens": 121, + "input_token_details": {"text_tokens": 119, "audio_tokens": 13, "image_tokens": 0, + "cached_tokens": 64, "cached_tokens_details": {"text_tokens": 64, "audio_tokens": 0, "image_tokens": 0}}, + "output_token_details": {"text_tokens": 30, "audio_tokens": 91} + }) +} + +#[derive(Default)] +struct VendorLog { + /// `(path?query, lower-cased headers)` per upgrade. + upgrades: Vec<(String, HashMap)>, + /// Every client frame the vendor received, in order. + frames: Vec, + connections: usize, +} + +#[derive(Clone)] +struct MockVendor { + addr: SocketAddr, + log: Arc>, +} + +impl MockVendor { + async fn start(b: Behaviour) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let log = Arc::new(Mutex::new(VendorLog::default())); + let shared = log.clone(); + tokio::spawn(async move { + loop { + let Ok((stream, _)) = listener.accept().await else { + return; + }; + shared.lock().unwrap().connections += 1; + let b = b.clone(); + let log = shared.clone(); + tokio::spawn(async move { + if !b.accept { + // Hold the TCP connection and never answer the upgrade. + tokio::time::sleep(Duration::from_secs(120)).await; + drop(stream); + return; + } + let captured = log.clone(); + let callback = move |req: &tokio_tungstenite::tungstenite::handshake::server::Request, + resp: tokio_tungstenite::tungstenite::handshake::server::Response| { + let headers = req + .headers() + .iter() + .map(|(k, v)| (k.as_str().to_ascii_lowercase(), v.to_str().unwrap_or("").to_string())) + .collect(); + let pq = req.uri().path_and_query().map(|p| p.to_string()).unwrap_or_default(); + captured.lock().unwrap().upgrades.push((pq, headers)); + Ok(resp) + }; + let Ok(mut ws) = tokio_tungstenite::accept_hdr_async(stream, callback).await + else { + return; + }; + let model = log + .lock() + .unwrap() + .upgrades + .last() + .and_then(|(pq, _)| pq.split("model=").nth(1).map(str::to_string)) + .unwrap_or_default(); + let created = json!({"type": "session.created", "event_id": "evt_v0", + "session": {"id": "sess_vendor_1", "object": "realtime.session", "type": "realtime", "model": model}}); + if ws + .send(Message::Text(created.to_string().into())) + .await + .is_err() + { + return; + } + let mut received = 0usize; + let mut response_n = 0usize; + while let Some(Ok(msg)) = ws.next().await { + let Message::Text(t) = msg else { continue }; + let v: Json = serde_json::from_str(t.as_str()).unwrap_or(Json::Null); + log.lock().unwrap().frames.push(v.clone()); + received += 1; + if b.go_silent_after.is_some_and(|n| received >= n) { + // Stop reading: no more pongs, no more answers. + tokio::time::sleep(Duration::from_secs(120)).await; + return; + } + let kind = v["type"].as_str().unwrap_or_default().to_string(); + let mut out: Vec = Vec::new(); + match kind.as_str() { + "session.update" if b.answer_updates => { + out.push(json!({"type": "session.updated", "event_id": "evt_vu", + "session": {"id": "sess_vendor_1", "model": model, "type": v["session"]["type"]}})); + } + "response.create" => { + response_n += 1; + let rid = format!("resp_{response_n}"); + out.push( + json!({"type": "response.created", "response": {"id": rid}}), + ); + for i in 0..b.audio_deltas { + out.push(json!({"type": "response.output_audio.delta", "response_id": rid, + "item_id": "item_1", "delta": format!("{:0>width$}", i, width = b.delta_bytes)})); + } + if b.rate_limits { + out.push(json!({"type": "rate_limits.updated", "rate_limits": [{"name": "tokens", "remaining": 1}]})); + } + out.push(json!({"type": "response.done", "event_id": format!("evt_done_{response_n}"), + "response": {"id": rid, "status": "completed", "usage": b.usage, + "output": [{"type": "message", "content": [{"type": "output_audio", "transcript": "hello there"}]}]}})); + } + "input_audio_buffer.commit" => { + if let Some(u) = &b.transcription_usage { + out.push(json!({"type": "conversation.item.input_audio_transcription.completed", + "item_id": "item_u", "content_index": 0, "transcript": "hi", "usage": u})); + } + } + _ => {} + } + for o in out { + if ws.send(Message::Text(o.to_string().into())).await.is_err() { + return; + } + } + } + }); + } + }); + Self { addr, log } + } + + fn base(&self) -> String { + format!("http://{}/v1", self.addr) + } + + fn frames(&self) -> Vec { + self.log.lock().unwrap().frames.clone() + } + + fn frames_of(&self, kind: &str) -> Vec { + self.frames() + .into_iter() + .filter(|f| f["type"] == kind) + .collect() + } + + fn upgrades(&self) -> Vec<(String, HashMap)> { + self.log.lock().unwrap().upgrades.clone() + } + + fn connections(&self) -> usize { + self.log.lock().unwrap().connections + } +} + +// ============================================================================================= +// The gateway +// ============================================================================================= + +fn config() -> ServerConfig { + ServerConfig { + host: "127.0.0.1".to_string(), + port: 0, + tls: None, + livekit_url: "ws://localhost:7880".to_string(), + livekit_public_url: "http://localhost:7880".to_string(), + livekit_api_key: None, + livekit_api_secret: None, + deepgram_api_key: None, + elevenlabs_api_key: None, + google_credentials: None, + azure_speech_subscription_key: None, + azure_speech_region: None, + cartesia_api_key: None, + openai_api_key: Some("sk-process-canary".to_string()), + azure_openai_api_key: None, + azure_openai_endpoint: None, + grok_api_key: None, + inworld_api_key: None, + gemini_api_key: None, + ultravox_api_key: None, + speechmatics_api_key: None, + yandex_api_key: None, + yandex_folder_id: None, + assemblyai_api_key: None, + hume_api_key: None, + groq_api_key: None, + ibm_watson_api_key: None, + ibm_watson_instance_id: None, + ibm_watson_region: None, + aws_access_key_id: None, + aws_secret_access_key: None, + aws_region: None, + gnani_token: None, + gnani_access_key: None, + gnani_certificate_path: None, + recording_s3_bucket: None, + recording_s3_region: None, + recording_s3_endpoint: None, + recording_s3_access_key: None, + recording_s3_secret_key: None, + recording_s3_prefix: None, + cache_path: None, + cache_ttl_seconds: Some(3600), + auth_service_url: None, + auth_signing_key_path: None, + auth_api_secrets: Vec::new(), + auth_timeout_seconds: 5, + auth_required: false, + sip: None, + cors_allowed_origins: None, + rate_limit_requests_per_second: 60, + rate_limit_burst_size: 10, + max_websocket_connections: None, + max_connections_per_ip: 1000, + ws_processing_timeout_secs: 10, + realtime_processing_timeout_secs: 30, + sip_max_participants: 3, + realtime_endpoint_overrides: Default::default(), + plugins: PluginConfig::default(), + dag_timeouts: DAGTimeoutsConfig::default(), + aliases: Default::default(), + } +} + +fn fast_timings() -> Timings { + Timings { + ping: Duration::from_millis(300), + max_missed_pongs: 3, + revalidate: Duration::from_millis(300), + connect: Duration::from_millis(1500), + hold: Duration::from_millis(700), + slow_client: Duration::from_millis(600), + max_session: Duration::from_secs(3600), + default_idle: Duration::from_secs(300), + warn_before: Duration::from_secs(1), + segment: Duration::from_secs(1), + upstream_send: Duration::from_secs(2), + } +} + +/// A realtime `voice_table` entry pointed at the mock vendor. +fn rt_entry(vendor: &MockVendor, extra: Json) -> Json { + let mut e = json!({ + "vendor": "openai", + "api_base": vendor.base(), + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["realtime_session"], + "model": "gpt-realtime-2.1", + "pricing": {"unit": "token", "per_units": 1000000, "currency": "USD", + "rates": {"input_text": 4.0, "input_audio": 32.0, "input_image": 5.0, + "cached_input_text": 0.4, "cached_input_audio": 0.4, "cached_input_image": 0.5, + "output_text": 24.0, "output_audio": 64.0, "transcription_per_minute": 0.003}} + }); + if let Some(obj) = extra.as_object() { + for (k, v) in obj { + e[k] = v.clone(); + } + } + e +} + +struct Gateway { + addr: SocketAddr, + store: Arc, + plane: Arc, + state: Arc, +} + +struct Setup { + /// `(alias, endpoint id, entry)` reachable by KEY in PROJECT. + endpoints: Vec<(String, String, Json)>, + /// Endpoints in the table that only OTHER_KEY reaches. + other_endpoints: Vec<(String, String, Json)>, + timings: Timings, + client_secret_keys: Option<&'static str>, + jwt_subjects: Vec<&'static str>, +} + +impl Default for Setup { + fn default() -> Self { + Self { + endpoints: Vec::new(), + other_endpoints: Vec::new(), + timings: fast_timings(), + client_secret_keys: Some("k1:BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc="), + jwt_subjects: Vec::new(), + } + } +} + +struct StaticJwks; +#[async_trait::async_trait] +impl bud_auth::JwksSource for StaticJwks { + async fn fetch(&self) -> Result { + Ok(JWKS.to_string()) + } +} + +async fn gateway(setup: Setup) -> Gateway { + allow_loopback(); + if std::env::var("RUST_LOG").is_ok() { + let _ = tracing_subscriber::fmt() + .with_env_filter(tracing_subscriber::EnvFilter::from_default_env()) + .with_test_writer() + .try_init(); + } + let store = Arc::new(bud_auth::MemoryStore::new()); + let blob = |eps: &[(String, String, Json)], project: &str| { + let mut m = serde_json::Map::new(); + for (alias, id, _) in eps { + m.insert(alias.clone(), json!({"endpoint_id": id, "model_id": MODEL_ID, "project_id": project, "kind": "model"})); + } + m.insert( + "__metadata__".into(), + json!({"api_key_id": API_KEY_ID, "user_id": USER, "api_key_project_id": project}), + ); + Json::Object(m).to_string() + }; + store.set( + &format!("api_key:{}", bud_auth::hash_api_key(KEY)), + &blob(&setup.endpoints, PROJECT), + ); + store.set( + &format!("api_key:{}", bud_auth::hash_api_key(OTHER_KEY)), + &blob(&setup.other_endpoints, OTHER_PROJECT), + ); + for (_, id, entry) in setup.endpoints.iter().chain(setup.other_endpoints.iter()) { + store.set( + &format!("voice_table:{id}"), + &json!({ id.as_str(): entry }).to_string(), + ); + } + // JWT subjects reach the KEY project's deployments. + let mut project_models = serde_json::Map::new(); + for (alias, id, _) in &setup.endpoints { + project_models.insert( + alias.clone(), + json!({"endpoint_id": id, "model_id": MODEL_ID, "project_id": PROJECT}), + ); + } + store.set( + &format!("project_models:{PROJECT}"), + &Json::Object(project_models).to_string(), + ); + for sub in &setup.jwt_subjects { + store.set( + &format!("user_projects:{sub}"), + &json!({"user_id": USER, "projects": [PROJECT]}).to_string(), + ); + } + + let jwt_cfg = bud_auth::JwtConfig::from_lookup(|k| match k { + "OIDC_ISSUER" => Some(JWT_ISSUER.into()), + "OIDC_ALLOWED_CLIENTS" => Some("bud-playground".into()), + "OIDC_AUTHZ_TTL_SECS" => Some("300".into()), + _ => None, + }) + .unwrap(); + let verifier = Arc::new(bud_auth::JwtVerifier::new(jwt_cfg, Arc::new(StaticJwks))); + verifier.prime().await.unwrap(); + let plane = Arc::new(bud_auth::BudPlane::with_decryptor( + store.clone() as Arc, + Some(verifier), + bud_auth::CredentialDecryptor::from_pem(&test_pem()).unwrap(), + )); + plane.boot().await.unwrap(); + + let mut state = AppState::new(config()).await; + { + let s = Arc::get_mut(&mut state).expect("unshared"); + s.bud_mode = Some(waav_gateway::auth::bud_mode::BudMode::for_plane(plane.clone()).unwrap()); + s.policies = Some(waav_gateway::core::deployment_policy::DeploymentPolicies::local()); + s.realtime = Arc::new(RealtimeRuntime { + timings: setup.timings, + client_secret_keys: setup + .client_secret_keys + .map(|k| waav_gateway::auth::ephemeral::ClientSecretKeys::parse(k).unwrap()), + }); + } + let app = waav_gateway::routes::openai_realtime::create_openai_realtime_router() + .layer(axum::middleware::from_fn_with_state( + state.clone(), + waav_gateway::middleware::connection_limit_middleware, + )) + .with_state(state.clone()); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + axum::serve( + listener, + app.into_make_service_with_connect_info::(), + ) + .await + .unwrap(); + }); + Gateway { + addr, + store, + plane, + state, + } +} + +fn ep(alias: &str, id: &str, entry: Json) -> (String, String, Json) { + (alias.to_string(), id.to_string(), entry) +} + +// ============================================================================================= +// The client +// ============================================================================================= + +type Client = + tokio_tungstenite::WebSocketStream>; + +struct Connected { + ws: Client, + protocol: Option, + extensions: Option, +} + +enum Auth<'a> { + Bearer(&'a str), + ApiKey(&'a str), + Subprotocol(&'a str), + None, +} + +async fn connect_with( + gw: &Gateway, + query: &str, + auth: Auth<'_>, + extra_headers: &[(&str, &str)], + extra_protocols: &[&str], +) -> Result)> { + let mut req = format!("ws://{}/v1/realtime?{query}", gw.addr) + .into_client_request() + .unwrap(); + let mut protocols: Vec = Vec::new(); + match auth { + Auth::Bearer(k) => { + req.headers_mut() + .insert("authorization", format!("Bearer {k}").parse().unwrap()); + } + Auth::ApiKey(k) => { + req.headers_mut().insert("api-key", k.parse().unwrap()); + } + Auth::Subprotocol(k) => { + protocols.push("realtime".into()); + protocols.push(format!("openai-insecure-api-key.{k}")); + } + Auth::None => {} + } + protocols.extend(extra_protocols.iter().map(|p| p.to_string())); + if !protocols.is_empty() { + req.headers_mut().insert( + "sec-websocket-protocol", + protocols.join(", ").parse().unwrap(), + ); + } + for (k, v) in extra_headers { + req.headers_mut().insert( + axum::http::HeaderName::from_bytes(k.as_bytes()).unwrap(), + v.parse().unwrap(), + ); + } + let connected = tokio::time::timeout( + Duration::from_secs(15), + tokio_tungstenite::connect_async(req), + ) + .await + .expect("the handshake did not complete within 15 s"); + match connected { + Ok((ws, resp)) => Ok(Connected { + protocol: resp + .headers() + .get("sec-websocket-protocol") + .map(|v| v.to_str().unwrap().to_string()), + extensions: resp + .headers() + .get("sec-websocket-extensions") + .map(|v| v.to_str().unwrap().to_string()), + ws, + }), + Err(tokio_tungstenite::tungstenite::Error::Http(resp)) => { + let status = resp.status().as_u16(); + let headers = resp + .headers() + .iter() + .map(|(k, v)| (k.as_str().to_string(), v.to_str().unwrap_or("").to_string())) + .collect(); + let body = resp + .body() + .as_ref() + .and_then(|b| serde_json::from_slice(b).ok()) + .unwrap_or(Json::Null); + Err((status, body, headers)) + } + Err(e) => panic!("unexpected connect error: {e}"), + } +} + +async fn connect(gw: &Gateway, model: &str) -> Client { + connect_with(gw, &format!("model={model}"), Auth::Bearer(KEY), &[], &[]) + .await + .unwrap_or_else(|(s, b, _)| panic!("connect refused {s}: {b}")) + .ws +} + +async fn next_json(ws: &mut Client) -> Json { + loop { + match tokio::time::timeout(Duration::from_secs(5), ws.next()).await { + Ok(Some(Ok(Message::Text(t)))) => return serde_json::from_str(t.as_str()).unwrap(), + Ok(Some(Ok(Message::Ping(_) | Message::Pong(_)))) => continue, + other => panic!("expected a text event, got {other:?}"), + } + } +} + +async fn until_type(ws: &mut Client, kind: &str) -> Json { + for _ in 0..200 { + let v = next_json(ws).await; + if v["type"] == kind { + return v; + } + } + panic!("no {kind}"); +} + +/// Read until the close frame; returns (error codes seen, close code). +async fn until_close(ws: &mut Client) -> (Vec, Option) { + let mut errors = Vec::new(); + loop { + match tokio::time::timeout(Duration::from_secs(10), ws.next()).await { + Ok(Some(Ok(Message::Text(t)))) => { + let v: Json = serde_json::from_str(t.as_str()).unwrap(); + if v["type"] == "error" { + errors.push(v["error"]["code"].as_str().unwrap_or_default().to_string()); + } + } + Ok(Some(Ok(Message::Close(frame)))) => { + return (errors, frame.map(|f| u16::from(f.code))); + } + Ok(Some(Ok(_))) => continue, + Ok(Some(Err(_))) | Ok(None) => return (errors, None), + Err(_) => panic!("no close within 10 s; errors so far {errors:?}"), + } + } +} + +/// Wait while still reading, as a live client does: the reads are what answer the server's +/// pings. A test that merely sleeps looks like a dead client and is closed (`client_timeout`). +async fn keep_alive(ws: &mut Client, dur: Duration) { + let deadline = tokio::time::Instant::now() + dur; + while tokio::time::Instant::now() < deadline { + let _ = tokio::time::timeout_at(deadline, ws.next()).await; + } +} + +async fn send(ws: &mut Client, v: Json) { + ws.send(Message::Text(v.to_string().into())).await.unwrap(); +} + +// ============================================================================================= +// Handshake — TC-HS +// ============================================================================================= + +const EP_HS: &str = "a1a1a1a1-0000-4000-8000-000000000001"; + +async fn hs_gateway() -> (MockVendor, Gateway) { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep("rt", EP_HS, rt_entry(&vendor, json!({})))], + jwt_subjects: vec!["sub-hs"], + ..Default::default() + }) + .await; + (vendor, gw) +} + +#[tokio::test] +async fn tc_hs_01_02_bearer_and_api_key_header() { + let (_v, gw) = hs_gateway().await; + let mut a = connect_with(&gw, "model=rt", Auth::Bearer(KEY), &[], &[]) + .await + .unwrap(); + assert_eq!( + until_type(&mut a.ws, "session.created").await["session"]["model"], + "rt" + ); + let mut b = connect_with(&gw, "model=rt", Auth::ApiKey(KEY), &[], &[]) + .await + .unwrap(); + until_type(&mut b.ws, "session.created").await; +} + +/// TC-HS-03 🔒 — `realtime` selected, the credential never echoed. +#[tokio::test] +async fn tc_hs_03_subprotocol_auth_selects_realtime_and_never_echoes_the_key() { + let (_v, gw) = hs_gateway().await; + let c = connect_with(&gw, "model=rt", Auth::Subprotocol(KEY), &[], &[]) + .await + .unwrap(); + assert_eq!(c.protocol.as_deref(), Some("realtime")); +} + +#[tokio::test] +async fn tc_hs_04_extra_subprotocols_are_tolerated() { + let (_v, gw) = hs_gateway().await; + let c = connect_with( + &gw, + "model=rt", + Auth::Subprotocol(KEY), + &[], + &["openai-agents-sdk.v0.18", "openai-project.p"], + ) + .await + .unwrap(); + assert_eq!(c.protocol.as_deref(), Some("realtime")); +} + +#[tokio::test] +async fn tc_hs_05_a_keycloak_jwt_through_the_subprotocol() { + let (_v, gw) = hs_gateway().await; + let token = jwt("sub-hs", 300); + let mut c = connect_with(&gw, "model=rt", Auth::Subprotocol(&token), &[], &[]) + .await + .unwrap(); + until_type(&mut c.ws, "session.created").await; +} + +#[tokio::test] +async fn tc_hs_06_to_09_pre_upgrade_refusals_are_openai_envelopes() { + let (_v, gw) = hs_gateway().await; + for (query, auth_key, headers, protocols, status, code) in [ + ("", Some(KEY), vec![], vec![], 400, "model_required"), + ( + "token=x&model=rt", + None, + vec![], + vec![], + 400, + "use_subprotocol", + ), + ( + "model=rt", + Some(KEY), + vec![("openai-beta", "realtime=v1")], + vec![], + 400, + "beta_api_shape_disabled", + ), + ( + "model=rt", + Some(KEY), + vec![], + vec!["openai-beta.realtime-v1"], + 400, + "beta_api_shape_disabled", + ), + ( + "model=rt&call_id=rtc_x", + Some(KEY), + vec![], + vec![], + 400, + "unsupported_parameter", + ), + ( + "model=rt", + Some("bud_not_a_key"), + vec![], + vec![], + 401, + "invalid_api_key", + ), + ] { + let auth = auth_key.map(Auth::Bearer).unwrap_or(Auth::None); + let Err((s, body, _)) = connect_with(&gw, query, auth, &headers, &protocols).await else { + panic!("{query} {code}: upgraded"); + }; + assert_eq!(s, status, "{code}: {body}"); + assert_eq!(body["error"]["code"], code, "{body}"); + assert!(body["error"]["message"].is_string()); + } +} + +#[tokio::test] +async fn tc_hs_10_intent_is_ignored() { + let (_v, gw) = hs_gateway().await; + let mut c = connect_with( + &gw, + "model=rt&intent=transcription", + Auth::Bearer(KEY), + &[], + &[], + ) + .await + .unwrap(); + until_type(&mut c.ws, "session.created").await; +} + +#[tokio::test] +async fn tc_hs_11_permessage_deflate_is_never_negotiated() { + let (_v, gw) = hs_gateway().await; + let c = connect_with( + &gw, + "model=rt", + Auth::Bearer(KEY), + &[( + "sec-websocket-extensions", + "permessage-deflate; client_max_window_bits", + )], + &[], + ) + .await + .unwrap(); + assert!(c.extensions.is_none(), "{:?}", c.extensions); +} + +#[tokio::test] +async fn tc_hs_13_to_16_resolution() { + let vendor = MockVendor::start(Behaviour::default()).await; + let tts = json!({"vendor": "openai", "endpoints": ["text_to_speech"], "model": "tts-1"}); + let gw = gateway(Setup { + endpoints: vec![ + ep("rt", EP_HS, rt_entry(&vendor, json!({}))), + ep("tts", "a1a1a1a1-0000-4000-8000-0000000000f1", tts), + ], + other_endpoints: vec![ep( + "theirs", + "a1a1a1a1-0000-4000-8000-0000000000f2", + rt_entry(&vendor, json!({})), + )], + ..Default::default() + }) + .await; + // TC-HS-13 unknown, TC-HS-14 another project's (by alias and by id), TC-HS-15 wrong capability. + for model in [ + "nope", + "theirs", + "a1a1a1a1-0000-4000-8000-0000000000f2", + "tts", + ] { + let Err((s, body, _)) = + connect_with(&gw, &format!("model={model}"), Auth::Bearer(KEY), &[], &[]).await + else { + panic!("{model} upgraded"); + }; + assert_eq!(s, 404, "{model}: {body}"); + assert_eq!(body["error"]["code"], "model_not_found"); + assert!( + body["error"]["message"] + .as_str() + .unwrap() + .contains("realtime_session") + ); + } + // TC-HS-16 the raw endpoint UUID the key reaches. + let mut c = connect(&gw, EP_HS).await; + until_type(&mut c, "session.created").await; +} + +// ============================================================================================= +// Upstream — TC-UP +// ============================================================================================= + +/// TC-UP-01 / TC-UP-05 — URL, auth header, no beta header; the model is the deployment's. +#[tokio::test] +async fn tc_up_01_05_the_vendor_request() { + let (vendor, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + send(&mut c, json!({"type": "session.update", "session": {"type": "realtime", "model": "gpt-4o-realtime-preview"}})).await; + until_type(&mut c, "session.updated").await; + + let (pq, headers) = vendor.upgrades()[0].clone(); + assert_eq!(pq, "/v1/realtime?model=gpt-realtime-2.1"); + assert_eq!( + headers.get("authorization").map(String::as_str), + Some(&*format!("Bearer {VENDOR_KEY}")) + ); + assert!(!headers.contains_key("openai-beta")); + assert!( + !headers.values().any(|v| v.contains(KEY)), + "the Bud key reached the vendor" + ); + let update = &vendor.frames_of("session.update")[0]; + assert!(update["session"].get("model").is_none(), "{update}"); +} + +/// TC-UP-07 — a vendor that never answers: `upstream_error` + 1011 within the deadline, and the +/// admission is released (a second session is admitted under `max_concurrent: 1`). +#[tokio::test] +async fn tc_up_07_connect_deadline_releases_the_admission() { + let vendor = MockVendor::start(Behaviour { + accept: false, + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000071", + rt_entry(&vendor, json!({"max_concurrent": 1})), + )], + ..Default::default() + }) + .await; + let started = std::time::Instant::now(); + let mut c = connect(&gw, "rt").await; + let (errors, code) = until_close(&mut c).await; + assert_eq!(code, Some(1011)); + assert_eq!(errors, vec!["upstream_error"]); + assert!(started.elapsed() < Duration::from_secs(5)); + // The slot came back. + tokio::time::sleep(Duration::from_millis(100)).await; + assert!( + connect_with(&gw, "model=rt", Auth::Bearer(KEY), &[], &[]) + .await + .is_ok() + ); +} + +/// TC-UP-06 — an open breaker refuses before connecting. +#[tokio::test] +async fn tc_up_06_an_open_breaker_refuses_with_503() { + let (vendor, gw) = hs_gateway().await; + let vkey = waav_gateway::core::deployment_policy::vendor_key("openai", Some(&vendor.base())); + gw.state + .policies + .as_ref() + .unwrap() + .breakers() + .vendor + .open_for(&vkey, Duration::from_secs(30)); + let Err((s, body, headers)) = connect_with(&gw, "model=rt", Auth::Bearer(KEY), &[], &[]).await + else { + panic!("upgraded through an open breaker"); + }; + assert_eq!(s, 503); + assert_eq!(body["error"]["code"], "circuit_open"); + assert!(headers.contains_key("retry-after")); + assert_eq!(vendor.connections(), 0, "no connect attempt"); +} + +// ============================================================================================= +// Event policy — TC-EVT +// ============================================================================================= + +async fn policy_gateway(extra: Json) -> (MockVendor, Gateway) { + let vendor = MockVendor::start(Behaviour { + rate_limits: true, + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-0000000000e1", + rt_entry(&vendor, extra), + )], + ..Default::default() + }) + .await; + (vendor, gw) +} + +/// TC-EVT-02 / 04 / 05 / 12 / 13 through a real session. +#[tokio::test] +async fn tc_evt_policy_through_a_live_session() { + let (vendor, gw) = policy_gateway(json!({})).await; + let mut c = connect(&gw, "rt").await; + let created = until_type(&mut c, "session.created").await; + assert_eq!(created["session"]["model"], "rt", "TC-EVT-13"); + + // TC-EVT-02 🔒: MCP refused, session continues. + send(&mut c, json!({"type": "session.update", "event_id": "c_mcp", "session": {"tools": [{"type": "mcp", "server_url": "https://x"}]}})).await; + let e = until_type(&mut c, "error").await; + assert_eq!(e["error"]["code"], "event_not_allowed"); + assert_eq!(e["error"]["param"], "session.tools.mcp"); + assert_eq!(e["error"]["event_id"], "c_mcp"); + + // TC-EVT-04 🔒: stored prompt refused. + send( + &mut c, + json!({"type": "session.update", "session": {"prompt": {"id": "pmpt_x"}}}), + ) + .await; + assert_eq!( + until_type(&mut c, "error").await["error"]["param"], + "session.prompt" + ); + + // TC-EVT-05: tracing stripped, forwarded. + send( + &mut c, + json!({"type": "session.update", "session": {"type": "realtime", "tracing": "auto"}}), + ) + .await; + until_type(&mut c, "session.updated").await; + + // TC-EVT-12 🔒: the vendor's rate_limits.updated never arrives. + send(&mut c, json!({"type": "response.create"})).await; + loop { + let v = next_json(&mut c).await; + assert_ne!(v["type"], "rate_limits.updated", "TC-EVT-12"); + if v["type"] == "response.done" { + break; + } + } + let updates = vendor.frames_of("session.update"); + assert_eq!( + updates.len(), + 1, + "refused frames never reached the vendor: {updates:?}" + ); + assert!(updates[0]["session"].get("tracing").is_none()); +} + +/// TC-EVT-14 — the vendor receives the defaults update BEFORE the client's, and the client's +/// voice wins (it arrives later). +#[tokio::test] +async fn tc_evt_14_defaults_first_then_the_client() { + let (vendor, gw) = policy_gateway(json!({"config": {"realtime": {"session_type": "realtime", + "defaults": {"voice": "marin", "instructions": "be brief"}}}})) + .await; + let mut c = connect(&gw, "rt").await; + // Sent immediately — before the vendor has even said session.created. + send(&mut c, json!({"type": "session.update", "session": {"type": "realtime", "audio": {"output": {"voice": "cedar"}}}})).await; + until_type(&mut c, "session.created").await; + for _ in 0..2 { + until_type(&mut c, "session.updated").await; + } + let updates = vendor.frames_of("session.update"); + assert_eq!(updates.len(), 2); + assert!( + updates[0]["event_id"] + .as_str() + .unwrap() + .starts_with("evt_bud_defaults_"), + "{updates:?}" + ); + assert_eq!(updates[0]["session"]["audio"]["output"]["voice"], "marin"); + assert_eq!(updates[0]["session"]["instructions"], "be brief"); + assert_eq!(updates[1]["session"]["audio"]["output"]["voice"], "cedar"); +} + +/// TC-EVT-15 — the vendor never answers the defaults update: `upstream_error` after the hold. +#[tokio::test] +async fn tc_evt_15_hold_timeout() { + let vendor = MockVendor::start(Behaviour { + answer_updates: false, + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-0000000000e5", + rt_entry( + &vendor, + json!({"config": {"realtime": {"defaults": {"voice": "marin"}}}}), + ), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + let (errors, code) = until_close(&mut c).await; + assert_eq!(code, Some(1011)); + assert_eq!(errors.last().map(String::as_str), Some("upstream_error")); +} + +/// TC-EVT-16 — a client that stops reading: 1011 `client_too_slow`, and every delta sent before +/// the close arrives in order. +#[tokio::test] +async fn tc_evt_16_slow_client_is_closed_without_silent_drops() { + let vendor = MockVendor::start(Behaviour { + audio_deltas: 2000, + delta_bytes: 4096, + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-0000000000e6", + rt_entry(&vendor, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + send(&mut c, json!({"type": "response.create"})).await; + // Stop reading for longer than the slow-client window. + tokio::time::sleep(Duration::from_millis(2500)).await; + let mut last = -1i64; + let mut close = None; + let mut errors = Vec::new(); + while let Ok(Some(Ok(msg))) = tokio::time::timeout(Duration::from_secs(10), c.next()).await { + match msg { + Message::Text(t) => { + let v: Json = serde_json::from_str(t.as_str()).unwrap(); + if v["type"] == "response.output_audio.delta" { + let i: i64 = v["delta"] + .as_str() + .unwrap() + .trim_start_matches('0') + .parse() + .unwrap_or(0); + assert!( + i == last + 1 || (last == -1 && i == 0), + "delta {i} after {last}: out of order or dropped" + ); + last = i; + } else if v["type"] == "error" { + errors.push(v["error"]["code"].as_str().unwrap().to_string()); + } + } + Message::Close(f) => { + close = f.map(|f| u16::from(f.code)); + break; + } + _ => {} + } + } + assert_eq!(close, Some(1011)); + assert_eq!(errors, vec!["client_too_slow"]); + assert!(last >= 0, "some deltas were delivered in order"); +} + +// ============================================================================================= +// Lifecycle — TC-LIFE +// ============================================================================================= + +/// TC-LIFE-01 — pings on both legs. +#[tokio::test] +async fn tc_life_01_server_pings() { + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + let mut pings = 0; + let deadline = tokio::time::Instant::now() + Duration::from_secs(2); + while tokio::time::Instant::now() < deadline { + if let Ok(Some(Ok(Message::Ping(_)))) = + tokio::time::timeout(Duration::from_millis(500), c.next()).await + { + pings += 1; + } + } + assert!(pings >= 2, "{pings} pings in 2 s at a 300 ms interval"); +} + +/// TC-LIFE-02 — a client that stops answering pongs is closed with 1011 after 3 missed. +#[tokio::test] +async fn tc_life_02_dead_client() { + let cap = Capture::install(); + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + // Stop reading: no pongs go back. + tokio::time::sleep(Duration::from_millis(2500)).await; + let session = cap.wait_for("voice.session", 1).await.remove(0); + assert_eq!( + text(&session, "bud.voice.session.end_reason").as_deref(), + Some("client_timeout") + ); + assert_eq!( + number(&session, "bud.voice.session.close_code"), + Some(1011.0) + ); +} + +/// TC-LIFE-03 — a vendor that stops answering pongs: `upstream_error` + 1011. +#[tokio::test] +async fn tc_life_03_dead_vendor() { + let vendor = MockVendor::start(Behaviour { + go_silent_after: Some(1), + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000103", + rt_entry(&vendor, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + send(&mut c, json!({"type": "input_audio_buffer.clear"})).await; + let (errors, code) = until_close(&mut c).await; + assert_eq!(code, Some(1011)); + assert!(errors.contains(&"upstream_error".to_string()), "{errors:?}"); +} + +/// TC-LIFE-05 — maximum length: a warning, then `session_expired` + 1000. +#[tokio::test] +async fn tc_life_05_maximum_length() { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000105", + rt_entry( + &vendor, + json!({"config": {"realtime": {"limits": {"max_session_seconds": 2}}}}), + ), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors, vec!["session_expiring", "session_expired"]); + assert_eq!(code, Some(1000)); +} + +/// TC-LIFE-04 — idle. +#[tokio::test] +async fn tc_life_04_idle() { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000104", + rt_entry( + &vendor, + json!({"config": {"realtime": {"limits": {"idle_timeout_seconds": 1}}}}), + ), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors.last().map(String::as_str), Some("session_expired")); + assert_eq!(code, Some(1000)); +} + +/// TC-LIFE-06 🔒 — the key is revoked mid-session: `session_revoked` + 1008 by the next timer. +#[tokio::test] +async fn tc_life_06_revoked_key_closes_on_the_timer() { + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + let key = format!("api_key:{}", bud_auth::hash_api_key(KEY)); + gw.store.remove(&key); + gw.plane + .on_key_event(&key, bud_auth::KeyEvent::Del) + .await + .unwrap(); + let started = std::time::Instant::now(); + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors, vec!["session_revoked"]); + assert_eq!(code, Some(1008)); + assert!(started.elapsed() < Duration::from_secs(2)); +} + +/// TC-LIFE-07 🔒 — revoked, then `response.create` at once: refused before forwarding. +#[tokio::test] +async fn tc_life_07_revoked_key_closes_on_response_create() { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000107", + rt_entry(&vendor, json!({})), + )], + timings: Timings { + revalidate: Duration::from_secs(3600), + ..fast_timings() + }, + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + let key = format!("api_key:{}", bud_auth::hash_api_key(KEY)); + gw.store.remove(&key); + gw.plane + .on_key_event(&key, bud_auth::KeyEvent::Del) + .await + .unwrap(); + send(&mut c, json!({"type": "response.create"})).await; + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors, vec!["session_revoked"]); + assert_eq!(code, Some(1008)); + assert!( + vendor.frames_of("response.create").is_empty(), + "forwarded after revocation" + ); +} + +/// TC-LIFE-08 — a JWT user removed from the project: closed by the next timer (Q-7 eviction). +#[tokio::test] +async fn tc_life_08_jwt_user_removed_from_the_project() { + let (_v, gw) = hs_gateway().await; + let token = jwt("sub-hs", 300); + let mut c = connect_with(&gw, "model=rt", Auth::Subprotocol(&token), &[], &[]) + .await + .unwrap() + .ws; + until_type(&mut c, "session.created").await; + gw.store.set( + "user_projects:sub-hs", + &json!({"user_id": USER, "projects": []}).to_string(), + ); + gw.plane + .on_key_event("user_projects:sub-hs", bud_auth::KeyEvent::Set) + .await + .unwrap(); + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors, vec!["session_revoked"]); + assert_eq!(code, Some(1008)); +} + +/// TC-LIFE-08 (second half) — JWT EXPIRY alone does not end a started session (D-17). +#[tokio::test] +async fn tc_life_08_jwt_expiry_does_not_close_a_started_session() { + let (_v, gw) = hs_gateway().await; + let token = jwt("sub-hs", 2); + let mut c = connect_with(&gw, "model=rt", Auth::Subprotocol(&token), &[], &[]) + .await + .unwrap() + .ws; + until_type(&mut c, "session.created").await; + keep_alive(&mut c, Duration::from_secs(3)).await; + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; +} + +/// TC-LIFE-09 — the endpoint is unpublished mid-session. +#[tokio::test] +async fn tc_life_09_unpublished_endpoint() { + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + let key = format!("voice_table:{EP_HS}"); + gw.store.remove(&key); + gw.plane + .on_key_event(&key, bud_auth::KeyEvent::Del) + .await + .unwrap(); + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors, vec!["session_revoked"]); + assert_eq!(code, Some(1008)); +} + +/// TC-LIFE-10 🔒 — one concurrency slot held per session and released at close. +#[tokio::test] +async fn tc_life_10_concurrency_held_and_released() { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000110", + rt_entry(&vendor, json!({"max_concurrent": 1})), + )], + ..Default::default() + }) + .await; + let mut first = connect(&gw, "rt").await; + until_type(&mut first, "session.created").await; + let Err((s, body, headers)) = connect_with(&gw, "model=rt", Auth::Bearer(KEY), &[], &[]).await + else { + panic!("a second session was admitted under max_concurrent: 1"); + }; + assert_eq!(s, 429); + assert_eq!(body["error"]["code"], "concurrency_limit_exceeded"); + assert!(headers.contains_key("retry-after")); + first.close(None).await.unwrap(); + let _ = until_close(&mut first).await; + tokio::time::sleep(Duration::from_millis(200)).await; + let mut again = connect(&gw, "rt").await; + until_type(&mut again, "session.created").await; +} + +/// TC-LIFE-11 — drain: `server_shutdown` then 1012. +#[tokio::test] +async fn tc_life_11_drain() { + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + gw.state.shutdown.cancel(); + let (errors, code) = until_close(&mut c).await; + assert_eq!(errors, vec!["server_shutdown"]); + assert_eq!(code, Some(1012)); +} + +/// TC-LIFE-12 — the FRD-022 connection slot returns to zero after sessions and refusals. +#[tokio::test] +async fn tc_life_12_connection_slots_return_to_zero() { + let (_v, gw) = hs_gateway().await; + for _ in 0..20 { + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + c.close(None).await.unwrap(); + let _ = until_close(&mut c).await; + let _ = connect_with(&gw, "model=nope", Auth::Bearer(KEY), &[], &[]).await; + } + for _ in 0..100 { + if gw.state.ws_connection_count() == 0 { + return; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("{} connection slots leaked", gw.state.ws_connection_count()); +} + +// ============================================================================================= +// Metering — TC-MET +// ============================================================================================= + +/// TC-MET-01 / 02 🔒 / 03 🔒 / 07 — three responses, three root turns with their own traces, each +/// linked to the session; the session span's totals are the sum of its turns. +#[tokio::test] +async fn tc_met_01_02_03_07_per_response_turns_and_the_session_record() { + let cap = Capture::install(); + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + for _ in 0..3 { + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; + } + c.close(None).await.unwrap(); + let _ = until_close(&mut c).await; + + let turns = cap.wait_for("voice.turn", 3).await; + let session = cap.wait_for("voice.session", 1).await.remove(0); + let mut traces = std::collections::HashSet::new(); + let mut indices = Vec::new(); + for t in &turns { + assert_eq!( + t.parent_span_id, + SpanId::INVALID, + "TC-MET-02: a turn has no parent" + ); + assert!( + traces.insert(t.span_context.trace_id()), + "TC-MET-02: turns share a trace" + ); + assert_ne!(t.span_context.trace_id(), session.span_context.trace_id()); + assert!( + t.links + .links + .iter() + .any(|l| l.span_context.span_id() == session.span_context.span_id()), + "TC-MET-02: no link to the session span" + ); + assert_eq!( + text(t, "bud.voice.capability").as_deref(), + Some("realtime_session") + ); + assert_eq!(text(t, "bud.voice.transport").as_deref(), Some("websocket")); + assert_eq!( + text(t, "bud.voice.rt.component").as_deref(), + Some("response") + ); + assert_eq!(text(t, "bud.endpoint_id").as_deref(), Some(EP_HS)); + assert_eq!(text(t, "bud.voice.endpoint_name").as_deref(), Some("rt")); + assert_eq!(text(t, "bud.project_id").as_deref(), Some(PROJECT)); + assert_eq!(text(t, "bud.model_id").as_deref(), Some(MODEL_ID)); + assert_eq!(text(t, "bud.api_key_id").as_deref(), Some(API_KEY_ID)); + assert_eq!(text(t, "bud.user_id").as_deref(), Some(USER)); + assert_eq!( + text(t, "bud.voice.vendor_session_id").as_deref(), + Some("sess_vendor_1") + ); + assert_eq!(text(t, "bud.voice.rt.vendor").as_deref(), Some("openai")); + assert_eq!( + text(t, "bud.voice.rt.model").as_deref(), + Some("gpt-realtime-2.1") + ); + assert_eq!(number(t, "bud.voice.rt.input_text_tokens"), Some(119.0)); + assert_eq!(number(t, "bud.voice.rt.cached_text_tokens"), Some(64.0)); + assert_eq!(number(t, "bud.voice.rt.output_audio_tokens"), Some(91.0)); + let cost = number(t, "bud.voice.cost").expect("TC-MET-03: priced"); + assert!((cost - 0.0072056).abs() < 1e-10, "TC-MET-03 cost {cost}"); + assert_eq!(text(t, "bud.voice.pricing_unit").as_deref(), Some("token")); + indices.push(number(t, "bud.voice.turn_index").unwrap() as u64); + } + indices.sort_unstable(); + assert_eq!(indices, vec![0, 1, 2], "TC-MET-01"); + let session_ids: std::collections::HashSet<_> = turns + .iter() + .map(|t| text(t, "bud.voice.session_id").unwrap()) + .collect(); + assert_eq!(session_ids.len(), 1); + + // TC-MET-07 + assert_eq!(number(&session, "bud.voice.session.turns"), Some(3.0)); + assert_eq!( + text(&session, "bud.voice.session.end_reason").as_deref(), + Some("client_close") + ); + assert_eq!( + number(&session, "bud.voice.rt.input_text_tokens"), + Some(357.0) + ); + let total = number(&session, "bud.voice.cost").unwrap(); + assert!((total - 3.0 * 0.0072056).abs() < 1e-9, "{total}"); + assert!(number(&session, "bud.voice.session.duration_ms").unwrap() > 0.0); + assert_eq!( + text(&session, "bud.voice.rt.session_type").as_deref(), + Some("realtime") + ); +} + +/// TC-MET-06 — an input transcription is its own priced turn. +#[tokio::test] +async fn tc_met_06_transcription_turn() { + let cap = Capture::install(); + let vendor = MockVendor::start(Behaviour { + transcription_usage: Some(json!({"type": "duration", "seconds": 12})), + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000206", + rt_entry(&vendor, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + send(&mut c, json!({"type": "input_audio_buffer.commit"})).await; + until_type( + &mut c, + "conversation.item.input_audio_transcription.completed", + ) + .await; + let t = cap.wait_for("voice.turn", 1).await.remove(0); + assert_eq!( + text(&t, "bud.voice.rt.component").as_deref(), + Some("input_transcription") + ); + assert_eq!(number(&t, "bud.voice.billed_seconds"), Some(12.0)); + assert!((number(&t, "bud.voice.cost").unwrap() - 0.0006).abs() < 1e-12); +} + +/// TC-MET-08 — the client's TCP dies after two responses: both turns are exported, and the +/// session record says how it ended. +#[tokio::test] +async fn tc_met_08_drop_mid_session_keeps_the_turns() { + let cap = Capture::install(); + let (_v, gw) = hs_gateway().await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + for _ in 0..2 { + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; + } + drop(c); + let turns = cap.wait_for("voice.turn", 2).await; + assert_eq!(turns.len(), 2); + let session = cap.wait_for("voice.session", 1).await.remove(0); + assert_eq!( + text(&session, "bud.voice.session.end_reason").as_deref(), + Some("client_close") + ); +} + +/// TC-XL-07 (duration billing) — a minute price bills 1 s segments here (shortened), with the +/// final partial one. +#[tokio::test] +async fn duration_priced_sessions_bill_segments() { + let cap = Capture::install(); + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000207", + rt_entry( + &vendor, + json!({"pricing": {"unit": "minute", "cost_per_unit": 0.06, "per_units": 1}}), + ), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "rt").await; + until_type(&mut c, "session.created").await; + keep_alive(&mut c, Duration::from_millis(2500)).await; + c.close(None).await.unwrap(); + let _ = until_close(&mut c).await; + let session = cap.wait_for("voice.session", 1).await.remove(0); + let segments: Vec = cap + .spans() + .into_iter() + .filter(|s| { + s.name == "voice.turn" + && text(s, "bud.voice.rt.component").as_deref() == Some("duration_segment") + }) + .collect(); + assert_eq!(segments.len(), 3, "1 s + 1 s + the partial remainder"); + let billed: f64 = segments + .iter() + .map(|s| number(s, "bud.voice.billed_seconds").unwrap()) + .sum(); + assert!((billed - number(&session, "bud.voice.billed_seconds").unwrap()).abs() < 1e-9); + assert!(billed > 2.4 && billed < 3.5, "{billed}"); + for s in &segments { + assert_eq!(text(s, "bud.voice.pricing_unit").as_deref(), Some("minute")); + } +} + +// ============================================================================================= +// Credentials never logged — TC-SEC-07 +// ============================================================================================= + +/// TC-SEC-07 🔒 — the Bud key and the vendor key appear in no log line and no span. +#[tokio::test] +async fn tc_sec_07_credentials_never_reach_logs_or_spans() { + let cap = Capture::install(); + let (_v, gw) = hs_gateway().await; + for auth in [Auth::Bearer(KEY), Auth::ApiKey(KEY), Auth::Subprotocol(KEY)] { + let mut c = connect_with(&gw, "model=rt", auth, &[], &[]) + .await + .unwrap() + .ws; + until_type(&mut c, "session.created").await; + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; + c.close(None).await.unwrap(); + let _ = until_close(&mut c).await; + } + cap.wait_for("voice.session", 3).await; + let logs = cap.logs(); + assert!(!logs.is_empty(), "the log capture is live"); + for secret in [KEY, VENDOR_KEY, "sk-process-canary"] { + assert!(!logs.contains(secret), "{secret} appeared in a log line"); + for s in cap.spans() { + for kv in s.attributes.iter() { + assert!( + !kv.value.as_str().contains(secret), + "{secret} on span {}", + s.name + ); + } + } + } +} + +// ============================================================================================= +// Client secrets — TC-EK +// ============================================================================================= + +async fn mint(gw: &Gateway, bearer: &str, body: Json) -> (u16, Json) { + let resp = reqwest::Client::new() + .post(format!("http://{}/v1/realtime/client_secrets", gw.addr)) + .bearer_auth(bearer) + .json(&body) + .send() + .await + .unwrap(); + let status = resp.status().as_u16(); + (status, resp.json().await.unwrap_or(Json::Null)) +} + +async fn ek_gateway() -> (MockVendor, Gateway) { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ + ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000301", + rt_entry(&vendor, json!({})), + ), + ep( + "rt2", + "a1a1a1a1-0000-4000-8000-000000000302", + rt_entry(&vendor, json!({})), + ), + ], + other_endpoints: vec![ep( + "theirs", + "a1a1a1a1-0000-4000-8000-000000000303", + rt_entry(&vendor, json!({})), + )], + jwt_subjects: vec!["sub-ek"], + ..Default::default() + }) + .await; + (vendor, gw) +} + +/// TC-EK-01 / 05 — mint, then connect with the secret; turns are attributed to the parent. +#[tokio::test] +async fn tc_ek_01_05_mint_and_connect() { + let cap = Capture::install(); + let (_v, gw) = ek_gateway().await; + let before = now(); + let (s, body) = mint(&gw, KEY, json!({"session": {"model": "rt"}})).await; + assert_eq!(s, 200, "{body}"); + let value = body["value"].as_str().unwrap().to_string(); + assert!(value.starts_with("ek_bud_")); + let exp = body["expires_at"].as_u64().unwrap(); + assert!( + (before + 600..=now() + 600).contains(&exp), + "default TTL 600: {exp}" + ); + assert_eq!(body["session"]["model"], "rt"); + + let c = connect_with(&gw, "model=rt", Auth::Subprotocol(&value), &[], &[]) + .await + .unwrap(); + assert_eq!(c.protocol.as_deref(), Some("realtime"), "TC-EK-15"); + let mut ws = c.ws; + until_type(&mut ws, "session.created").await; + send(&mut ws, json!({"type": "response.create"})).await; + until_type(&mut ws, "response.done").await; + let t = cap.wait_for("voice.turn", 1).await.remove(0); + assert_eq!(text(&t, "bud.api_key_id").as_deref(), Some(API_KEY_ID)); + assert_eq!(text(&t, "bud.project_id").as_deref(), Some(PROJECT)); +} + +#[tokio::test] +async fn tc_ek_02_03_04_08_mint_and_connect_refusals() { + let (_v, gw) = ek_gateway().await; + for (body, status) in [ + ( + json!({"expires_after": {"anchor": "created_at", "seconds": 5}, "session": {"model": "rt"}}), + 400, + ), + ( + json!({"expires_after": {"anchor": "created_at", "seconds": 7201}, "session": {"model": "rt"}}), + 400, + ), + (json!({"session": {}}), 400), + (json!({"session": {"model": "theirs"}}), 403), + ] { + let (s, b) = mint(&gw, KEY, body.clone()).await; + assert_eq!(s, status, "{body} → {b}"); + } + let (_, ok) = mint(&gw, KEY, json!({"session": {"model": "rt"}})).await; + let secret = ok["value"].as_str().unwrap().to_string(); + // TC-EK-04: no chaining. + let (s, _) = mint(&gw, &secret, json!({"session": {"model": "rt"}})).await; + assert_eq!(s, 401); + // TC-EK-08: a secret for `rt` cannot open `rt2`. + let Err((s, body, _)) = + connect_with(&gw, "model=rt2", Auth::Subprotocol(&secret), &[], &[]).await + else { + panic!("a secret opened a deployment it was not minted for"); + }; + assert_eq!(s, 403); + assert_eq!(body["error"]["code"], "model_mismatch"); +} + +/// TC-EK-09 🔒 / TC-EK-14 🔒 — the parent key revoked: connect refused, a live session closed. +#[tokio::test] +async fn tc_ek_09_14_parent_revocation() { + let (_v, gw) = ek_gateway().await; + let (_, ok) = mint(&gw, KEY, json!({"session": {"model": "rt"}})).await; + let secret = ok["value"].as_str().unwrap().to_string(); + let mut live = connect_with(&gw, "model=rt", Auth::Subprotocol(&secret), &[], &[]) + .await + .unwrap() + .ws; + until_type(&mut live, "session.created").await; + + // TC-EK-14: the parent is exactly the snapshot key — deleting `api_key:{hash_api_key(K)}` is + // what revokes the secret. + let key = format!("api_key:{}", bud_auth::hash_api_key(KEY)); + gw.store.remove(&key); + gw.plane + .on_key_event(&key, bud_auth::KeyEvent::Del) + .await + .unwrap(); + + let (errors, code) = until_close(&mut live).await; + assert_eq!(errors, vec!["session_revoked"]); + assert_eq!(code, Some(1008)); + let Err((s, _, _)) = connect_with(&gw, "model=rt", Auth::Subprotocol(&secret), &[], &[]).await + else { + panic!("a secret of a revoked parent connected"); + }; + assert_eq!(s, 401); +} + +/// TC-EK-12 🔒 — a JWT parent caps the lifetime; under 10 s left, minting is refused. +#[tokio::test] +async fn tc_ek_12_jwt_parent_caps_the_lifetime() { + let (_v, gw) = ek_gateway().await; + let token = jwt("sub-ek", 120); + let (s, body) = mint(&gw, &token, json!({"expires_after": {"anchor": "created_at", "seconds": 600}, "session": {"model": "rt"}})).await; + assert_eq!(s, 200, "{body}"); + let exp = body["expires_at"].as_u64().unwrap(); + assert!( + exp <= now() + 121 && exp >= now() + 110, + "capped at the JWT's exp: {exp}" + ); + + let short = jwt("sub-ek", 5); + let (s, body) = mint(&gw, &short, json!({"session": {"model": "rt"}})).await; + assert_eq!(s, 401); + assert_eq!(body["error"]["code"], "credential_expiring"); +} + +/// TC-EK-17 — a JWT parent removed from the project: connect refused. +#[tokio::test] +async fn tc_ek_17_jwt_parent_revoked() { + let (_v, gw) = ek_gateway().await; + let (_, ok) = mint( + &gw, + &jwt("sub-ek", 300), + json!({"session": {"model": "rt"}}), + ) + .await; + let secret = ok["value"].as_str().unwrap().to_string(); + gw.store.set( + "user_projects:sub-ek", + &json!({"user_id": USER, "projects": []}).to_string(), + ); + gw.plane + .on_key_event("user_projects:sub-ek", bud_auth::KeyEvent::Set) + .await + .unwrap(); + let Err((s, _, _)) = connect_with(&gw, "model=rt", Auth::Subprotocol(&secret), &[], &[]).await + else { + panic!("connected on a revoked JWT parent"); + }; + assert_eq!(s, 401); +} + +/// No keys configured: the mint route answers 501. +#[tokio::test] +async fn client_secrets_unconfigured_is_501() { + let vendor = MockVendor::start(Behaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "rt", + "a1a1a1a1-0000-4000-8000-000000000304", + rt_entry(&vendor, json!({})), + )], + client_secret_keys: None, + ..Default::default() + }) + .await; + let (s, body) = mint(&gw, KEY, json!({"session": {"model": "rt"}})).await; + assert_eq!(s, 501); + assert_eq!(body["error"]["code"], "client_secrets_not_configured"); +} diff --git a/gateway/tests/voice_span_contract.json b/gateway/tests/voice_span_contract.json index f3d14e6c..c11cf462 100644 --- a/gateway/tests/voice_span_contract.json +++ b/gateway/tests/voice_span_contract.json @@ -1,40 +1,48 @@ { - "_comment": "Span attributes VoiceTurnFact reads. WaaV EMITS these; budmetrics READS them. Neither repo can see the other at build time, so both assert against this file. A name changed on one side only is invisible: the column simply stays NULL, which looks exactly like a feature nobody uses.", - "_generated_from": "services/budmetrics/budmetrics/observability/voice_turn_fact_ddl.py", + "_comment": "Span attributes VoiceTurnFact (roles turn, leg: voice.turn) and VoiceSessionFact (role session: voice.session) read. WaaV EMITS these; budmetrics READS them. Neither repo can see the other at build time, so both assert against this file. A name changed on one side only is invisible: the column simply stays NULL, which looks exactly like a feature nobody uses.", + "_generated_from": [ + "services/budmetrics/budmetrics/observability/voice_turn_fact_ddl.py", + "services/budmetrics/budmetrics/observability/voice_session_fact_ddl.py" + ], "attributes": [ { "column": "api_key_id", "attribute": "bud.api_key_id", "roles": [ - "turn" + "turn", + "session" ] }, { "column": "api_key_project_id", "attribute": "bud.api_key_project_id", "roles": [ - "turn" + "turn", + "session" ] }, { "column": "endpoint_id", "attribute": "bud.endpoint_id", "roles": [ - "turn" + "turn", + "session" ] }, { "column": "model_id", "attribute": "bud.model_id", "roles": [ - "turn" + "turn", + "session" ] }, { "column": "project_id", "attribute": "bud.project_id", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -48,7 +56,8 @@ "column": "user_id", "attribute": "bud.user_id", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -72,11 +81,20 @@ "turn" ] }, + { + "column": "billed_seconds", + "attribute": "bud.voice.billed_seconds", + "roles": [ + "turn", + "session" + ] + }, { "column": "capability", "attribute": "bud.voice.capability", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -90,7 +108,8 @@ "column": "cost", "attribute": "bud.voice.cost", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -104,7 +123,8 @@ "column": "endpoint_name", "attribute": "bud.voice.endpoint_name", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -153,7 +173,8 @@ "column": "pricing_unit", "attribute": "bud.voice.pricing_unit", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -170,6 +191,126 @@ "turn" ] }, + { + "column": "cached_audio_tokens", + "attribute": "bud.voice.rt.cached_audio_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "cached_image_tokens", + "attribute": "bud.voice.rt.cached_image_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "cached_text_tokens", + "attribute": "bud.voice.rt.cached_text_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "rt_component", + "attribute": "bud.voice.rt.component", + "roles": [ + "turn" + ] + }, + { + "column": "input_audio_tokens", + "attribute": "bud.voice.rt.input_audio_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "input_image_tokens", + "attribute": "bud.voice.rt.input_image_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "input_text_tokens", + "attribute": "bud.voice.rt.input_text_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "model", + "attribute": "bud.voice.rt.model", + "roles": [ + "session" + ] + }, + { + "column": "rt_model", + "attribute": "bud.voice.rt.model", + "roles": [ + "turn" + ] + }, + { + "column": "output_audio_tokens", + "attribute": "bud.voice.rt.output_audio_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "output_text_tokens", + "attribute": "bud.voice.rt.output_text_tokens", + "roles": [ + "turn", + "session" + ] + }, + { + "column": "response_id", + "attribute": "bud.voice.rt.response_id", + "roles": [ + "turn" + ] + }, + { + "column": "response_status", + "attribute": "bud.voice.rt.response_status", + "roles": [ + "turn" + ] + }, + { + "column": "session_type", + "attribute": "bud.voice.rt.session_type", + "roles": [ + "session" + ] + }, + { + "column": "rt_vendor", + "attribute": "bud.voice.rt.vendor", + "roles": [ + "turn" + ] + }, + { + "column": "vendor", + "attribute": "bud.voice.rt.vendor", + "roles": [ + "session" + ] + }, { "column": "sample_rate", "attribute": "bud.voice.sample_rate", @@ -184,11 +325,40 @@ "turn" ] }, + { + "column": "close_code", + "attribute": "bud.voice.session.close_code", + "roles": [ + "session" + ] + }, + { + "column": "duration_ms", + "attribute": "bud.voice.session.duration_ms", + "roles": [ + "session" + ] + }, + { + "column": "end_reason", + "attribute": "bud.voice.session.end_reason", + "roles": [ + "session" + ] + }, + { + "column": "turns", + "attribute": "bud.voice.session.turns", + "roles": [ + "session" + ] + }, { "column": "session_id", "attribute": "bud.voice.session_id", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -251,7 +421,8 @@ "column": "transport", "attribute": "bud.voice.transport", "roles": [ - "turn" + "turn", + "session" ] }, { @@ -303,6 +474,13 @@ "turn" ] }, + { + "column": "unpriced_components", + "attribute": "bud.voice.unpriced_components", + "roles": [ + "turn" + ] + }, { "column": "vendor_request_id", "attribute": "bud.voice.vendor_request_id", @@ -310,6 +488,14 @@ "turn" ] }, + { + "column": "vendor_session_id", + "attribute": "bud.voice.vendor_session_id", + "roles": [ + "turn", + "session" + ] + }, { "column": "vendor_status_code", "attribute": "bud.voice.vendor_status_code", From c85e5d364cf0fa7adc066920c6a58ae47536b014 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 05:01:20 +0530 Subject: [PATCH 05/17] fix(gateway): clippy findings in the /v1/realtime relay Two let-else blocks become ? and to_vendor becomes send_to_vendor (a to_* method taking &mut self trips wrong_self_convention under -D warnings). Co-Authored-By: Claude Opus 5.5 --- gateway/src/handlers/openai_realtime/policy.rs | 8 ++------ gateway/src/handlers/openai_realtime/session.rs | 6 +++--- 2 files changed, 5 insertions(+), 9 deletions(-) diff --git a/gateway/src/handlers/openai_realtime/policy.rs b/gateway/src/handlers/openai_realtime/policy.rs index fb8d77fc..0c5e540f 100644 --- a/gateway/src/handlers/openai_realtime/policy.rs +++ b/gateway/src/handlers/openai_realtime/policy.rs @@ -171,9 +171,7 @@ fn apply_session_update( rules: &ClientRules, event_id: Option<&str>, ) -> Option> { - let Some(session) = event.get_mut("session").and_then(Value::as_object_mut) else { - return None; - }; + let session = event.get_mut("session").and_then(Value::as_object_mut)?; // Routine: SDKs send both on every update. `model` is the deployment's (D-3); `tracing` // would open a trace in the vendor org every project on the credential shares (D-11). session.remove("model"); @@ -217,9 +215,7 @@ fn apply_response_create( rules: &ClientRules, event_id: Option<&str>, ) -> Option> { - let Some(response) = event.get_mut("response").and_then(Value::as_object_mut) else { - return None; - }; + let response = event.get_mut("response").and_then(Value::as_object_mut)?; response.remove("model"); if let Some(refused) = check_overridable(response, "response", rules, event_id) { return Some(refused); diff --git a/gateway/src/handlers/openai_realtime/session.rs b/gateway/src/handlers/openai_realtime/session.rs index 0ee71f17..64b3cdc8 100644 --- a/gateway/src/handlers/openai_realtime/session.rs +++ b/gateway/src/handlers/openai_realtime/session.rs @@ -593,7 +593,7 @@ impl Relay<'_> { } } - async fn to_vendor(&mut self, text: String) -> Result<(), End> { + async fn send_to_vendor(&mut self, text: String) -> Result<(), End> { let started = std::time::Instant::now(); match tokio::time::timeout( self.timings.upstream_send, @@ -651,7 +651,7 @@ impl Relay<'_> { .await; } } - self.to_vendor(text).await + self.send_to_vendor(text).await } async fn release_held(&mut self) -> Result<(), End> { @@ -731,7 +731,7 @@ impl Relay<'_> { &event_id, ) { Some(update) => { - self.to_vendor(update).await?; + self.send_to_vendor(update).await?; self.awaiting_defaults = Some(event_id); self.hold_deadline = Some(Instant::now() + self.timings.hold); Ok(()) From cc3c559cd018d3f79946bea9777d7e0e87a31321 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 05:01:29 +0530 Subject: [PATCH 06/17] feat(gateway): /ws legs address Bud deployments (FRD-023 RT6) Under the Bud control plane a /ws session's legs are deployments, resolved through the caller's own allowlist, exactly as the REST routes resolve them. - STT/TTS legs (WP-RT6.1): stt_config.model / tts_config.model name a transcription / text-to-speech deployment; vendor, model, credential, api_base, voice and settings come from voice_table (request > deployment). provider is optional and ignored. Refused by name: provider-only legs (deployment_required), unreachable names (model_not_found), upload-only STT (unsupported_deployment), deployments that cannot reach their vendor (deployment_misconfigured, e.g. AWS without its key pair), and a client's own api_key (client_key_not_accepted). Client extras are replaced by the deployment's: Groq and Azure Speech read a destination from them. - One admission per leg deployment, held for the session; a capped leg closes 1013 and releases the other leg (MessageRoute::CloseWith). - Voice agent LLM leg (WP-RT6.2): conversation_config.model names a Bud chat deployment reached through WAAV_LLM_BASE_URL with the caller's credential, read per call; an auth message on an authenticated session refreshes it (same API key or user only); an expired credential fails the turn with auth_expired and the session stays open. Conversation turns are attributed. - DAG templates (WP-RT6.3): TTS nodes bind to deployments (admitted, metered, revalidated); LLM/translate nodes go to the Bud gateway as the caller; realtime nodes are refused until RT7. - stt.streaming (WP-RT6.4) is modelled in bud-auth and applied on /ws only. - Metering and revocation (WP-RT6.5): a voice.turn per vendor final (audio seconds since the previous one, the remainder at close) and per speak (characters), fully attributed; revalidation every 30 s closes 1008. Tests: TC-WS-01..13 (lib + conversation_loop); every guard seen red with its check removed. Live: 21/21 on pde-ditto against Deepgram, ElevenLabs and a chat deployment through budgateway. Co-Authored-By: Claude Opus 5.5 --- bud-auth/src/endpoint_config.rs | 49 +- gateway/src/auth/context.rs | 50 + gateway/src/auth/mod.rs | 2 +- gateway/src/core/conversation/mod.rs | 107 +- gateway/src/core/voice_manager/manager.rs | 59 +- gateway/src/dag/compiler.rs | 23 +- gateway/src/dag/definition.rs | 7 + gateway/src/dag/nodes/llm.rs | 31 +- gateway/src/dag/nodes/mod.rs | 4 +- gateway/src/dag/nodes/provider.rs | 75 +- gateway/src/handlers/openai_audio.rs | 34 +- gateway/src/handlers/ws/audio_handler.rs | 7 +- gateway/src/handlers/ws/bud_legs.rs | 1694 +++++++++++++++++++++ gateway/src/handlers/ws/config.rs | 19 +- gateway/src/handlers/ws/config_handler.rs | 567 ++++++- gateway/src/handlers/ws/handler.rs | 54 +- gateway/src/handlers/ws/messages.rs | 6 + gateway/src/handlers/ws/mod.rs | 1 + gateway/src/handlers/ws/processor.rs | 202 ++- gateway/src/handlers/ws/state.rs | 17 + gateway/src/middleware/auth.rs | 8 + gateway/src/test_support.rs | 36 + gateway/tests/conversation_loop.rs | 326 ++++ 23 files changed, 3268 insertions(+), 110 deletions(-) create mode 100644 gateway/src/handlers/ws/bud_legs.rs diff --git a/bud-auth/src/endpoint_config.rs b/bud-auth/src/endpoint_config.rs index 31e83fc4..00192e8c 100644 --- a/bud-auth/src/endpoint_config.rs +++ b/bud-auth/src/endpoint_config.rs @@ -117,9 +117,9 @@ pub struct TtsSettings { /// Transcription defaults for a deployment. /// -/// The five streaming-only canonical features are absent by construction: they need a continuous -/// stream, budapp refuses them at publish naming the transport, and a field here would suggest -/// otherwise to the next person reading this struct. +/// The five streaming-only canonical features are not fields of their own: they need a continuous +/// stream, so budapp refuses them at the top of `stt` and publishes them under `stt.streaming` +/// ([`SttStreaming`]), which only the `/ws` transport applies (FRD-023 WP-RT6.4, FR-WS-4). #[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)] pub struct SttSettings { // --- the four the handler used to hardcode, plus model and prompt --- @@ -170,6 +170,26 @@ pub struct SttSettings { /// request and changes the audio the vendor bills against, so it defaults off. #[serde(default)] pub noise_suppression: Option, + + /// The streaming-only features, applied on `/ws` and ignored by the prerecorded upload. + #[serde(default)] + pub streaming: Option, +} + +/// `stt.streaming`: the five canonical features that need a continuous audio stream +/// (FRD-023 WP-RT6.4). budapp validates the closed key set and the millisecond ranges. +#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)] +pub struct SttStreaming { + #[serde(default)] + pub interim_results: Option, + #[serde(default)] + pub vad_events: Option, + #[serde(default)] + pub endpointing_ms: Option, + #[serde(default)] + pub utterance_end_ms: Option, + #[serde(default)] + pub speech_begin_event: Option, } /// Translation defaults. @@ -386,6 +406,7 @@ const KNOWN_STT: &[&str] = &[ "alternatives", "sentiment", "noise_suppression", + "streaming", ]; const KNOWN_TRANSLATION: &[&str] = &["target_languages", "translate_to_english", "partials"]; @@ -539,6 +560,28 @@ mod tests { ); } + #[test] + fn streaming_features_are_modelled_under_stt_streaming() { + // TC-WS-11: budapp publishes the five streaming-only features here (WP-RT6.4). + let settings = parse( + r#"{"stt": {"diarization": true, "streaming": {"interim_results": true, + "vad_events": false, "endpointing_ms": 300, "utterance_end_ms": 1000, + "speech_begin_event": true}}}"#, + ); + let stt = settings.stt.expect("stt parses"); + assert_eq!(stt.diarization, Some(true)); + let streaming = stt.streaming.expect("stt.streaming is modelled"); + assert_eq!(streaming.interim_results, Some(true)); + assert_eq!(streaming.vad_events, Some(false)); + assert_eq!(streaming.endpointing_ms, Some(300)); + assert_eq!(streaming.utterance_end_ms, Some(1000)); + assert_eq!(streaming.speech_begin_event, Some(true)); + assert!( + KNOWN_STT.contains(&"streaming"), + "no unmodelled-key warning for it" + ); + } + #[test] fn an_unknown_field_is_kept_not_fatal() { // TC-CFG-06. This is the property that lets budapp and WaaV ship in either order. diff --git a/gateway/src/auth/context.rs b/gateway/src/auth/context.rs index eda16c1a..6f9e317f 100644 --- a/gateway/src/auth/context.rs +++ b/gateway/src/auth/context.rs @@ -37,6 +37,40 @@ pub struct Auth { // pub metadata: Option, } +/// A socket session's raw credential, kept so the session can act AS its caller (FRD-023 RT6): +/// resolve its legs' deployments, reach budgateway for the voice agent's LLM leg, and be +/// revalidated while it lives. +/// +/// Shared (`Clone` shares the value), because an `auth` message REFRESHES it mid-session — a +/// Keycloak token lives about five minutes, a voice call longer — and the LLM leg must send the +/// current token on its next call. `Debug` never prints it. +#[derive(Clone)] +pub struct SessionCredential(std::sync::Arc>); + +impl SessionCredential { + pub fn new(raw: impl Into) -> Self { + Self(std::sync::Arc::new(std::sync::RwLock::new(raw.into()))) + } + + /// The credential as it is now. + pub fn current(&self) -> String { + self.0.read().map(|g| g.clone()).unwrap_or_default() + } + + /// Replace it (an `auth` refresh); every holder sees the new value. + pub fn replace(&self, raw: impl Into) { + if let Ok(mut g) = self.0.write() { + *g = raw.into(); + } + } +} + +impl std::fmt::Debug for SessionCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("SessionCredential([redacted])") + } +} + impl Auth { /// Create a new Auth with the given id pub fn new(id: impl Into) -> Self { @@ -120,6 +154,22 @@ impl Auth { mod tests { use super::*; + #[test] + fn a_session_credential_is_shared_refreshable_and_never_printed() { + let a = SessionCredential::new("bud_secret_one"); + let b = a.clone(); + b.replace("bud_secret_two"); + assert_eq!( + a.current(), + "bud_secret_two", + "a refresh reaches every holder" + ); + assert!( + !format!("{a:?}").contains("bud_secret"), + "Debug must redact" + ); + } + #[test] fn test_normalize_room_name_basic_prefix() { let auth = Auth::new("project1"); diff --git a/gateway/src/auth/mod.rs b/gateway/src/auth/mod.rs index 91b0282d..f4b74b13 100644 --- a/gateway/src/auth/mod.rs +++ b/gateway/src/auth/mod.rs @@ -8,7 +8,7 @@ pub mod jwt; // Re-export commonly used items pub use api_secret::match_api_secret_id; pub use client::AuthClient; -pub use context::Auth; +pub use context::{Auth, SessionCredential}; pub use jwt::{ AuthClaims, AuthPayload, detect_algorithm, filter_headers, load_private_key, sign_auth_request, sign_auth_request_with_key, diff --git a/gateway/src/core/conversation/mod.rs b/gateway/src/core/conversation/mod.rs index 2e4166bd..7912934d 100644 --- a/gateway/src/core/conversation/mod.rs +++ b/gateway/src/core/conversation/mod.rs @@ -67,6 +67,14 @@ pub struct ConversationConfig { pub system_prompt: Option, /// API key (literal or `${ENV_VAR}`); falls back to `OPENAI_API_KEY`. pub api_key: Option, + /// FRD-023 RT6: `base_url` is the OPERATOR's (budgateway, `WAAV_LLM_BASE_URL`), not a client's, + /// so it is not SSRF-validated — it is an in-cluster address by design. Set only by the server. + pub server_llm_endpoint: bool, + /// FRD-023 RT6: the caller's live credential, sent on every LLM call in place of `api_key`, so + /// an `auth` refresh reaches the next call and a JWT caller's leg outlives one token. + pub credential: Option, + /// FRD-023 RT6: who each turn is attributed to, recorded on its `voice.turn` span. + pub attribution: Option, /// Sampling temperature. pub temperature: Option, /// Max tokens per completion. @@ -188,6 +196,9 @@ impl Default for ConversationConfig { model: "gpt-4o-mini".to_string(), system_prompt: None, api_key: None, + server_llm_endpoint: false, + credential: None, + attribution: None, temperature: None, max_tokens: None, streaming: true, @@ -697,7 +708,31 @@ impl std::fmt::Debug for ConversationOrchestrator { } } +/// A Bud-mode session's caller, for its conversation turns' `voice.turn` records (FRD-023 RT6). +/// +/// Attribution only: the legs' costs are their own records (the transcription and speech +/// deployments', and budgateway's for the chat deployment), so a turn carrying them would count +/// them twice. +#[derive(Debug, Clone, Default)] +pub struct TurnAttribution { + pub project_id: Option, + pub api_key_id: Option, + pub api_key_project_id: Option, + pub user_id: Option, + /// The chat deployment the agent answers with. + pub endpoint_name: Option, +} + impl ConversationOrchestrator { + /// The key for the next LLM call: the session's live credential when it has one (FRD-023 RT6, + /// a Bud caller reaching budgateway), else the configured `api_key`. + fn llm_key(&self) -> Option { + match &self.config.credential { + Some(credential) => Some(credential.current()), + None => self.config.api_key.clone(), + } + } + /// Create a new orchestrator for `session_id`. /// /// Validates the LLM `base_url` for SSRF (resolve-then-validate, with the @@ -708,7 +743,10 @@ impl ConversationOrchestrator { config: ConversationConfig, voice_manager: Arc, ) -> Result { - validate_llm_url(&config.base_url).map_err(ConversationOrchestratorError::InvalidLlmUrl)?; + if !config.server_llm_endpoint { + validate_llm_url(&config.base_url) + .map_err(ConversationOrchestratorError::InvalidLlmUrl)?; + } // S1/S2: the reasoning tier's base_url is ALSO client-supplied — validate // it for SSRF before it is ever used for a request. if let Some(rb) = &config.reasoning_base_url { @@ -865,9 +903,17 @@ impl ConversationOrchestrator { } else { DEFAULT_REASONING_BUDGET_MS }; + // FRD-023 RT6: both tiers are Bud chat deployments behind one budgateway, reached + // with the session's own credential; otherwise the tier resolves its own key. + let fb_key = self.config.credential.as_ref().map(|c| c.current()); let fb_result = tokio::time::timeout( Duration::from_millis(bound_ms), - fallback.continue_from_history(&self.session_id, None, &fb_token, None), + fallback.continue_from_history( + &self.session_id, + fb_key.as_deref(), + &fb_token, + None, + ), ) .await; match fb_result { @@ -1042,12 +1088,7 @@ impl ConversationOrchestrator { let epoch = self.voice_manager.clear_epoch(); let result = self .llm - .continue_from_history( - &self.session_id, - self.config.api_key.as_deref(), - &token, - None, - ) + .continue_from_history(&self.session_id, self.llm_key().as_deref(), &token, None) .await; match result { Ok(resp) if !resp.content.trim().is_empty() => { @@ -1200,6 +1241,20 @@ impl ConversationOrchestrator { crate::observability::voice_attrs::turn::TRANSCRIPT, transcript, ); + if let Some(a) = &self.config.attribution { + use crate::observability::voice_attrs::turn; + for (key, value) in [ + (turn::PROJECT_ID, &a.project_id), + (turn::API_KEY_ID, &a.api_key_id), + (turn::API_KEY_PROJECT_ID, &a.api_key_project_id), + (turn::USER_ID, &a.user_id), + (turn::ENDPOINT_NAME, &a.endpoint_name), + ] { + if let Some(v) = value.as_deref().filter(|v| !v.is_empty()) { + span.record(key, v); + } + } + } tracing::Instrument::instrument(self.run_turn_inner(transcript), span).await } @@ -1470,10 +1525,11 @@ impl ConversationOrchestrator { let reasoner_token = token.child_token(); let req_start = crate::core::observability::now_monotonic_ns(); let budget_ns = budget_ms.saturating_mul(1_000_000); + let llm_key = self.llm_key(); let complete_fut = llm.complete( &self.session_id, transcript, - self.config.api_key.as_deref(), + llm_key.as_deref(), &reasoner_token, on_token, ); @@ -1505,7 +1561,7 @@ impl ConversationOrchestrator { llm.complete( &self.session_id, transcript, - self.config.api_key.as_deref(), + self.llm_key().as_deref(), &token, on_token, ) @@ -1565,7 +1621,7 @@ impl ConversationOrchestrator { ®istry, &self.session_id, response, - self.config.api_key.as_deref(), + self.llm_key().as_deref(), &token, tool_opts, ) @@ -1764,7 +1820,7 @@ impl ConversationOrchestrator { let llm = Arc::clone(&self.llm); let session_id = self.session_id.clone(); let target_tokens = self.config.summarize_target_tokens; - let api_key = self.config.api_key.clone(); + let api_key = self.llm_key(); tokio::spawn(async move { let cfg = crate::core::llm::SummaryConfig { target_tokens, @@ -1842,7 +1898,7 @@ impl ConversationOrchestrator { let response: Arc>>> = Arc::new(SyncMutex::new(None)); let llm = self.llm.clone(); let session_id = self.session_id.clone(); - let api_key = self.config.api_key.clone(); + let api_key = self.llm_key(); let text_owned = text.to_string(); let response_store = response.clone(); let task_token = token.clone(); @@ -2085,6 +2141,20 @@ impl ConversationOrchestrator { StageErrorClass::Recoverable => { warn!(session = %self.session_id, error = %e, "conversation turn failed (recoverable; call continues)"); } + StageErrorClass::Fatal + if self.config.credential.is_some() && is_auth_failure(&e.to_string()) => + { + // FRD-023 RT6: the session's credential is refreshable (an `auth` message + // replaces it), so a refused credential — typically an expired Keycloak token — + // is not the end of the call. Tell the client, keep listening; the next turn + // sends whatever credential is current. + warn!(session = %self.session_id, error = %e, + "LLM leg refused the session credential; waiting for an auth refresh"); + let handler = self.fatal_handler.lock().clone(); + if let Some(handler) = handler { + handler(format!("auth_expired: {e}")); + } + } StageErrorClass::Fatal => { tracing::error!(session = %self.session_id, error = %e, "FATAL turn error (auth/config) — stopping the session"); @@ -2252,6 +2322,17 @@ fn eager_transcript_matches(speculation: &str, final_transcript: &str) -> bool { norm(speculation) == norm(final_transcript) } +/// Whether an LLM error is the provider refusing the credential (401/403 or an auth message). +fn is_auth_failure(message: &str) -> bool { + let e = message.to_ascii_lowercase(); + e.contains("http 401") + || e.contains("http 403") + || e.contains("invalid api key") + || e.contains("invalid_api_key") + || e.contains("unauthorized") + || e.contains("authentication") +} + /// Validate a client-supplied LLM base URL for SSRF. /// /// Thin wrapper over the canonical [`crate::core::net::validate_url_for_ssrf`] diff --git a/gateway/src/core/voice_manager/manager.rs b/gateway/src/core/voice_manager/manager.rs index 6749c3dc..57fb261a 100644 --- a/gateway/src/core/voice_manager/manager.rs +++ b/gateway/src/core/voice_manager/manager.rs @@ -62,6 +62,12 @@ fn uninterruptible_playback_from_env() -> VoiceManagerResult { } } +/// See [`VoiceManager::set_speak_observer`]. +pub type SpeakObserver = Arc; + +/// See [`VoiceManager::set_stt_final_observer`]. +pub type SttFinalObserver = Arc; + /// VoiceManager provides a unified interface for managing STT and TTS providers /// Optimized for extreme low-latency with lock-free atomics and pre-allocated buffers pub struct VoiceManager { @@ -118,6 +124,13 @@ pub struct VoiceManager { // None ⇒ a single relaxed read per call site, zero work. Set once at // session setup via `set_observers`; read-mostly thereafter. observers: Arc>>>, + /// FRD-023 RT6: told the text of every synthesis request, so a Bud deployment's TTS leg is + /// metered per `speak` — the client's and the voice agent's alike. + speak_observer: Arc>>, + /// FRD-023 RT6: told every FINAL transcript the vendor returns, so a Bud deployment's STT leg + /// is metered per utterance. Separate from the result callback, which the voice agent + /// replaces; read per result, so it holds across every `on_stt_result` registration. + stt_final_observer: Arc>>, /// A-G6: when uninterruptible playback is enabled, TTS chunks are metered to /// the transport through this pump's queue so a barge-in can selectively @@ -244,6 +257,8 @@ impl VoiceManager { clear_notify: Arc::new(Notify::new()), clear_epoch: Arc::new(AtomicUsize::new(0)), observers: Arc::new(SyncRwLock::new(None)), + speak_observer: Arc::new(SyncRwLock::new(None)), + stt_final_observer: Arc::new(SyncRwLock::new(None)), playback_pump: Arc::new(SyncRwLock::new(None)), uninterruptible_playback: AtomicBool::new(uninterruptible_playback), }) @@ -705,16 +720,34 @@ impl VoiceManager { // D-G9 (review wf_d43814c3): count chars on THIS path too — the // orchestrator speaks via speak_if_epoch, not speak(), so the // counter was never incremented on the production conversation path. - crate::core::metrics::bridge::count_tts_chars( - &self.config.tts_config.provider, - text.chars().count(), - ); + self.note_speak(text); tts.speak(text, flush) .await .map_err(VoiceManagerError::TTSError)?; Ok(true) } + /// Every synthesis request passes here: the D-G9 cost proxy and the RT6 speak observer. + fn note_speak(&self, text: &str) { + crate::core::metrics::bridge::count_tts_chars( + &self.config.tts_config.provider, + text.chars().count(), + ); + if let Some(observer) = self.speak_observer.read().clone() { + observer(text); + } + } + + /// Observe the text of every synthesis request (FRD-023 RT6 TTS-leg metering). + pub fn set_speak_observer(&self, observer: SpeakObserver) { + *self.speak_observer.write() = Some(observer); + } + + /// Observe every final transcript the STT vendor returns (FRD-023 RT6 STT-leg metering). + pub fn set_stt_final_observer(&self, observer: SttFinalObserver) { + *self.stt_final_observer.write() = Some(observer); + } + pub async fn speak(&self, text: &str, flush: bool) -> VoiceManagerResult<()> { // Per-turn profiling anchor: TTS synthesis requested (first-of-turn // dedup happens inside the profiler). @@ -722,10 +755,7 @@ impl VoiceManager { obs.notify_tts_request(crate::core::observability::now_monotonic_ns()); } // D-G9: synthesis cost proxy. - crate::core::metrics::bridge::count_tts_chars( - &self.config.tts_config.provider, - text.chars().count(), - ); + self.note_speak(text); // Send text to TTS provider { let mut tts = self.tts.write().await; @@ -786,10 +816,7 @@ impl VoiceManager { // D-G9: synthesis cost proxy (covers the non-interruptible / // speak_if_epoch-delegated path). - crate::core::metrics::bridge::count_tts_chars( - &self.config.tts_config.provider, - text.chars().count(), - ); + self.note_speak(text); // Send text to TTS provider { let mut tts = self.tts.write().await; @@ -986,6 +1013,7 @@ impl VoiceManager { let interruption_state_clone = self.interruption_state.clone(); let turn_detector_clone = self.turn_detector.clone(); let observers_clone = self.observers.clone(); + let stt_final_observer_clone = self.stt_final_observer.clone(); // Create STT processor with configured timeouts from VoiceManagerConfig, // plus the provider's measured TTFS p99 (A-G2/D-G8: a slow provider's @@ -1010,6 +1038,13 @@ impl VoiceManager { let turn_detector = turn_detector_clone.clone(); let stt_processor = stt_processor.clone(); let observers = observers_clone.read().clone(); + // Before any suppression: the vendor transcribed (and billed) this audio whether or + // not the session acts on the result. + if result.is_final + && let Some(observer) = stt_final_observer_clone.read().clone() + { + observer(&result.transcript); + } Box::pin(async move { // Fast synchronous check for interruption - execute before any async ops diff --git a/gateway/src/dag/compiler.rs b/gateway/src/dag/compiler.rs index 365990a7..de8f6b50 100644 --- a/gateway/src/dag/compiler.rs +++ b/gateway/src/dag/compiler.rs @@ -269,6 +269,9 @@ impl DAGCompiler { if let Some(m) = model { node = node.with_model(m); } + if let Some(bud) = &def.bud { + node = node.with_bud(bud.clone()); + } Arc::new(node) } NodeType::RealtimeProvider { provider, model } => { @@ -381,8 +384,7 @@ impl DAGCompiler { llm_config.tools = Some(parsed_tools); } - // Use try_new() for SSRF protection on the client-supplied base_url. - Arc::new(LlmEndpointNode::try_new(&def.id, llm_config)?) + bud_or_checked_llm_node(def, llm_config)? } NodeType::Translate { target_language, @@ -433,7 +435,7 @@ impl DAGCompiler { headers: headers.clone(), ..Default::default() }; - Arc::new(LlmEndpointNode::try_new(&def.id, llm_config)?) + bud_or_checked_llm_node(def, llm_config)? } NodeType::WebhookOutput { url, headers } => { // Use try_new() for SSRF protection (S6): webhook URLs are client-supplied. @@ -547,6 +549,21 @@ impl Default for DAGCompiler { /// excludes reconvergence/join nodes, which have at least one predecessor from a /// different branch and therefore execute once in the main sweep rather than per /// branch. This is the data needed for single-execution split handling. +/// An LLM node: SSRF-checked when its `base_url` came with the template, or pointed at the Bud +/// gateway with the caller's credential when the server bound it (FRD-023 WP-RT6.3). +fn bud_or_checked_llm_node( + def: &NodeDefinition, + llm_config: LlmEndpointConfig, +) -> DAGResult> { + match def.bud.as_ref().and_then(|b| b.session_credential.clone()) { + Some(credential) => Ok(Arc::new(LlmEndpointNode::for_bud_gateway( + &def.id, llm_config, credential, + ))), + // Use try_new() for SSRF protection on the client-supplied base_url. + None => Ok(Arc::new(LlmEndpointNode::try_new(&def.id, llm_config)?)), + } +} + fn compute_split_plans( graph: &DiGraph, topo_order: &[NodeIndex], diff --git a/gateway/src/dag/definition.rs b/gateway/src/dag/definition.rs index 1190815f..8a73dc93 100644 --- a/gateway/src/dag/definition.rs +++ b/gateway/src/dag/definition.rs @@ -460,6 +460,12 @@ pub struct NodeDefinition { /// Maximum retry attempts #[serde(default = "default_max_retries")] pub max_retries: u32, + + /// FRD-023 RT6 (WP-RT6.3): the Bud deployment this node was bound to by the SERVER — its + /// vendor credential, address and metering, or the caller's credential for an LLM node. + /// Never deserialized, so neither a template nor a client can supply one. + #[serde(skip)] + pub bud: Option>, } fn default_max_retries() -> u32 { @@ -476,6 +482,7 @@ impl NodeDefinition { timeout_ms: None, retry_on_failure: false, max_retries: default_max_retries(), + bud: None, } } diff --git a/gateway/src/dag/nodes/llm.rs b/gateway/src/dag/nodes/llm.rs index 8c3d7f0b..d77a9379 100644 --- a/gateway/src/dag/nodes/llm.rs +++ b/gateway/src/dag/nodes/llm.rs @@ -238,6 +238,9 @@ pub struct LlmEndpointNode { id: String, streaming: bool, client: LlmClient, + /// FRD-023 WP-RT6.3: a node bound to a Bud chat deployment sends the CALLER's credential, + /// read per call, instead of `ctx.api_key`. + session_credential: Option, } impl std::fmt::Debug for LlmEndpointNode { @@ -261,6 +264,28 @@ impl LlmEndpointNode { id: id.into(), streaming, client, + session_credential: None, + } + } + + /// A node the server pointed at the Bud gateway (FRD-023 WP-RT6.3): its `base_url` is the + /// operator's in-cluster address (so not SSRF-checked, as a client's would be) and every call + /// carries the session's credential. + pub fn for_bud_gateway( + id: impl Into, + config: LlmEndpointConfig, + credential: crate::auth::SessionCredential, + ) -> Self { + let mut node = Self::new(id, config); + node.session_credential = Some(credential); + node + } + + /// The key for this call. + fn call_key(&self, ctx: &DAGContext) -> Option { + match &self.session_credential { + Some(credential) => Some(credential.current()), + None => ctx.api_key.clone(), } } @@ -362,12 +387,13 @@ impl DAGNode for LlmEndpointNode { // Per-connection API key from the DAG context takes priority (matches the // old node behavior); otherwise the client falls back to config/env. + let key = self.call_key(ctx); let response = self .client .complete( &ctx.stream_id, &input_text, - ctx.api_key.as_deref(), + key.as_deref(), &ctx.cancel_token, None, ) @@ -438,10 +464,11 @@ impl DAGNode for LlmEndpointNode { // Drop reasoning chain-of-thought before it streams downstream. let mut think = crate::core::text::ThinkStripper::default(); let mut emitted_any = false; + let key = self.call_key(ctx); let completion = self.client.complete( &ctx.stream_id, &input_text, - ctx.api_key.as_deref(), + key.as_deref(), &ctx.cancel_token, Some(on_token), ); diff --git a/gateway/src/dag/nodes/mod.rs b/gateway/src/dag/nodes/mod.rs index 5e3f3e48..6a1c55df 100644 --- a/gateway/src/dag/nodes/mod.rs +++ b/gateway/src/dag/nodes/mod.rs @@ -21,8 +21,8 @@ pub use llm::{ChatMessage, LlmEndpointConfig, LlmEndpointNode, ResponseFormat, T pub use output::{AudioOutputNode, TextOutputNode, WebhookOutputNode}; pub use processor::ProcessorNode; pub use provider::{ - RealtimeProviderNode, RealtimeSessionMap, STTProviderNode, SessionRealtime, TTSProviderNode, - disconnect_realtime_sessions, realtime_resilience_key, realtime_sessions_key, + BudNodeBinding, RealtimeProviderNode, RealtimeSessionMap, STTProviderNode, SessionRealtime, + TTSProviderNode, disconnect_realtime_sessions, realtime_resilience_key, realtime_sessions_key, }; pub use router::{JoinNode, RouterNode, SplitNode}; pub use transform::{PassthroughNode, TransformNode}; diff --git a/gateway/src/dag/nodes/provider.rs b/gateway/src/dag/nodes/provider.rs index 69cfc1fd..c631636a 100644 --- a/gateway/src/dag/nodes/provider.rs +++ b/gateway/src/dag/nodes/provider.rs @@ -121,6 +121,32 @@ pub(crate) fn bud_mode_node_credential_error( )) } +/// A provider or LLM node bound by the server to a Bud deployment (FRD-023 WP-RT6.3). +/// +/// Built only by the `/ws` session from the caller's own allowlist; never deserialized. +pub struct BudNodeBinding { + /// The deployment, for logs. + pub endpoint_id: String, + /// A provider node's vendor credential, from the deployment's `voice_table` entry. + pub vendor_credential: Option, + /// The deployment's own address (self-hosted, Azure OpenAI). + pub api_base: Option, + /// The deployment's own vendor parameters (AWS region and keys, Google project, …). + pub extras: serde_json::Map, + /// An LLM node's credential: the CALLER's, read per call so a refresh reaches it. + pub session_credential: Option, + /// Told the text of every synthesis, to meter it against the deployment. + pub on_synthesis: Option>, +} + +impl std::fmt::Debug for BudNodeBinding { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BudNodeBinding") + .field("endpoint_id", &self.endpoint_id) + .finish_non_exhaustive() + } +} + /// Callback bridge for TTS provider to DAG node /// /// This struct implements the `AudioCallback` trait and bridges @@ -553,6 +579,8 @@ pub struct TTSProviderNode { config: serde_json::Value, /// Maximum total audio bytes to collect (prevents memory exhaustion) max_audio_bytes: usize, + /// FRD-023 WP-RT6.3: bound to a Bud deployment by the server. + bud: Option>, } impl TTSProviderNode { @@ -565,6 +593,7 @@ impl TTSProviderNode { model: None, config: serde_json::Value::Null, max_audio_bytes: DEFAULT_MAX_TTS_AUDIO_BYTES, + bud: None, } } @@ -589,6 +618,12 @@ impl TTSProviderNode { /// Set maximum audio bytes limit (default: 100MB) /// /// This prevents memory exhaustion from abnormally long TTS audio. + /// Bind the node to a Bud deployment (FRD-023 WP-RT6.3). + pub fn with_bud(mut self, binding: Arc) -> Self { + self.bud = Some(binding); + self + } + pub fn with_max_audio_bytes(mut self, max_bytes: usize) -> Self { self.max_audio_bytes = max_bytes; self @@ -673,26 +708,44 @@ impl DAGNode for TTSProviderNode { // Get TTS provider from registry let registry = crate::plugin::global_registry(); - // Build TTS configuration. A configured credential must resolve; when no - // DAG credential is supplied, provider-specific fallback may still apply. - let api_key = resolve_configured_node_credential( - &self.config, - "api_key", - &self.id, - &self.provider, - "TTS", - )? - .unwrap_or_default(); + // Build TTS configuration. A node bound to a Bud deployment uses the deployment's + // credential, address and parameters (FRD-023 WP-RT6.3). Otherwise a configured + // credential must resolve; when no DAG credential is supplied, provider-specific fallback + // may still apply. + let api_key = match &self.bud { + Some(bud) => bud.vendor_credential.clone().unwrap_or_default(), + None => resolve_configured_node_credential( + &self.config, + "api_key", + &self.id, + &self.provider, + "TTS", + )? + .unwrap_or_default(), + }; let tts_config = crate::core::tts::TTSConfig { provider: self.provider.clone(), voice_id: self.voice_id.clone(), model: self.model.clone().unwrap_or_default(), api_key, + api_base: self.bud.as_ref().and_then(|b| b.api_base.clone()), ..Default::default() }; // Create TTS provider - let mut tts = match registry.create_tts(&self.provider, tts_config) { + let created = match &self.bud { + Some(bud) => { + if let Some(meter) = &bud.on_synthesis { + meter(&text); + } + let mut standard = + crate::core::tts::standard::StandardTTSConfig::from_base(tts_config); + standard.extras.0 = bud.extras.clone(); + crate::core::tts::standard::create_tts_standard(&self.provider, standard) + } + None => registry.create_tts(&self.provider, tts_config), + }; + let mut tts = match created { Ok(tts) => tts, Err(e) => { return Err(DAGError::TTSProviderError { diff --git a/gateway/src/handlers/openai_audio.rs b/gateway/src/handlers/openai_audio.rs index 22075a16..896dbbcb 100644 --- a/gateway/src/handlers/openai_audio.rs +++ b/gateway/src/handlers/openai_audio.rs @@ -1462,7 +1462,7 @@ fn is_aws_vendor(vendor: &str) -> bool { /// GATEWAY's identity in us-east-1), a Google project and location, the Azure Speech host, the /// Azure OpenAI api-version. Only the keys named here are copied — never `endpoint_override`, /// which is a destination for the vendor's credential. -fn deployment_extras( +pub(crate) fn deployment_extras( endpoint: &bud_auth::credentials::VoiceEndpoint, ) -> serde_json::Map { let mut extras = serde_json::Map::new(); @@ -1524,17 +1524,25 @@ fn endpoint_misconfiguration( name: &str, advisories: &mut Advisories, ) -> Option { - let refuse = |why: String| { - Some(openai_error( - StatusCode::INTERNAL_SERVER_ERROR, - "api_error", - format!( - "Endpoint '{name}' is misconfigured for vendor '{}': {why}", - endpoint.vendor - ), - None, - )) - }; + let why = endpoint_misconfiguration_reason(endpoint, advisories)?; + Some(openai_error( + StatusCode::INTERNAL_SERVER_ERROR, + "api_error", + format!( + "Endpoint '{name}' is misconfigured for vendor '{}': {why}", + endpoint.vendor + ), + None, + )) +} + +/// Why [`endpoint_misconfiguration`] refuses a deployment, shared with the `/ws` legs (FRD-023 +/// RT6), which reach the same vendors with the same credentials. +pub(crate) fn endpoint_misconfiguration_reason( + endpoint: &bud_auth::credentials::VoiceEndpoint, + advisories: &mut Advisories, +) -> Option { + let refuse = |why: String| Some(why); if is_azure_speech(&endpoint.vendor) && let Some(api_base) = endpoint .api_base @@ -1982,7 +1990,7 @@ fn describing_voice( /// the resolver's `style`, because that is what the control above it promises — "words matched /// against the vendor's own voice metadata: warm, gravelly, bright". The blob key keeps its /// original spelling so an entry written before this still parses. -async fn resolve_described_voice( +pub(crate) async fn resolve_described_voice( state: &Arc, endpoint: &bud_auth::credentials::VoiceEndpoint, advisories: &mut Advisories, diff --git a/gateway/src/handlers/ws/audio_handler.rs b/gateway/src/handlers/ws/audio_handler.rs index cc150173..f025e17b 100644 --- a/gateway/src/handlers/ws/audio_handler.rs +++ b/gateway/src/handlers/ws/audio_handler.rs @@ -80,7 +80,7 @@ pub async fn handle_audio_message( } // Fast path: read lock to check state, get the voice manager, and (D8) opus-decode the frame. - let (voice_manager, audio_data) = { + let (voice_manager, audio_data, leg_meter) = { let state_guard = state.read().await; // Check if audio processing is enabled (atomic read, no lock overhead) @@ -120,10 +120,13 @@ pub async fn handle_audio_message( None => audio_data, }; - (voice_manager, audio_data) + (voice_manager, audio_data, state_guard.leg_meter.clone()) }; // Send the (decoded) PCM audio to the STT provider. Bytes gives O(1) clones. + if let Some(meter) = &leg_meter { + meter.add_stt_audio(audio_data.len()); + } if let Err(e) = voice_manager.receive_audio(audio_data).await { error!("Failed to process audio: {}", e); send_error(message_tx, format!("Failed to process audio: {e}")).await; diff --git a/gateway/src/handlers/ws/bud_legs.rs b/gateway/src/handlers/ws/bud_legs.rs new file mode 100644 index 00000000..d23cb74d --- /dev/null +++ b/gateway/src/handlers/ws/bud_legs.rs @@ -0,0 +1,1694 @@ +//! `/ws` on Bud deployments (FRD-023 RT6, §5.9, FR-WS-1…5). +//! +//! Under the Bud control plane a `/ws` session's legs address DEPLOYMENTS, exactly as the REST +//! routes do: `stt_config.model` names a transcription deployment and `tts_config.model` a +//! text-to-speech deployment, each resolved through the caller's own allowlist. The vendor, its +//! model, its credential, its `api_base` and its non-secret parameters come from the deployment's +//! `voice_table` entry — never from the client and never from WaaV's environment (D-5, X-4). +//! +//! Each leg takes the deployment's admission once, at session start, and holds it for the session +//! (FRD-022 §6.2). Each leg is metered as its own `voice.turn` records: a final transcription +//! bills the audio seconds streamed since the previous one; a `speak` bills its characters. And the +//! session is revalidated every 30 s, so a revoked key, a user removed from the project, or an +//! unpublished deployment ends it (D-17). + +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; + +use bud_auth::{AliasMetadata, VoiceEndpoint, VoicePricing}; +use tracing::Span; + +use crate::core::deployment_policy::{Admission, Rejection}; +use crate::core::voice_cost::voice_cost; +use crate::handlers::advisories::Advisories; +use crate::handlers::endpoint_settings; +use crate::handlers::openai_realtime::session::{Caller, CallerCheck}; +use crate::observability::voice_attrs::{leg as leg_attr, turn}; +use crate::state::AppState; + +use super::config::{STTWebSocketConfig, TTSWebSocketConfig}; + +pub const STT_CAPABILITY: &str = "audio_transcription"; +pub const TTS_CAPABILITY: &str = "text_to_speech"; + +/// A leg resolved to a Bud deployment, with its admission held. +pub struct BudLeg { + pub endpoint_id: String, + /// The name the client used. + pub endpoint_name: String, + pub endpoint: VoiceEndpoint, + pub alias: Option, + pub admission: Admission, +} + +impl std::fmt::Debug for BudLeg { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BudLeg") + .field("endpoint_id", &self.endpoint_id) + .field("endpoint_name", &self.endpoint_name) + .field("vendor", &self.endpoint.vendor) + .finish() + } +} + +/// Why a leg could not be served. `message` goes to the client. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct LegRefusal { + pub code: &'static str, + pub message: String, + /// Close the socket with this code. `None` leaves it open for a corrected `config`. + pub close: Option, +} + +impl LegRefusal { + fn new(code: &'static str, message: impl Into) -> Self { + Self { + code, + message: message.into(), + close: None, + } + } + + fn closing(mut self, code: u16) -> Self { + self.close = Some(code); + self + } +} + +/// The close code for a session its caller may no longer hold (revoked, removed, unpublished). +pub const CLOSE_REVOKED: u16 = 1008; +/// The close code for a deployment at its rate or concurrency limit: retry later. +pub const CLOSE_TRY_LATER: u16 = 1013; + +fn kind(capability: &str) -> &'static str { + if capability == STT_CAPABILITY { + "transcription" + } else { + "text-to-speech" + } +} + +/// Resolve `model` (named at `field`, e.g. `stt_config.model`) as a deployment serving `capability` +/// and admit the session to it. +pub async fn resolve_leg( + state: &AppState, + credential: &str, + field: &str, + model: &str, + capability: &'static str, +) -> Result { + let model = model.trim(); + if model.is_empty() { + return Err(LegRefusal::new( + "deployment_required", + format!( + "{field} must name your Bud {} deployment: this gateway serves Bud \ + deployments only, and uses the deployment's own vendor credential; `provider` \ + is taken from the deployment (FRD-023 RT6).", + kind(capability) + ), + )); + } + let Some(resolved) = state.resolve_voice_endpoint(model, capability, Some(credential)) else { + return Err(LegRefusal::new( + "model_not_found", + format!( + "{field} '{model}' is not a {} deployment this credential can reach.", + kind(capability) + ), + )); + }; + let vendor = resolved.endpoint.vendor.as_str(); + if capability == STT_CAPABILITY + && (crate::core::tts::self_hosted::is_self_hosted(vendor) + || crate::core::tts::self_hosted::is_azure_openai(vendor)) + { + // These answer one uploaded file over HTTP; there is no stream to hold open. + return Err(LegRefusal::new( + "unsupported_deployment", + format!( + "Deployment '{model}' ({vendor}) transcribes uploaded files and cannot stream; use \ + it through /v1/audio/transcriptions, or name a streaming transcription deployment." + ), + )); + } + let mut ignored = Advisories::new(); + if let Some(why) = crate::handlers::openai_audio::endpoint_misconfiguration_reason( + &resolved.endpoint, + &mut ignored, + ) { + return Err(LegRefusal::new( + "deployment_misconfigured", + format!("Deployment '{model}' is misconfigured for vendor '{vendor}': {why}"), + )); + } + let admission = state + .admit_deployment(&resolved.endpoint_id) + .await + .map_err(|rej| { + let (code, what) = match rej { + Rejection::Rate(_) => ("rate_limit_exceeded", "rate limit"), + Rejection::Concurrency(_) => { + ("concurrency_limit_exceeded", "concurrent-session limit") + } + }; + LegRefusal::new( + code, + format!( + "Deployment '{model}' reached its {what}; retry in {} s.", + rej.retry_after().as_secs().max(1) + ), + ) + .closing(CLOSE_TRY_LATER) + })?; + Ok(BudLeg { + endpoint_id: resolved.endpoint_id, + endpoint_name: model.to_string(), + endpoint: resolved.endpoint, + alias: resolved.alias, + admission, + }) +} + +/// The deployment's own vendor parameters replace whatever the client sent: some providers read a +/// DESTINATION from `extras` (Groq's `url`, Azure Speech's host), and on a Bud leg the credential +/// that would travel there is the deployment's. +fn replace_extras( + target: &mut serde_json::Map, + endpoint: &VoiceEndpoint, +) { + target.clear(); + target.extend(crate::handlers::openai_audio::deployment_extras(endpoint)); +} + +/// `base`, with every field the client set explicitly laid over it: request > deployment. +fn overlay(base: T, client: &T) -> T +where + T: serde::Serialize + serde::de::DeserializeOwned, +{ + let (Ok(serde_json::Value::Object(mut merged)), Ok(serde_json::Value::Object(ours))) = + (serde_json::to_value(&base), serde_json::to_value(client)) + else { + return base; + }; + for (k, v) in ours { + if !v.is_null() { + merged.insert(k, v); + } + } + match serde_json::from_value(serde_json::Value::Object(merged)) { + Ok(v) => v, + Err(_) => base, + } +} + +fn non_empty(v: &Option) -> Option { + v.as_deref() + .map(str::trim) + .filter(|v| !v.is_empty()) + .map(str::to_string) +} + +/// Point the STT leg at its deployment (FR-WS-1, FR-WS-4). Returns the vendor credential. +/// +/// The deployment's published transcription settings are the defaults and the client's explicit +/// choices win, as on REST; `stt.streaming` applies here and only here. +pub fn apply_stt( + cfg: &mut STTWebSocketConfig, + leg: &BudLeg, + advisories: &mut Advisories, +) -> String { + let ep = &leg.endpoint; + let settings = ep.config.stt(); + cfg.provider = ep.vendor.clone(); + cfg.model = non_empty(&settings.model) + .or_else(|| ep.model.clone()) + .unwrap_or_default(); + if cfg.language.trim().is_empty() + && let Some(lang) = non_empty(&settings.language).or_else(|| ep.language.clone()) + { + cfg.language = lang; + } + + let mut base = endpoint_settings::stt_features_for(&settings, &ep.vendor, advisories); + // `stt_features_for` pins Deepgram's `vad_events` to the REST plane's historical default; + // on this plane unset has always meant unset. + base.vad_events = None; + if let Some(streaming) = &settings.streaming { + base.interim_results = streaming.interim_results; + base.vad_events = streaming.vad_events; + base.endpointing_ms = streaming.endpointing_ms; + base.utterance_end_ms = streaming.utterance_end_ms; + base.speech_begin_event = streaming.speech_begin_event; + } + cfg.features = overlay(base, &cfg.features); + if cfg.translation.is_none() { + cfg.translation = + endpoint_settings::translation_for(ep.config.translation.as_ref(), false, advisories); + } + replace_extras(&mut cfg.extras.0, ep); + cfg.api_key = None; + ep.credential.clone().unwrap_or_default() +} + +/// Point the TTS leg at its deployment (FR-WS-1). Returns the vendor credential; the deployment's +/// `api_base` is the leg's [`BudLeg::api_base`]. +pub fn apply_tts( + cfg: &mut TTSWebSocketConfig, + leg: &BudLeg, + advisories: &mut Advisories, +) -> String { + let ep = &leg.endpoint; + let settings = ep.config.tts(); + cfg.provider = ep.vendor.clone(); + cfg.model = ep.model.clone().unwrap_or_default(); + if cfg.sample_rate.is_none() { + cfg.sample_rate = settings.sample_rate; + } + if cfg.connection_timeout.is_none() { + cfg.connection_timeout = settings.connection_timeout; + } + if cfg.request_timeout.is_none() { + cfg.request_timeout = settings.request_timeout; + } + if cfg.pronunciations.is_empty() + && let Some(list) = &settings.pronunciations + { + cfg.pronunciations = list + .iter() + .map(|p| crate::core::tts::Pronunciation { + word: p.word.clone(), + pronunciation: p.pronunciation.clone(), + }) + .collect(); + } + let language = ep.language.clone().or_else(|| settings.language.clone()); + let base = + endpoint_settings::tts_features_for(&settings, &ep.vendor, language.as_deref(), advisories); + cfg.features = overlay(base, &cfg.features); + replace_extras(&mut cfg.extras.0, ep); + cfg.api_key = None; + ep.credential.clone().unwrap_or_default() +} + +/// Choose the TTS leg's voice as REST does (request > deployment), with every catalogue read made +/// with the DEPLOYMENT's credential: the client's voice id; the client's descriptor matched in the +/// deployment vendor's catalogue; the deployment's voice; the deployment's descriptor. +pub async fn resolve_tts_voice( + state: &Arc, + cfg: &mut TTSWebSocketConfig, + leg: &BudLeg, + advisories: &mut Advisories, +) { + if cfg + .voice_id + .as_deref() + .is_some_and(|v| !v.trim().is_empty()) + { + return; + } + let ep = &leg.endpoint; + if let Some(described) = cfg.voice_descriptor.clone().filter(|d| d.is_set()) { + let catalog = crate::handlers::voices::fetch_provider_catalog_with_key( + state, + &ep.vendor, + ep.credential.as_deref(), + ) + .await; + let resolved = crate::core::voice::resolve_voice( + &described, + &catalog, + crate::handlers::voices::provider_default_voice(&ep.vendor), + ); + if let Some(warning) = resolved.warning { + advisories.warn(warning); + } + if !resolved.voice_id.trim().is_empty() { + cfg.voice_id = Some(resolved.voice_id); + } + return; + } + if let Some(voice) = non_empty(&ep.voice) { + cfg.voice_id = Some(voice); + return; + } + cfg.voice_id = + crate::handlers::openai_audio::resolve_described_voice(state, ep, advisories).await; + // Nothing named or described a voice: where the vendor REQUIRES one, its default (as REST + // answers, rather than a session that cannot speak); elsewhere the vendor picks. + if cfg.voice_id.is_none() && crate::handlers::voices::voice_required(&ep.vendor) { + let default = crate::handlers::voices::provider_default_voice(&ep.vendor); + if !default.is_empty() { + advisories.warn(format!( + "deployment '{}' has no voice configured, so {}'s default voice was used; set one \ + on the deployment or send tts_config.voice_id to choose it", + leg.endpoint_name, ep.vendor + )); + cfg.voice_id = Some(default.to_string()); + } + } +} + +/// Who a leg's records are attributed to and what they cost. +#[derive(Debug, Clone)] +struct LegBilling { + capability: &'static str, + endpoint_id: String, + endpoint_name: String, + model_id: Option, + project_id: Option, + vendor: String, + vendor_model: Option, + pricing: Option, +} + +impl LegBilling { + fn from_leg(capability: &'static str, leg: &BudLeg, caller: &Caller) -> Self { + Self { + capability, + endpoint_id: leg.endpoint_id.clone(), + endpoint_name: leg.endpoint_name.clone(), + model_id: leg.alias.as_ref().and_then(|a| a.model_id.clone()), + project_id: leg + .alias + .as_ref() + .and_then(|a| a.project_id.clone()) + .or_else(|| caller.principal.project_id.clone()), + vendor: leg.endpoint.vendor.clone(), + vendor_model: leg.endpoint.model.clone(), + pricing: leg.endpoint.pricing.clone(), + } + } +} + +fn record_text(span: &Span, key: &'static str, value: Option<&str>) { + if let Some(v) = value.filter(|v| !v.is_empty()) { + span.record(key, v); + } +} + +/// Meters a Bud-mode `/ws` session's STT and TTS legs (FR-WS-5, TC-WS-12). +pub struct LegMeter { + session_id: String, + caller: Caller, + stt: Option, + tts: Option, + /// PCM16 bytes streamed to the STT leg since the last final transcript. + stt_bytes: AtomicU64, + /// Settable: the codec negotiation may move the rate after the legs are resolved. + stt_bytes_per_second: AtomicU64, + turn_index: AtomicU64, + /// Further deployments the session holds (a DAG template's bound nodes), revalidated with it. + held: parking_lot::Mutex>, +} + +impl std::fmt::Debug for LegMeter { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LegMeter") + .field("session_id", &self.session_id) + .finish() + } +} + +impl LegMeter { + pub fn new( + session_id: String, + caller: Caller, + stt: Option<(&BudLeg, u32, u16)>, + tts: Option<&BudLeg>, + ) -> Self { + let meter = Self { + session_id, + stt: stt.map(|(leg, _, _)| LegBilling::from_leg(STT_CAPABILITY, leg, &caller)), + tts: tts.map(|leg| LegBilling::from_leg(TTS_CAPABILITY, leg, &caller)), + caller, + stt_bytes: AtomicU64::new(0), + stt_bytes_per_second: AtomicU64::new(0), + turn_index: AtomicU64::new(0), + held: parking_lot::Mutex::new(Vec::new()), + }; + if let Some((_, rate, channels)) = stt { + meter.set_stt_format(rate, channels); + } + meter + } + + /// The PCM16 format the STT leg receives, as finally negotiated. + pub fn set_stt_format(&self, sample_rate: u32, channels: u16) { + let bps = u64::from(sample_rate) * 2 * u64::from(channels.max(1)); + self.stt_bytes_per_second.store(bps, Ordering::Relaxed); + } + + /// Count audio streamed to the STT leg (after any codec decode). One atomic add: hot path. + pub fn add_stt_audio(&self, bytes: usize) { + self.stt_bytes.fetch_add(bytes as u64, Ordering::Relaxed); + } + + fn open(&self, billing: &LegBilling) -> Span { + let span = crate::voice_turn_span!( + parent: None, + capability = billing.capability, + transport = "websocket" + ); + let p = &self.caller.principal; + record_text(&span, turn::PROJECT_ID, billing.project_id.as_deref()); + record_text(&span, turn::ENDPOINT_ID, Some(&billing.endpoint_id)); + record_text(&span, turn::ENDPOINT_NAME, Some(&billing.endpoint_name)); + record_text(&span, turn::MODEL_ID, billing.model_id.as_deref()); + record_text(&span, turn::API_KEY_ID, p.api_key_id.as_deref()); + record_text(&span, turn::USER_ID, p.user_id.as_deref()); + record_text(&span, turn::API_KEY_PROJECT_ID, p.project_id.as_deref()); + record_text(&span, turn::SESSION_ID, Some(&self.session_id)); + span.record( + turn::TURN_INDEX, + self.turn_index.fetch_add(1, Ordering::Relaxed), + ); + span + } + + /// A final transcript: bill the audio streamed since the previous one. + pub fn stt_final(&self, transcript: &str) { + let Some(billing) = &self.stt else { return }; + let bps = self.stt_bytes_per_second.load(Ordering::Relaxed); + if bps == 0 { + return; + } + let bytes = self.stt_bytes.swap(0, Ordering::Relaxed); + if bytes == 0 { + return; + } + let seconds = bytes as f64 / bps as f64; + let span = self.open(billing); + record_text(&span, leg_attr::STT_VENDOR, Some(&billing.vendor)); + record_text(&span, leg_attr::STT_MODEL, billing.vendor_model.as_deref()); + span.record(turn::AUDIO_SECONDS, seconds); + if let Some((cost, unit)) = voice_cost( + billing.pricing.as_ref(), + STT_CAPABILITY, + None, + Some(seconds), + None, + ) { + span.record(turn::COST, cost); + span.record(turn::PRICING_UNIT, unit); + } + if crate::observability::trace_redact::capture_content() { + span.record( + turn::TRANSCRIPT, + crate::observability::trace_redact::sanitize_body(transcript).as_str(), + ); + } + } + + /// A `speak`: bill its characters. + pub fn tts_spoken(&self, text: &str) { + let Some(billing) = &self.tts else { return }; + let characters = text.chars().count() as u64; + if characters == 0 { + return; + } + let span = self.open(billing); + record_text(&span, leg_attr::TTS_VENDOR, Some(&billing.vendor)); + record_text(&span, leg_attr::TTS_MODEL, billing.vendor_model.as_deref()); + span.record(turn::CHARACTERS, characters); + if let Some((cost, unit)) = voice_cost( + billing.pricing.as_ref(), + TTS_CAPABILITY, + Some(characters), + None, + None, + ) { + span.record(turn::COST, cost); + span.record(turn::PRICING_UNIT, unit); + } + } + + /// At close: bill audio streamed after the last final transcript, so nothing streamed goes + /// unbilled. + pub fn finish(&self) { + self.stt_final(""); + } + + /// How to re-check this session's caller. + pub fn check(&self) -> &CallerCheck { + &self.caller.check + } + + /// The deployments this session holds and what each serves it, for revalidation. + pub fn endpoints(&self) -> Vec<(String, &'static str)> { + let mut all: Vec<(String, &'static str)> = self + .stt + .iter() + .chain(self.tts.iter()) + .map(|b| (b.endpoint_id.clone(), b.capability)) + .collect(); + all.extend(self.held.lock().iter().cloned()); + all + } + + /// Revalidate `endpoint_id` with the session from now on. + pub fn hold(&self, endpoint_id: String, capability: &'static str) { + self.held.lock().push((endpoint_id, capability)); + } + + /// The caller the session acts as. + pub fn caller(&self) -> &Caller { + &self.caller + } + + /// Whether `other` is the same caller: an `auth` refresh may renew a credential, never swap + /// the principal a live session bills and authorizes against. + pub fn same_caller(&self, other: &CallerCheck) -> bool { + self.caller.check == *other + } +} + +/// Every deployment a Bud-mode `/ws` session holds is still reachable by its caller and still +/// serves what the session uses it for (D-17). +pub async fn session_still_allowed(state: &AppState, meter: &LegMeter) -> bool { + let Some(bud) = state.bud_mode.as_ref() else { + return false; + }; + let plane = bud.plane(); + for (ep, capability) in meter.endpoints() { + let serving = plane + .voice_endpoint(&ep) + .is_some_and(|e| e.serves(capability)); + let reachable = match meter.check() { + CallerCheck::ApiKey { hashed, client_key } => { + plane.hash_reaches(hashed, &ep, *client_key).is_some() + } + CallerCheck::Jwt { sub } => plane.subject_reaches(sub, &ep).await.is_some(), + }; + if !serving || !reachable { + return false; + } + } + true +} + +/// A Bud-mode `/ws` session's legs, resolved and admitted. +pub struct PreparedLegs { + pub stt_key: String, + pub tts_key: String, + /// The TTS deployment's own address (self-hosted, Azure OpenAI); never the client's (§5.9). + pub tts_api_base: Option, + pub meter: Arc, + /// One per leg deployment, held for the session (FRD-022 §6.2). + pub admissions: Vec, + /// What the client should hear about settings that did not apply. + pub advisories: Advisories, +} + +impl std::fmt::Debug for PreparedLegs { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PreparedLegs") + .field("meter", &self.meter) + .field("admissions", &self.admissions.len()) + .finish() + } +} + +/// Authenticate the session's caller, resolve both legs as deployments through the caller's +/// allowlist, admit each, and point the configs at them (FR-WS-1, TC-WS-01…05). +/// +/// A leg refused after the other was admitted releases that admission as it returns: nothing is +/// held for a session that never starts (TC-WS-05). +pub async fn prepare( + state: &Arc, + credential: Option<&crate::auth::SessionCredential>, + stt: &mut STTWebSocketConfig, + tts: &mut TTSWebSocketConfig, + session_id: &str, +) -> Result { + let Some(credential) = credential.map(|c| c.current()) else { + return Err(LegRefusal::new( + "authentication_required", + "This gateway serves Bud deployments, which need your Bud API key or token.", + ) + .closing(CLOSE_REVOKED)); + }; + let bearer = crate::handlers::openai_realtime::handshake::Credential::new( + credential.clone(), + crate::handlers::openai_realtime::handshake::CredentialSource::Bearer, + ); + let caller = crate::handlers::openai_realtime::session::authenticate(state, &bearer) + .await + .map_err(|e| { + let close = if e.status == axum::http::StatusCode::SERVICE_UNAVAILABLE + || e.status == axum::http::StatusCode::TOO_MANY_REQUESTS + { + CLOSE_TRY_LATER + } else { + CLOSE_REVOKED + }; + LegRefusal::new(e.code, e.message).closing(close) + })?; + + // A leg's credential is its deployment's. A key of the client's own would bypass attribution, + // quota and billing (FRD-018 §5.3.7), and is refused rather than silently replaced. + for (field, key) in [ + ("stt_config.api_key", &stt.api_key), + ("tts_config.api_key", &tts.api_key), + ] { + if key.as_deref().is_some_and(|k| !k.trim().is_empty()) { + return Err(LegRefusal::new( + "client_key_not_accepted", + format!( + "{field} is not accepted by this gateway: each leg uses its deployment's own \ + credential. Remove the field and name the deployment in `model`." + ), + )); + } + } + + let stt_leg = resolve_leg( + state, + &credential, + "stt_config.model", + &stt.model, + STT_CAPABILITY, + ) + .await?; + let tts_leg = resolve_leg( + state, + &credential, + "tts_config.model", + &tts.model, + TTS_CAPABILITY, + ) + .await?; + + let mut advisories = Advisories::new(); + let stt_key = apply_stt(stt, &stt_leg, &mut advisories); + let tts_key = apply_tts(tts, &tts_leg, &mut advisories); + resolve_tts_voice(state, tts, &tts_leg, &mut advisories).await; + + let meter = Arc::new(LegMeter::new( + session_id.to_string(), + caller, + Some((&stt_leg, stt.sample_rate, stt.channels)), + Some(&tts_leg), + )); + let tts_api_base = tts_leg.endpoint.api_base.clone(); + Ok(PreparedLegs { + stt_key, + tts_key, + tts_api_base, + meter, + admissions: vec![stt_leg.admission, tts_leg.admission], + advisories, + }) +} + +/// Headers a template may not send to the Bud gateway: the caller's own credential is the one. +const LLM_CREDENTIAL_HEADERS: &[&str] = &[ + "authorization", + "api-key", + "x-api-key", + "proxy-authorization", +]; + +/// Make a server DAG template's nodes address Bud deployments as the caller (FRD-023 WP-RT6.3, +/// TC-WS-10). +/// +/// * A TTS provider node's `model` names a text-to-speech deployment: resolved through the +/// caller's allowlist, admitted for the session, metered per synthesis and revalidated. +/// * An LLM or translate node's `model` names a chat deployment, reached through the Bud gateway +/// with the caller's credential; the template's `base_url`, `api_key` and credential headers are +/// dropped. +/// * An STT provider node marks where the session's STT leg injects its transcript; it never runs. +/// * A realtime provider node is refused: GA realtime on Bud deployments is the `/v1/realtime` +/// relay, and the native engines address deployments from RT7. +/// +/// Returns the admissions to hold for the session. +pub async fn bind_dag( + state: &Arc, + definition: &mut crate::dag::definition::DAGDefinition, + credential: &crate::auth::SessionCredential, + session_meter: &Arc, + session_id: &str, +) -> Result, LegRefusal> { + use crate::dag::definition::{NodeDefinition, NodeType}; + let raw = credential.current(); + let mut admissions = Vec::new(); + for node in definition.nodes.iter_mut() { + let NodeDefinition { + id, + node_type, + config, + bud, + .. + } = node; + match node_type { + NodeType::RealtimeProvider { .. } => { + return Err(LegRefusal::new( + "unsupported_node", + format!( + "DAG node '{id}' is a realtime provider, which this gateway does not run on \ + Bud deployments yet; use /v1/realtime with a realtime deployment." + ), + )); + } + NodeType::TtsProvider { + provider, + voice_id, + model, + } => { + let name = model.clone().unwrap_or_default(); + let field = format!("DAG node '{id}' model"); + let leg = resolve_leg(state, &raw, &field, &name, TTS_CAPABILITY).await?; + let ep = &leg.endpoint; + *provider = ep.vendor.clone(); + *model = ep.model.clone(); + if voice_id.as_deref().is_none_or(|v| v.trim().is_empty()) { + *voice_id = match non_empty(&ep.voice) { + Some(voice) => Some(voice), + None => { + crate::handlers::openai_audio::resolve_described_voice( + state, + ep, + &mut Advisories::new(), + ) + .await + } + }; + } + if let Some(obj) = config.as_object_mut() { + obj.remove("api_key"); + } + let meter = Arc::new(LegMeter::new( + session_id.to_string(), + session_meter.caller().clone(), + None, + Some(&leg), + )); + session_meter.hold(leg.endpoint_id.clone(), TTS_CAPABILITY); + *bud = Some(Arc::new(crate::dag::nodes::BudNodeBinding { + endpoint_id: leg.endpoint_id.clone(), + vendor_credential: ep.credential.clone(), + api_base: ep.api_base.clone(), + extras: crate::handlers::openai_audio::deployment_extras(ep), + session_credential: None, + on_synthesis: Some(Arc::new(move |text: &str| meter.tts_spoken(text))), + })); + admissions.push(leg.admission); + } + NodeType::LlmEndpoint { + base_url, + model, + api_key, + headers, + .. + } + | NodeType::Translate { + base_url, + model, + api_key, + headers, + .. + } => { + let Some(url) = llm_base_url() else { + return Err(LegRefusal::new( + "llm_unavailable", + format!( + "DAG node '{id}' needs the Bud gateway (WAAV_LLM_BASE_URL), which this \ + gateway is not configured with." + ), + )); + }; + *base_url = url; + *api_key = None; + headers.retain(|k, _| { + !LLM_CREDENTIAL_HEADERS.contains(&k.to_ascii_lowercase().as_str()) + }); + *bud = Some(Arc::new(crate::dag::nodes::BudNodeBinding { + endpoint_id: model.clone(), + vendor_credential: None, + api_base: None, + extras: serde_json::Map::new(), + session_credential: Some(credential.clone()), + on_synthesis: None, + })); + } + _ => {} + } + } + Ok(admissions) +} + +/// The budgateway base URL for the voice agent's LLM leg (`WAAV_LLM_BASE_URL`, F-8). +pub fn llm_base_url() -> Option { + std::env::var("WAAV_LLM_BASE_URL") + .ok() + .map(|v| v.trim().trim_end_matches('/').to_string()) + .filter(|v| !v.is_empty()) +} + +#[allow(dead_code)] +fn _assert_send_sync() { + fn is() {} + is::(); + is::>(); +} + +#[cfg(test)] +mod tests { + //! FRD-023 RT6 over an in-memory control plane: TC-WS-01…05, 10…13 and the leg security + //! rules. The vendors are never dialled — a leg's job ends at a config pointed at the + //! deployment, which is what these assert; the live streams are the pde-ditto E2E. + + use super::*; + use crate::test_support::{TEST_CREDENTIAL, TEST_CREDENTIAL_PLAIN, bud_state_with_credentials}; + use serde_json::{Value as Json, json}; + use std::collections::HashMap; + use std::sync::Mutex; + use tracing::field::{Field, Visit}; + use tracing::span::{Attributes, Id, Record}; + use tracing_subscriber::Layer; + use tracing_subscriber::layer::{Context, SubscriberExt}; + use tracing_subscriber::registry::LookupSpan; + + const KEY: &str = "bud_ws_legs_test_key"; + const OTHER_KEY: &str = "bud_ws_legs_other_key"; + const PROJECT: &str = "7c0c7e1d-0000-4000-8000-00000000aa01"; + const OTHER_PROJECT: &str = "7c0c7e1d-0000-4000-8000-00000000aa02"; + const USER: &str = "7c0c7e1d-0000-4000-8000-00000000bb01"; + const API_KEY_ID: &str = "7c0c7e1d-0000-4000-8000-00000000cc01"; + const MODEL_ID: &str = "7c0c7e1d-0000-4000-8000-00000000dd01"; + + fn stt_entry(extra: Json) -> Json { + merge( + json!({ + "vendor": "deepgram", + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["audio_transcription"], + "model": "nova-3", + "language": "en-US", + "pricing": {"unit": "minute", "per_units": 1, "cost_per_unit": 0.0043, "currency": "USD"}, + "config": {"stt": {"diarization": true, + "streaming": {"interim_results": true, "endpointing_ms": 300, "utterance_end_ms": 1000}}} + }), + extra, + ) + } + + fn tts_entry(extra: Json) -> Json { + merge( + json!({ + "vendor": "elevenlabs", + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["text_to_speech"], + "model": "eleven_flash_v2_5", + "voice": "JBFqnCBsd6RMkjVDRZzb", + "pricing": {"unit": "character", "per_units": 1000, "cost_per_unit": 0.3, "currency": "USD"}, + "config": {"tts": {"pronunciations": [{"word": "Bud", "pronunciation": "bʌd"}]}} + }), + extra, + ) + } + + fn merge(mut base: Json, extra: Json) -> Json { + if let Some(obj) = extra.as_object() { + for (k, v) in obj { + base[k] = v.clone(); + } + } + base + } + + /// `(alias, endpoint id, entry)`; KEY reaches `mine`, OTHER_KEY reaches `theirs`. + async fn plane( + mine: &[(&str, &str, Json)], + theirs: &[(&str, &str, Json)], + ) -> (Arc, Arc) { + let blob = |eps: &[(&str, &str, Json)], project: &str| { + let mut m = serde_json::Map::new(); + for (alias, id, _) in eps { + m.insert( + alias.to_string(), + json!({"endpoint_id": id, "model_id": MODEL_ID, "project_id": project, "kind": "model"}), + ); + } + m.insert( + "__metadata__".into(), + json!({"api_key_id": API_KEY_ID, "user_id": USER, "api_key_project_id": project}), + ); + Json::Object(m).to_string() + }; + let mut keys: Vec<(String, String)> = vec![ + ( + format!("api_key:{}", bud_auth::hash_api_key(KEY)), + blob(mine, PROJECT), + ), + ( + format!("api_key:{}", bud_auth::hash_api_key(OTHER_KEY)), + blob(theirs, OTHER_PROJECT), + ), + ]; + for (_, id, entry) in mine.iter().chain(theirs.iter()) { + keys.push(( + format!("voice_table:{id}"), + json!({ *id: entry }).to_string(), + )); + } + let refs: Vec<(&str, &str)> = keys.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect(); + bud_state_with_credentials(&refs).await + } + + fn stt_cfg(v: Json) -> STTWebSocketConfig { + serde_json::from_value(merge( + json!({"language": "en-US", "sample_rate": 16000, "channels": 1, "punctuation": true}), + v, + )) + .expect("a Bud client may omit provider") + } + + fn tts_cfg(v: Json) -> TTSWebSocketConfig { + serde_json::from_value(merge( + json!({"voice_id": null, "speaking_rate": null, "audio_format": "linear16", + "sample_rate": 24000, "connection_timeout": null, "request_timeout": null}), + v, + )) + .expect("a Bud client may omit provider") + } + + fn cred(raw: &str) -> crate::auth::SessionCredential { + crate::auth::SessionCredential::new(raw) + } + + async fn default_plane() -> (Arc, Arc) { + plane( + &[ + ("stt-dg", "ep-stt", stt_entry(json!({}))), + ("tts-el", "ep-tts", tts_entry(json!({}))), + ], + &[("their-tts", "ep-their-tts", tts_entry(json!({})))], + ) + .await + } + + /// TC-WS-01 🔒 / TC-WS-02 🔒 — each leg is its deployment: vendor, model, credential and voice + /// from voice_table; the client's `provider` is informational. + #[tokio::test] + async fn tc_ws_01_02_legs_take_the_deployment_vendor_model_and_credential() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"provider": "assemblyai", "model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .expect("both legs resolve"); + + assert_eq!(stt.provider, "deepgram"); + assert_eq!(stt.model, "nova-3"); + assert_eq!(legs.stt_key, TEST_CREDENTIAL_PLAIN); + assert_eq!(tts.provider, "elevenlabs"); + assert_eq!(tts.model, "eleven_flash_v2_5"); + assert_eq!(tts.voice_id.as_deref(), Some("JBFqnCBsd6RMkjVDRZzb")); + assert_eq!(legs.tts_key, TEST_CREDENTIAL_PLAIN); + assert_eq!( + tts.pronunciations.len(), + 1, + "the deployment's pronunciations apply" + ); + assert!(stt.api_key.is_none() && tts.api_key.is_none()); + assert_eq!(legs.admissions.len(), 2, "one admission per leg"); + } + + /// A deployment with no voice, on a vendor that requires one, speaks with the vendor's default + /// and says so (as REST does). + #[tokio::test] + async fn a_voiceless_deployment_uses_the_vendor_default_and_says_so() { + let (state, _) = plane( + &[ + ("stt-dg", "ep-stt", stt_entry(json!({}))), + ("tts-el", "ep-tts", tts_entry(json!({"voice": null}))), + ], + &[], + ) + .await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap(); + assert_eq!( + tts.voice_id.as_deref(), + Some(crate::handlers::voices::provider_default_voice( + "elevenlabs" + )) + ); + assert!( + legs.advisories + .as_slice() + .iter() + .any(|a| a.contains("no voice configured")), + "{:?}", + legs.advisories.as_slice() + ); + } + + /// TC-WS-02 — request > deployment: the client's own voice wins. + #[tokio::test] + async fn tc_ws_02_the_clients_voice_wins() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el", "voice_id": "EXAVITQu4vr4xnSDxMaL"})); + prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap(); + assert_eq!(tts.voice_id.as_deref(), Some("EXAVITQu4vr4xnSDxMaL")); + } + + /// TC-WS-03 🔒 — a self-hosted TTS deployment's `api_base` is used (it was forced `None` on + /// the socket path). + #[tokio::test] + async fn tc_ws_03_the_deployment_api_base_is_used() { + let (state, _) = plane( + &[ + ("stt-dg", "ep-stt", stt_entry(json!({}))), + ( + "tts-own", + "ep-own", + tts_entry( + json!({"vendor": "self_hosted", "api_base": "http://tts.internal:8000", + "model": "kokoro", "voice": "af_heart"}), + ), + ), + ], + &[], + ) + .await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-own"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap(); + assert_eq!( + legs.tts_api_base.as_deref(), + Some("http://tts.internal:8000") + ); + assert_eq!(tts.provider, "self_hosted"); + } + + /// 🔒 A client's `extras` never ride a Bud leg: Groq reads a destination from `extras.url`, + /// Azure Speech its host, and the credential that would travel there is the deployment's. + #[tokio::test] + async fn client_extras_are_replaced_by_the_deployments() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg", + "extras": {"url": "https://attacker.example/listen", "endpoint_override": "wss://attacker.example"}})); + let mut tts = + tts_cfg(json!({"model": "tts-el", "extras": {"url": "https://attacker.example"}})); + prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap(); + assert!( + !serde_json::to_string(&stt.extras.0) + .unwrap() + .contains("attacker") + && !serde_json::to_string(&tts.extras.0) + .unwrap() + .contains("attacker"), + "stt extras {:?}, tts extras {:?}", + stt.extras.0, + tts.extras.0 + ); + } + + /// TC-WS-04 🔒 — `provider` without a deployment `model` is refused with the addressing hint; + /// a name the caller cannot reach is not found. + #[tokio::test] + async fn tc_ws_04_provider_only_and_unreachable_names_are_refused() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"provider": "deepgram"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "deployment_required"); + assert!( + refusal.message.contains("stt_config.model"), + "{}", + refusal.message + ); + assert_eq!( + refusal.close, None, + "the client may send a corrected config" + ); + + // Another project's deployment, by alias or by endpoint id. + for name in ["their-tts", "ep-their-tts"] { + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": name})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "model_not_found", "{name}"); + } + } + + /// 🔒 A client's own vendor key is refused on a Bud leg, not silently replaced. + #[tokio::test] + async fn a_client_vendor_key_is_refused_by_name() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el", "api_key": "sk-byok"})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "client_key_not_accepted"); + assert!( + refusal.message.contains("tts_config.api_key"), + "{}", + refusal.message + ); + } + + /// A leg's capability is checked: a TTS deployment is not an STT leg. + #[tokio::test] + async fn a_deployment_of_the_wrong_capability_is_not_found() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "tts-el"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "model_not_found"); + } + + /// A self-hosted or Azure OpenAI transcription deployment answers uploads, not streams. + #[tokio::test] + async fn an_upload_only_stt_deployment_is_refused_by_name() { + let (state, _) = plane( + &[ + ( + "whisper", + "ep-whisper", + stt_entry(json!({"vendor": "self_hosted", "api_base": "http://whisper:8000"})), + ), + ("tts-el", "ep-tts", tts_entry(json!({}))), + ], + &[], + ) + .await; + let mut stt = stt_cfg(json!({"model": "whisper"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "unsupported_deployment"); + assert!(refusal.message.contains("/v1/audio/transcriptions")); + } + + /// 🔒 An AWS deployment without its key pair would authenticate as the GATEWAY's identity; it + /// is refused as REST refuses it. + #[tokio::test] + async fn an_aws_deployment_without_its_key_pair_is_refused() { + let (state, _) = plane( + &[ + ("stt-dg", "ep-stt", stt_entry(json!({}))), + ( + "polly", + "ep-polly", + tts_entry(json!({"vendor": "aws_polly", "credential": null, + "provider_params": {"region": "us-east-1"}})), + ), + ], + &[], + ) + .await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "polly"})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "deployment_misconfigured", "{refusal:?}"); + } + + /// TC-WS-05 🔒 — a leg at its concurrency cap refuses the session with 1013, and the other + /// leg's admission is released rather than leaked. + #[tokio::test] + async fn tc_ws_05_per_leg_admission_refuses_1013_and_releases_the_other_leg() { + let (state, _) = plane( + &[ + ("stt-dg", "ep-stt", stt_entry(json!({"max_concurrent": 1}))), + ("tts-el", "ep-tts", tts_entry(json!({"max_concurrent": 1}))), + ], + &[], + ) + .await; + let held_tts = state.admit_deployment("ep-tts").await.expect("free"); + + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let refusal = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.code, "concurrency_limit_exceeded"); + assert_eq!(refusal.close, Some(CLOSE_TRY_LATER)); + let stt_again = state.admit_deployment("ep-stt").await; + assert!(stt_again.is_ok(), "the STT leg's slot was released"); + drop(stt_again); + drop(held_tts); + + // And a session holds both slots for its whole life. + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-2") + .await + .unwrap(); + assert!(state.admit_deployment("ep-stt").await.is_err()); + drop(legs); + assert!(state.admit_deployment("ep-stt").await.is_ok()); + } + + /// No credential, or one the plane rejects: refused and closed. + #[tokio::test] + async fn an_unauthenticated_session_is_refused_and_closed() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let refusal = prepare(&state, None, &mut stt, &mut tts, "s-1") + .await + .unwrap_err(); + assert_eq!(refusal.close, Some(CLOSE_REVOKED)); + let refusal = prepare( + &state, + Some(&cred("bud_not_a_key")), + &mut stt, + &mut tts, + "s-1", + ) + .await + .unwrap_err(); + assert_eq!(refusal.code, "invalid_api_key"); + assert_eq!(refusal.close, Some(CLOSE_REVOKED)); + } + + /// TC-WS-11 🔒 — `stt.streaming` applies on `/ws`; the client's explicit choice wins; the + /// batch features apply too. + #[tokio::test] + async fn tc_ws_11_streaming_features_apply_on_ws() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg", "features": {"endpointing_ms": 800}})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s-1") + .await + .unwrap(); + assert_eq!(stt.features.interim_results, Some(true)); + assert_eq!(stt.features.utterance_end_ms, Some(1000)); + assert_eq!( + stt.features.endpointing_ms, + Some(800), + "request > deployment" + ); + assert_eq!(stt.features.diarization, Some(true)); + assert_eq!(stt.features.vad_events, None, "unset stays unset on /ws"); + } + + /// TC-WS-11 (REST half) — the upload path ignores `stt.streaming`. + #[test] + fn tc_ws_11_rest_ignores_streaming_features() { + let settings: bud_auth::endpoint_config::SttSettings = serde_json::from_value( + json!({"streaming": {"interim_results": true, "endpointing_ms": 300}}), + ) + .unwrap(); + let features = + endpoint_settings::stt_features_for(&settings, "deepgram", &mut Advisories::new()); + assert_eq!(features.interim_results, None); + assert_eq!(features.endpointing_ms, None); + } + + // --------------------------------------------------------------------------------------- + // TC-WS-12 — metering + // --------------------------------------------------------------------------------------- + + #[derive(Clone, Default)] + struct Spans(Arc)>>>); + + struct V<'a>(&'a mut HashMap); + impl Visit for V<'_> { + fn record_debug(&mut self, f: &Field, v: &dyn std::fmt::Debug) { + self.0.insert(f.name().to_string(), format!("{v:?}")); + } + fn record_str(&mut self, f: &Field, v: &str) { + self.0.insert(f.name().to_string(), v.to_string()); + } + } + + struct Capture(Spans); + impl Layer for Capture + where + S: tracing::Subscriber + for<'a> LookupSpan<'a>, + { + fn on_new_span(&self, attrs: &Attributes<'_>, id: &Id, ctx: Context<'_, S>) { + let name = ctx + .span(id) + .map(|s| s.name().to_string()) + .unwrap_or_default(); + let mut fields = HashMap::new(); + attrs.record(&mut V(&mut fields)); + self.0.0.lock().unwrap().push((id.into_u64(), name, fields)); + } + fn on_record(&self, id: &Id, values: &Record<'_>, _ctx: Context<'_, S>) { + let mut all = self.0.0.lock().unwrap(); + if let Some((_, _, fields)) = all.iter_mut().rev().find(|(i, _, _)| *i == id.into_u64()) + { + values.record(&mut V(fields)); + } + } + } + + fn turns(spans: &Spans) -> Vec> { + spans + .0 + .lock() + .unwrap() + .iter() + .filter(|(_, name, _)| name == "voice.turn") + .map(|(_, _, f)| f.clone()) + .collect() + } + + /// TC-WS-12 🔒 — one utterance and one `speak`: an STT record with seconds and cost, a TTS + /// record with characters and cost, each attributed to the caller and the deployment. + #[tokio::test] + async fn tc_ws_12_each_leg_is_metered_and_attributed() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "sess-12") + .await + .unwrap(); + let meter = legs.meter.clone(); + + let spans = Spans::default(); + let subscriber = tracing_subscriber::registry().with(Capture(spans.clone())); + tracing::subscriber::with_default(subscriber, || { + // 1.5 s of 16 kHz mono PCM16, then its final transcript. + meter.add_stt_audio(16_000 * 2 * 3 / 2); + meter.stt_final("hello there"); + meter.stt_final(""); + meter.tts_spoken("Hello from Bud!"); + meter.tts_spoken(""); + }); + + let records = turns(&spans); + assert_eq!( + records.len(), + 2, + "an empty final or speak bills nothing: {records:?}" + ); + let stt_rec = &records[0]; + assert_eq!(stt_rec["bud.voice.capability"], "audio_transcription"); + assert_eq!(stt_rec["bud.voice.audio_seconds"], "1.5"); + let cost: f64 = stt_rec["bud.voice.cost"].parse().unwrap(); + assert!((cost - 1.5 / 60.0 * 0.0043).abs() < 1e-12, "{cost}"); + assert_eq!(stt_rec["bud.voice.pricing_unit"], "minute"); + assert_eq!(stt_rec["bud.endpoint_id"], "ep-stt"); + assert_eq!(stt_rec["bud.voice.endpoint_name"], "stt-dg"); + assert_eq!(stt_rec["bud.project_id"], PROJECT); + assert_eq!(stt_rec["bud.api_key_id"], API_KEY_ID); + assert_eq!(stt_rec["bud.voice.session_id"], "sess-12"); + + let tts_rec = &records[1]; + assert_eq!(tts_rec["bud.voice.capability"], "text_to_speech"); + assert_eq!(tts_rec["bud.voice.characters"], "15"); + let cost: f64 = tts_rec["bud.voice.cost"].parse().unwrap(); + assert!((cost - 15.0 / 1000.0 * 0.3).abs() < 1e-12, "{cost}"); + assert_eq!(tts_rec["bud.endpoint_id"], "ep-tts"); + assert_ne!( + stt_rec["bud.voice.turn_index"], + tts_rec["bud.voice.turn_index"] + ); + } + + /// TC-WS-12 — audio streamed after the last final is billed at close; nothing is billed twice. + #[tokio::test] + async fn tc_ws_12_audio_after_the_last_final_is_billed_at_close() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + let meter = legs.meter.clone(); + // The codec negotiation moved the session to 48 kHz stereo. + meter.set_stt_format(48_000, 2); + let spans = Spans::default(); + let subscriber = tracing_subscriber::registry().with(Capture(spans.clone())); + tracing::subscriber::with_default(subscriber, || { + meter.add_stt_audio(48_000 * 2 * 2); + meter.finish(); + meter.finish(); + }); + let records = turns(&spans); + assert_eq!(records.len(), 1); + assert_eq!(records[0]["bud.voice.audio_seconds"], "1.0"); + } + + // --------------------------------------------------------------------------------------- + // TC-WS-13 — revalidation + // --------------------------------------------------------------------------------------- + + /// TC-WS-13 🔒 — a revoked key, and an unpublished deployment, fail revalidation. + #[tokio::test] + async fn tc_ws_13_revocation_and_unpublishing_fail_revalidation() { + let (state, store) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + assert!(session_still_allowed(&state, &legs.meter).await); + + let plane = state.bud_mode.as_ref().unwrap().plane().clone(); + store.remove("voice_table:ep-tts"); + plane + .on_key_event("voice_table:ep-tts", bud_auth::KeyEvent::Del) + .await + .unwrap(); + assert!( + !session_still_allowed(&state, &legs.meter).await, + "an unpublished leg deployment ends the session" + ); + + let (state, store) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + let key = format!("api_key:{}", bud_auth::hash_api_key(KEY)); + store.remove(&key); + state + .bud_mode + .as_ref() + .unwrap() + .plane() + .on_key_event(&key, bud_auth::KeyEvent::Del) + .await + .unwrap(); + assert!( + !session_still_allowed(&state, &legs.meter).await, + "a revoked key" + ); + } + + /// TC-WS-13 🔒 — a leg deployment republished without the capability the session uses it for + /// (still in the table, no longer a text-to-speech deployment) fails revalidation. + #[tokio::test] + async fn tc_ws_13_a_leg_that_no_longer_serves_its_capability_fails_revalidation() { + let (state, store) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + assert!(session_still_allowed(&state, &legs.meter).await); + + let republished = tts_entry(json!({"endpoints": ["audio_transcription"]})); + store.set( + "voice_table:ep-tts", + &json!({ "ep-tts": republished }).to_string(), + ); + state + .bud_mode + .as_ref() + .unwrap() + .plane() + .on_key_event("voice_table:ep-tts", bud_auth::KeyEvent::Set) + .await + .unwrap(); + assert!( + state + .bud_mode + .as_ref() + .unwrap() + .plane() + .voice_endpoint("ep-tts") + .is_some(), + "the deployment is still published" + ); + assert!( + !session_still_allowed(&state, &legs.meter).await, + "but no longer serves text_to_speech" + ); + } + + // --------------------------------------------------------------------------------------- + // TC-WS-10 — DAG templates + // --------------------------------------------------------------------------------------- + + fn template(nodes: Json) -> crate::dag::definition::DAGDefinition { + serde_json::from_value(json!({ + "id": "t", "name": "t", "nodes": nodes, "edges": [], + "entry_node": "in", "exit_nodes": [] + })) + .expect("template parses") + } + + /// TC-WS-10 🔒 — a template's TTS node resolves a deployment (vendor, model, voice, + /// credential, admission, revalidation); its LLM node goes to the Bud gateway with the + /// caller's credential and none of the template's own. + #[tokio::test] + #[serial_test::serial] + async fn tc_ws_10_template_nodes_address_deployments() { + unsafe { std::env::set_var("WAAV_LLM_BASE_URL", "http://bud-gateway:3000/v1/") }; + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + let mut def = template(json!([ + {"id": "in", "type": "audio_input"}, + {"id": "speak", "type": "tts_provider", "provider": "cartesia", "model": "tts-el", + "config": {"api_key": "${CARTESIA_API_KEY}"}}, + {"id": "think", "type": "llm_endpoint", "base_url": "https://api.openai.com/v1", + "model": "chat-deployment", "api_key": "${OPENAI_API_KEY}", + "headers": {"Authorization": "Bearer sk-template", "X-Trace": "keep"}} + ])); + let session_credential = cred(KEY); + let admissions = bind_dag(&state, &mut def, &session_credential, &legs.meter, "s") + .await + .expect("binds"); + unsafe { std::env::remove_var("WAAV_LLM_BASE_URL") }; + + assert_eq!(admissions.len(), 1, "the TTS node's deployment is admitted"); + let speak = &def.nodes[1]; + match &speak.node_type { + crate::dag::definition::NodeType::TtsProvider { + provider, + model, + voice_id, + } => { + assert_eq!(provider, "elevenlabs"); + assert_eq!(model.as_deref(), Some("eleven_flash_v2_5")); + assert_eq!(voice_id.as_deref(), Some("JBFqnCBsd6RMkjVDRZzb")); + } + other => panic!("{other:?}"), + } + let binding = speak.bud.as_ref().expect("bound"); + assert_eq!( + binding.vendor_credential.as_deref(), + Some(TEST_CREDENTIAL_PLAIN) + ); + assert!(speak.config.get("api_key").is_none()); + assert!( + legs.meter + .endpoints() + .iter() + .filter(|(id, _)| id == "ep-tts") + .count() + == 2, + "the node's deployment is revalidated with the session" + ); + + let think = &def.nodes[2]; + match &think.node_type { + crate::dag::definition::NodeType::LlmEndpoint { + base_url, + api_key, + headers, + model, + .. + } => { + assert_eq!(base_url, "http://bud-gateway:3000/v1"); + assert!(api_key.is_none()); + assert_eq!(model, "chat-deployment"); + assert!( + !headers + .keys() + .any(|k| k.eq_ignore_ascii_case("authorization")) + ); + assert_eq!(headers.get("X-Trace").map(String::as_str), Some("keep")); + } + other => panic!("{other:?}"), + } + let binding = think.bud.as_ref().expect("bound"); + session_credential.replace("bud_refreshed"); + assert_eq!( + binding + .session_credential + .as_ref() + .map(|c| c.current()) + .as_deref(), + Some("bud_refreshed"), + "the node reads the session's credential per call" + ); + } + + /// TC-WS-10 — a template node naming a deployment the caller cannot reach, or a realtime + /// provider node, is refused. + #[tokio::test] + #[serial_test::serial] + async fn tc_ws_10_unreachable_and_realtime_nodes_are_refused() { + unsafe { std::env::set_var("WAAV_LLM_BASE_URL", "http://bud-gateway:3000/v1") }; + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + let mut def = template(json!([ + {"id": "speak", "type": "tts_provider", "provider": "elevenlabs", "model": "their-tts"} + ])); + let refusal = bind_dag(&state, &mut def, &cred(KEY), &legs.meter, "s") + .await + .unwrap_err(); + assert_eq!(refusal.code, "model_not_found"); + + let mut def = template(json!([ + {"id": "rt", "type": "realtime_provider", "provider": "openai", "model": "gpt-realtime"} + ])); + let refusal = bind_dag(&state, &mut def, &cred(KEY), &legs.meter, "s") + .await + .unwrap_err(); + unsafe { std::env::remove_var("WAAV_LLM_BASE_URL") }; + assert_eq!(refusal.code, "unsupported_node"); + } + + /// The same-caller rule an `auth` refresh is held to (TC-WS-08). + #[tokio::test] + async fn the_meter_knows_its_caller() { + let (state, _) = default_plane().await; + let mut stt = stt_cfg(json!({"model": "stt-dg"})); + let mut tts = tts_cfg(json!({"model": "tts-el"})); + let legs = prepare(&state, Some(&cred(KEY)), &mut stt, &mut tts, "s") + .await + .unwrap(); + let same = CallerCheck::ApiKey { + hashed: bud_auth::hash_api_key(KEY), + client_key: false, + }; + let other = CallerCheck::ApiKey { + hashed: bud_auth::hash_api_key(OTHER_KEY), + client_key: false, + }; + assert!(legs.meter.same_caller(&same)); + assert!(!legs.meter.same_caller(&other)); + } +} diff --git a/gateway/src/handlers/ws/config.rs b/gateway/src/handlers/ws/config.rs index 1150b818..2f1c07cf 100644 --- a/gateway/src/handlers/ws/config.rs +++ b/gateway/src/handlers/ws/config.rs @@ -62,8 +62,11 @@ pub struct DAGWebSocketConfig { #[derive(Debug, Deserialize, Serialize, Clone)] #[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] pub struct ConversationWebSocketConfig { - /// OpenAI-compatible base URL for the LLM (e.g. `https://api.openai.com/v1`). + /// OpenAI-compatible base URL for the LLM (e.g. `https://api.openai.com/v1`). Omitted — and + /// refused if present — under the Bud control plane, where the LLM leg is a Bud chat + /// deployment reached through the Bud gateway (FRD-023 RT6). #[cfg_attr(feature = "openapi", schema(example = "https://api.openai.com/v1"))] + #[serde(default)] pub base_url: String, /// Model identifier. @@ -251,6 +254,9 @@ impl ConversationWebSocketConfig { model: self.model.clone(), system_prompt: self.system_prompt.clone(), api_key: self.api_key.clone(), + server_llm_endpoint: false, + credential: None, + attribution: None, temperature: self.temperature, max_tokens: self.max_tokens, streaming: self.streaming, @@ -307,7 +313,9 @@ pub fn default_stt_encoding() -> String { #[derive(Debug, Deserialize, Serialize, Clone)] #[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] pub struct STTWebSocketConfig { - /// Provider name (e.g., "deepgram") + /// Provider name (e.g., "deepgram"). Under the Bud control plane it is taken from the + /// deployment `model` names, and may be omitted (FRD-023 RT6). + #[serde(default)] #[cfg_attr(feature = "openapi", schema(example = "deepgram"))] pub provider: String, /// Language code for transcription (e.g., "en-US", "es-ES") @@ -533,7 +541,9 @@ impl LiveKitWebSocketConfig { #[derive(Debug, Deserialize, Serialize, Clone)] #[cfg_attr(feature = "openapi", derive(utoipa::ToSchema))] pub struct TTSWebSocketConfig { - /// Provider name (e.g., "deepgram", "hume", "elevenlabs") + /// Provider name (e.g., "deepgram", "hume", "elevenlabs"). Under the Bud control plane it is + /// taken from the deployment `model` names, and may be omitted (FRD-023 RT6). + #[serde(default)] #[cfg_attr(feature = "openapi", schema(example = "deepgram"))] pub provider: String, /// Voice ID or name to use for synthesis. @@ -591,7 +601,8 @@ pub struct TTSWebSocketConfig { /// Request timeout in seconds #[cfg_attr(feature = "openapi", schema(example = 60))] pub request_timeout: Option, - /// Model to use for TTS + /// Model to use for TTS. Under the Bud control plane: the text-to-speech DEPLOYMENT. + #[serde(default)] #[cfg_attr(feature = "openapi", schema(example = "aura-asteria-en"))] pub model: String, /// Pronunciation replacements to apply before TTS diff --git a/gateway/src/handlers/ws/config_handler.rs b/gateway/src/handlers/ws/config_handler.rs index 31abbb25..3bfdc06c 100644 --- a/gateway/src/handlers/ws/config_handler.rs +++ b/gateway/src/handlers/ws/config_handler.rs @@ -246,10 +246,6 @@ pub async fn handle_config_message( // wins; on no catalog match the resolver returns the provider default + a // non-fatal `config_warning` (never a 400). The resolved id is set on the config // and thus echoed in the `ready` ack. - if let Some(tts) = tts_ws_config.as_mut() { - resolve_voice_descriptor(tts, app_state, message_tx).await; - } - // Generate stream_id if not provided by client let stream_id = resolve_stream_id(stream_id); info!("Session stream_id: {}", stream_id); @@ -263,6 +259,59 @@ pub async fn handle_config_message( livekit_ws_config.is_some() ); + // FRD-023 RT6: under the Bud control plane each leg addresses a DEPLOYMENT, resolved through + // the caller's allowlist and admitted once for the session; its vendor credential, model, + // api_base and voice come from voice_table (the voice read with the deployment's own key, so + // P4 below has nothing left to do for it). + let mut bud_legs = if app_state.bud_mode.is_some() && audio_enabled { + match prepare_bud_legs( + app_state, + state, + &mut stt_ws_config, + &mut tts_ws_config, + &stream_id, + ) + .await + { + Ok(legs) => legs, + Err(refusal) => { + warn!(code = refusal.code, "Refusing a Bud-mode /ws leg"); + send_error(message_tx, format!("{}: {}", refusal.code, refusal.message)).await; + if let Some(code) = refusal.close { + send_critical( + message_tx, + MessageRoute::CloseWith { + code, + reason: refusal.code.to_string(), + }, + ) + .await; + return false; + } + return true; + } + } + } else { + None + }; + if let Some(legs) = bud_legs.as_mut() { + for advisory in std::mem::take(&mut legs.advisories).as_slice() { + send_config_warning( + message_tx, + "deployment_setting_not_applied", + advisory.clone(), + None, + ) + .await; + } + } + + if bud_legs.is_none() + && let Some(tts) = tts_ws_config.as_mut() + { + resolve_voice_descriptor(tts, app_state, message_tx).await; + } + // Validate required configurations when audio is enabled if audio_enabled && !validate_audio_configs(&stt_ws_config, &tts_ws_config, message_tx).await { return true; @@ -316,8 +365,32 @@ pub async fn handle_config_message( return true; }; - match initialize_voice_manager(stt_config, tts_config, app_state, message_tx).await { + match initialize_voice_manager( + stt_config, + tts_config, + app_state, + message_tx, + bud_legs.as_ref(), + ) + .await + { Some(vm) => { + // FRD-023 RT6: the session now holds its legs' admissions; meter each leg (STT + // per final transcript, TTS per synthesis) and revalidate the caller. + if let Some(legs) = bud_legs.take() { + let meter = Arc::clone(&legs.meter); + meter.set_stt_format(stt_config.sample_rate, stt_config.channels); + { + let mut guard = state.write().await; + guard.leg_meter = Some(Arc::clone(&meter)); + guard.leg_admissions = legs.admissions; + } + let m = Arc::clone(&meter); + vm.set_speak_observer(Arc::new(move |text: &str| m.tts_spoken(text))); + let m = Arc::clone(&meter); + vm.set_stt_final_observer(Arc::new(move |text: &str| m.stt_final(text))); + spawn_bud_revalidation(app_state, state, message_tx, meter).await; + } let heartbeat_period = match heartbeat_period_from_env() { Ok(period) => period, Err(e) => { @@ -525,6 +598,7 @@ pub async fn handle_config_message( app_state.core_state.profiler.clone(), egress_audio.clone(), Some(app_state.core_state.resilience().clone()), + app_state, ) .await { @@ -575,7 +649,31 @@ pub async fn handle_config_message( } else if let (Some(conv_config), Some(vm)) = (conversation_ws_config.as_ref(), voice_manager.as_ref()) { - match initialize_conversation_loop(conv_config, &stream_id, vm, message_tx).await { + let (credential, attribution) = { + let guard = state.read().await; + let attribution = guard.leg_meter.as_ref().map(|m| { + let p = &m.caller().principal; + crate::core::conversation::TurnAttribution { + project_id: p.project_id.clone(), + api_key_id: p.api_key_id.clone(), + api_key_project_id: p.project_id.clone(), + user_id: p.user_id.clone(), + endpoint_name: Some(conv_config.model.clone()), + } + }); + (guard.credential.clone(), attribution) + }; + match initialize_conversation_loop( + conv_config, + &stream_id, + vm, + message_tx, + app_state, + credential, + attribution, + ) + .await + { Ok(true) => { info!("Conversation loop initialized for stream {}", stream_id); emit_reasoning_config_warnings(conv_config, message_tx).await; @@ -678,6 +776,163 @@ pub(crate) fn bud_mode_config_refusal( None } +/// Resolve a Bud-mode session's STT and TTS legs as deployments (FRD-023 RT6, FR-WS-1). +/// +/// `Ok(None)` when a leg config is missing: `validate_audio_configs` refuses that, by name. +async fn prepare_bud_legs( + app_state: &Arc, + state: &Arc>, + stt: &mut Option, + tts: &mut Option, + stream_id: &str, +) -> Result, super::bud_legs::LegRefusal> { + let (Some(stt), Some(tts)) = (stt.as_mut(), tts.as_mut()) else { + return Ok(None); + }; + let credential = state.read().await.credential.clone(); + super::bud_legs::prepare(app_state, credential.as_ref(), stt, tts, stream_id) + .await + .map(Some) +} + +/// Re-check a Bud-mode session's caller against its legs' deployments every +/// `WAAV_REALTIME_REVALIDATE_SECS` (D-17, TC-WS-13): a revoked key, a user removed from the +/// project, or an unpublished deployment ends the session with 1008. +async fn spawn_bud_revalidation( + app_state: &Arc, + state: &Arc>, + message_tx: &mpsc::Sender, + meter: Arc, +) { + let every = app_state.realtime.timings.revalidate; + let app_state = Arc::clone(app_state); + let session = Arc::downgrade(state); + let tx = message_tx.clone(); + let tracker = state.read().await.task_tracker.clone(); + let handle = tokio::spawn(async move { + let mut tick = tokio::time::interval(every); + tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + tick.tick().await; + loop { + tokio::select! { + _ = tick.tick() => {} + _ = tx.closed() => return, + } + // The session ended (or, defensively, holds other legs now). + let Some(session) = session.upgrade() else { + return; + }; + let current = session + .read() + .await + .leg_meter + .as_ref() + .is_some_and(|m| Arc::ptr_eq(m, &meter)); + drop(session); + if !current { + return; + } + if !super::bud_legs::session_still_allowed(&app_state, &meter).await { + warn!("Bud-mode /ws session no longer authorized; closing"); + send_error( + &tx, + "session_revoked: this session's credential or deployment is no longer \ + authorized (key revoked, access removed, or deployment unpublished)", + ) + .await; + send_critical( + &tx, + MessageRoute::CloseWith { + code: super::bud_legs::CLOSE_REVOKED, + reason: "session_revoked".to_string(), + }, + ) + .await; + return; + } + } + }); + tracker.track("bud-revalidation", handle); +} + +/// Bind a Bud-mode session's DAG template to the caller's deployments (FRD-023 WP-RT6.3). +#[cfg(feature = "dag-routing")] +async fn bind_bud_dag( + app_state: &Arc, + state: &Arc>, + message_tx: &mpsc::Sender, + definition: &mut DAGDefinition, + stream_id: &str, +) -> Result<(), String> { + let (credential, meter) = { + let guard = state.read().await; + (guard.credential.clone(), guard.leg_meter.clone()) + }; + let credential = credential.ok_or_else(|| { + "this gateway serves Bud deployments, which need your Bud API key or token".to_string() + })?; + let meter = match meter { + Some(meter) => meter, + None => { + // No audio legs: the session meter anchors the caller, the held deployments and + // their revalidation. + let bearer = crate::handlers::openai_realtime::handshake::Credential::new( + credential.current(), + crate::handlers::openai_realtime::handshake::CredentialSource::Bearer, + ); + let caller = + crate::handlers::openai_realtime::session::authenticate(app_state, &bearer) + .await + .map_err(|e| format!("{}: {}", e.code, e.message))?; + let meter = Arc::new(super::bud_legs::LegMeter::new( + stream_id.to_string(), + caller, + None, + None, + )); + state.write().await.leg_meter = Some(Arc::clone(&meter)); + spawn_bud_revalidation(app_state, state, message_tx, Arc::clone(&meter)).await; + meter + } + }; + let admissions = + super::bud_legs::bind_dag(app_state, definition, &credential, &meter, stream_id) + .await + .map_err(|r| format!("{}: {}", r.code, r.message))?; + state.write().await.leg_admissions.extend(admissions); + Ok(()) +} + +/// Point a Bud-mode voice agent's LLM leg at the Bud gateway (FRD-023 RT6, FR-WS-2, TC-WS-06). +/// +/// `model` (and `reasoning_model`) name Bud chat deployments; the address is the operator's +/// `WAAV_LLM_BASE_URL` and the credential is the caller's own, read per call so an `auth` refresh +/// reaches the next one. The client's endpoint and key were refused earlier (RT0). +fn bud_llm_leg( + config: &mut crate::core::conversation::ConversationConfig, + credential: Option, +) -> Result<(), String> { + let base_url = super::bud_legs::llm_base_url().ok_or_else(|| { + "the voice agent's LLM leg is not available on this gateway: WAAV_LLM_BASE_URL (the Bud \ + gateway) is not configured" + .to_string() + })?; + let credential = credential.ok_or_else(|| { + "the voice agent's LLM leg needs your Bud API key or token, and this session has none" + .to_string() + })?; + config.base_url = base_url; + config.server_llm_endpoint = true; + config.api_key = None; + config.credential = Some(credential); + // budgateway speaks the OpenAI wire format whatever the deployment's vendor. + config.provider_kind = Some(crate::core::llm::AdapterKind::OpenAi); + config.reasoning_base_url = None; + config.reasoning_api_key = None; + config.reasoning_provider_kind = Some(crate::core::llm::AdapterKind::OpenAi); + Ok(()) +} + /// Initialize the built-in conversation loop for a session (plan W-O2). /// /// Constructs a [`ConversationOrchestrator`] (validating the client-supplied LLM @@ -1056,13 +1311,18 @@ async fn initialize_conversation_loop( stream_id: &str, voice_manager: &Arc, message_tx: &mpsc::Sender, + app_state: &Arc, + credential: Option, + attribution: Option, ) -> Result { - let orchestrator = ConversationOrchestrator::new( - stream_id.to_string(), - conv_config.to_conversation_config(), - voice_manager.clone(), - ) - .map_err(|e| e.to_string())?; + let mut config = conv_config.to_conversation_config(); + if app_state.bud_mode.is_some() { + bud_llm_leg(&mut config, credential)?; + config.attribution = attribution; + } + let orchestrator = + ConversationOrchestrator::new(stream_id.to_string(), config, voice_manager.clone()) + .map_err(|e| e.to_string())?; let orchestrator = Arc::new(orchestrator); @@ -1091,11 +1351,16 @@ async fn initialize_conversation_loop( crate::core::observability::spawn_observed_detached( "conversation.fatal-handler", async move { - send_error( - &message_tx, - format!("fatal provider error (session cannot recover): {error}"), - ) - .await; + let message = if error.starts_with("auth_expired") { + // FRD-023 RT6: the session stays; the next turn uses a refreshed token. + format!( + "{error}. The Bud gateway refused this session's credential; send \ + {{\"type\":\"auth\",\"token\":\"\"}} to continue." + ) + } else { + format!("fatal provider error (session cannot recover): {error}") + }; + send_error(&message_tx, message).await; }, ); })); @@ -1374,6 +1639,14 @@ fn validate_audio_config_values( stt_config: &STTWebSocketConfig, tts_config: &TTSWebSocketConfig, ) -> Result<(), String> { + // Optional on the wire since FRD-023 RT6 (a Bud leg takes its vendor from the deployment, and + // by this point has it); a standalone gateway still needs to be told. + if stt_config.provider.trim().is_empty() { + return Err("STT provider is required when audio=true".to_string()); + } + if tts_config.provider.trim().is_empty() { + return Err("TTS provider is required when audio=true".to_string()); + } if stt_config.sample_rate == 0 { return Err("STT sample_rate must be greater than 0 when audio=true".to_string()); } @@ -1394,6 +1667,7 @@ async fn initialize_voice_manager( tts_ws_config: &TTSWebSocketConfig, app_state: &Arc, message_tx: &mpsc::Sender, + bud_legs: Option<&super::bud_legs::PreparedLegs>, ) -> Option> { info!( "Initializing voice manager with STT provider: {} and TTS provider: {}", @@ -1403,25 +1677,31 @@ async fn initialize_voice_manager( // Get API keys - prefer client-provided keys where the deployment allows them, fall back to // server config let allow_client_keys = app_state.allows_client_supplied_keys(); - let stt_api_key = resolve_provider_api_key( - stt_ws_config.api_key.as_deref(), - &stt_ws_config.provider, - "stt", - allow_client_keys, - &app_state.config, - message_tx, - ) - .await?; - - let tts_api_key = resolve_provider_api_key( - tts_ws_config.api_key.as_deref(), - &tts_ws_config.provider, - "tts", - allow_client_keys, - &app_state.config, - message_tx, - ) - .await?; + // FRD-023 RT6: a Bud deployment's legs bring their own credentials (from voice_table). + let (stt_api_key, tts_api_key) = match bud_legs { + Some(legs) => (legs.stt_key.clone(), legs.tts_key.clone()), + None => { + let stt_api_key = resolve_provider_api_key( + stt_ws_config.api_key.as_deref(), + &stt_ws_config.provider, + "stt", + allow_client_keys, + &app_state.config, + message_tx, + ) + .await?; + let tts_api_key = resolve_provider_api_key( + tts_ws_config.api_key.as_deref(), + &tts_ws_config.provider, + "tts", + allow_client_keys, + &app_state.config, + message_tx, + ) + .await?; + (stt_api_key, tts_api_key) + } + }; // P4 VOICE-DESCRIPTOR resolution: when the client supplied a canonical // `voice_descriptor` but NO raw `voice_id`, resolve it SERVER-SIDE to a concrete @@ -1500,7 +1780,11 @@ async fn initialize_voice_manager( // on the flat factory path. The flat `stt_config`/`tts_config` are still derived (== the // standardized bases) for cache hashing and other flat consumers below. let standard_stt = stt_ws_config.to_standard_stt(stt_api_key); - let standard_tts = tts_ws_config.to_standard_tts(tts_api_key); + let mut standard_tts = tts_ws_config.to_standard_tts(tts_api_key); + if let Some(legs) = bud_legs { + // The deployment's address, published by budapp — never the client's (§5.9). + standard_tts.base.api_base = legs.tts_api_base.clone(); + } // P2 language standardization: the client's canonical `language` was just mapped to each // provider's native notation inside `to_standard_stt`/`to_standard_tts` (so no provider sees a @@ -2787,9 +3071,10 @@ async fn initialize_dag_routing( profiler: Arc, egress_audio: Option>, resilience: Option>, + app_state: &Arc, ) -> Result { // Get DAG definition from template or inline - let dag_definition: DAGDefinition = if let Some(ref def) = dag_config.definition { + let mut dag_definition: DAGDefinition = if let Some(ref def) = dag_config.definition { // Parse inline definition serde_json::from_value(def.clone()).map_err(|e| format!("Invalid DAG definition: {}", e))? } else if let Some(ref template_name) = dag_config.template { @@ -2803,6 +3088,12 @@ async fn initialize_dag_routing( return Ok(false); }; + // FRD-023 WP-RT6.3: under the Bud control plane the template's provider and LLM nodes address + // Bud deployments as the caller (inline definitions were refused before this, RT0). + if app_state.bud_mode.is_some() { + bind_bud_dag(app_state, state, message_tx, &mut dag_definition, stream_id).await?; + } + info!( dag_id = %dag_definition.id, dag_name = %dag_definition.name, @@ -3874,4 +4165,204 @@ mod frd023_bud_mode_tests { "nothing was built" ); } + + // ----------------------------------------------------------------------------------------- + // FRD-023 RT6 + // ----------------------------------------------------------------------------------------- + + const KEY: &str = "bud_ws_config_test_key"; + + async fn leg_plane(stt_extra: serde_json::Value) -> Arc { + let mut stt = serde_json::json!({ + "vendor": "deepgram", "credential": crate::test_support::TEST_CREDENTIAL.trim(), + "endpoints": ["audio_transcription"], "model": "nova-3" + }); + for (k, v) in stt_extra.as_object().unwrap() { + stt[k] = v.clone(); + } + let tts = serde_json::json!({ + "vendor": "elevenlabs", "credential": crate::test_support::TEST_CREDENTIAL.trim(), + "endpoints": ["text_to_speech"], "model": "eleven_flash_v2_5", "voice": "v1" + }); + let blob = serde_json::json!({ + "stt-dg": {"endpoint_id": "ep-stt", "project_id": "p1", "kind": "model"}, + "tts-el": {"endpoint_id": "ep-tts", "project_id": "p1", "kind": "model"}, + "__metadata__": {"api_key_id": "k1", "user_id": "u1", "api_key_project_id": "p1"} + }) + .to_string(); + let keys = [ + (format!("api_key:{}", bud_auth::hash_api_key(KEY)), blob), + ( + "voice_table:ep-stt".to_string(), + serde_json::json!({"ep-stt": stt}).to_string(), + ), + ( + "voice_table:ep-tts".to_string(), + serde_json::json!({"ep-tts": tts}).to_string(), + ), + ]; + let refs: Vec<(&str, &str)> = keys.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect(); + crate::test_support::bud_state_with_credentials(&refs) + .await + .0 + } + + fn keyed_state() -> Arc> { + let mut state = ConnectionState::with_auth(crate::auth::Auth::new("p1")); + state.credential = Some(crate::auth::SessionCredential::new(KEY)); + Arc::new(RwLock::new(state)) + } + + fn leg(v: serde_json::Value) -> (STTWebSocketConfig, TTSWebSocketConfig) { + let mut stt = serde_json::json!({"language": "en-US", "sample_rate": 16000, "channels": 1, + "punctuation": true}); + for (k, val) in v["stt"].as_object().unwrap() { + stt[k] = val.clone(); + } + let mut tts = serde_json::json!({"voice_id": null, "speaking_rate": null, + "audio_format": "linear16", "sample_rate": 24000, "connection_timeout": null, + "request_timeout": null}); + for (k, val) in v["tts"].as_object().unwrap() { + tts[k] = val.clone(); + } + ( + serde_json::from_value(stt).unwrap(), + serde_json::from_value(tts).unwrap(), + ) + } + + fn drain(rx: &mut mpsc::Receiver) -> Vec { + let mut out = Vec::new(); + while let Ok(route) = rx.try_recv() { + out.push(route); + } + out + } + + /// TC-WS-04 🔒 — a provider-only leg is refused before anything is built; the socket stays + /// open for a corrected config. + #[tokio::test] + async fn tc_ws_04_a_provider_only_leg_is_refused_with_the_addressing_hint() { + let app_state = leg_plane(serde_json::json!({})).await; + let state = keyed_state(); + let (tx, mut rx) = mpsc::channel(16); + let (stt, tts) = + leg(serde_json::json!({"stt": {"provider": "deepgram"}, "tts": {"model": "tts-el"}})); + let keep = handle_config_message( + None, + Some(true), + Some(stt), + Some(tts), + None, + None, + None, + None, + &state, + &tx, + &app_state, + ) + .await; + assert!(keep); + let routes = drain(&mut rx); + assert!( + matches!(routes.first(), Some(MessageRoute::Outgoing(OutgoingMessage::Error { message })) + if message.starts_with("deployment_required")), + "{routes:?}" + ); + assert!( + !routes + .iter() + .any(|r| matches!(r, MessageRoute::CloseWith { .. })) + ); + let guard = state.read().await; + assert!( + guard.stream_id.is_none() && guard.voice_manager.is_none() && guard.leg_meter.is_none() + ); + } + + /// TC-WS-05 🔒 — a leg deployment at its cap refuses the session with close code 1013. + #[tokio::test] + async fn tc_ws_05_a_capped_leg_closes_the_session_with_1013() { + let app_state = leg_plane(serde_json::json!({"max_concurrent": 1})).await; + let _held = app_state.admit_deployment("ep-stt").await.expect("free"); + let state = keyed_state(); + let (tx, mut rx) = mpsc::channel(16); + let (stt, tts) = + leg(serde_json::json!({"stt": {"model": "stt-dg"}, "tts": {"model": "tts-el"}})); + let keep = handle_config_message( + None, + Some(true), + Some(stt), + Some(tts), + None, + None, + None, + None, + &state, + &tx, + &app_state, + ) + .await; + assert!(!keep, "the session ends"); + let routes = drain(&mut rx); + assert!( + routes + .iter() + .any(|r| matches!(r, MessageRoute::CloseWith { code: 1013, .. })), + "{routes:?}" + ); + assert!( + app_state.admit_deployment("ep-tts").await.is_ok(), + "nothing held" + ); + } + + /// TC-WS-06 🔒 — the LLM leg is the Bud gateway with the caller's credential, OpenAI wire. + #[test] + #[serial_test::serial] + fn tc_ws_06_the_llm_leg_targets_the_bud_gateway_as_the_caller() { + let mut config = crate::core::conversation::ConversationConfig { + base_url: "https://api.openai.com/v1".into(), + api_key: Some("sk-client".into()), + provider_kind: Some(crate::core::llm::AdapterKind::Anthropic), + ..Default::default() + }; + let credential = crate::auth::SessionCredential::new(KEY); + + unsafe { std::env::remove_var("WAAV_LLM_BASE_URL") }; + let err = bud_llm_leg(&mut config.clone(), Some(credential.clone())).unwrap_err(); + assert!(err.contains("WAAV_LLM_BASE_URL"), "{err}"); + + unsafe { std::env::set_var("WAAV_LLM_BASE_URL", "http://ditto-budgateway:3000/v1") }; + let err = bud_llm_leg(&mut config.clone(), None).unwrap_err(); + assert!(err.contains("credential") || err.contains("key"), "{err}"); + + bud_llm_leg(&mut config, Some(credential)).expect("configured"); + unsafe { std::env::remove_var("WAAV_LLM_BASE_URL") }; + assert_eq!(config.base_url, "http://ditto-budgateway:3000/v1"); + assert!(config.server_llm_endpoint); + assert!(config.api_key.is_none()); + assert_eq!( + config.credential.as_ref().map(|c| c.current()).as_deref(), + Some(KEY) + ); + assert_eq!( + config.provider_kind, + Some(crate::core::llm::AdapterKind::OpenAi) + ); + assert_eq!( + config.reasoning_provider_kind, + Some(crate::core::llm::AdapterKind::OpenAi) + ); + } + + /// `provider` became optional on the wire for Bud clients; a standalone gateway still refuses + /// a leg without one, by name. + #[test] + fn a_standalone_leg_still_needs_a_provider() { + let (stt, tts) = + leg(serde_json::json!({"stt": {"model": "nova-3"}, "tts": {"provider": "deepgram"}})); + let err = validate_audio_config_values(&stt, &tts).unwrap_err(); + assert!(err.contains("STT provider"), "{err}"); + } } diff --git a/gateway/src/handlers/ws/handler.rs b/gateway/src/handlers/ws/handler.rs index 504f82fd..2e942212 100644 --- a/gateway/src/handlers/ws/handler.rs +++ b/gateway/src/handlers/ws/handler.rs @@ -71,6 +71,7 @@ pub async fn ws_voice_handler( Extension(auth): Extension, client_ip: Option>, slot: Option>, + credential: Option>, ) -> Response { info!( auth_id = ?auth.id, @@ -83,6 +84,7 @@ pub async fn ws_voice_handler( // The connection slot rides into the session and is released when it ends; if the upgrade // never happens, the closure (and the slot) is dropped with it. let slot = slot.map(|Extension(s)| s); + let credential = credential.map(|Extension(c)| c); // Apply message size limits to prevent memory exhaustion attacks let response = ws @@ -90,7 +92,7 @@ pub async fn ws_voice_handler( .max_message_size(MAX_WS_MESSAGE_SIZE) .on_upgrade(move |socket| { debug!("WebSocket upgrade callback triggered"); - handle_voice_socket(socket, state, auth, ip, slot) + handle_voice_socket(socket, state, auth, ip, slot, credential) }); debug!("WebSocket upgrade response created"); @@ -125,6 +127,7 @@ async fn handle_voice_socket( auth: Auth, client_ip: Option, slot: Option, + credential: Option, ) { // Multi-tenant panic isolation (W-E1 / E6). // @@ -140,8 +143,9 @@ async fn handle_voice_socket( // exposing a logically-torn invariant to another session. let _connection_slot = slot; - let session = - std::panic::AssertUnwindSafe(run_voice_socket_session(socket, app_state, auth, client_ip)); + let session = std::panic::AssertUnwindSafe(run_voice_socket_session( + socket, app_state, auth, client_ip, credential, + )); if futures::FutureExt::catch_unwind(session).await.is_err() { // A panic was caught and contained to this session. The process and all // other sessions remain alive. The connection guard above still releases @@ -160,6 +164,7 @@ async fn run_voice_socket_session( app_state: Arc, auth: Auth, client_ip: Option, + credential: Option, ) { debug!("handle_voice_socket started"); info!( @@ -177,6 +182,21 @@ async fn run_voice_socket_session( // Connection state with RwLock for rare writes, frequent reads // Initialize with auth context for room name normalization let state = Arc::new(RwLock::new(ConnectionState::with_auth(auth.clone()))); + // FRD-023 RT6: fix the identity the upgrade authenticated, for `auth` refreshes to match. + if app_state.bud_mode.is_some() + && let Some(credential) = credential.as_ref() + { + let bearer = crate::handlers::openai_realtime::handshake::Credential::new( + credential.current(), + crate::handlers::openai_realtime::handshake::CredentialSource::Bearer, + ); + state.write().await.caller_check = + crate::handlers::openai_realtime::session::authenticate(&app_state, &bearer) + .await + .ok() + .map(|caller| caller.check); + } + state.write().await.credential = credential; let (message_tx, mut message_rx) = mpsc::channel::(CHANNEL_BUFFER_SIZE); @@ -204,7 +224,8 @@ async fn run_voice_socket_session( // Channel closed, exit gracefully break; }; - let should_close = matches!(route, MessageRoute::Close); + let should_close = + matches!(route, MessageRoute::Close | MessageRoute::CloseWith { .. }); let result = match route { MessageRoute::Outgoing(message) => { @@ -222,6 +243,10 @@ async fn run_voice_socket_session( info!("Closing WebSocket connection"); sender.send(Message::Close(None)).await } + MessageRoute::CloseWith { code, reason } => { + info!(code, reason = %reason, "Closing WebSocket connection"); + sender.send(coded_close(code, &reason)).await + } }; if let Err(e) = result { @@ -246,6 +271,9 @@ async fn run_voice_socket_session( } MessageRoute::Binary(data) => sender.send(Message::Binary(data)).await, MessageRoute::Close => sender.send(Message::Close(None)).await, + MessageRoute::CloseWith { code, reason } => { + sender.send(coded_close(code, &reason)).await + } }; if result.is_err() { break; @@ -307,6 +335,12 @@ async fn run_voice_socket_session( // Signal shutdown to sender task shutdown_voice_sender_task(shutdown_tx, &mut sender_task).await; + // FRD-023 RT6: bill audio streamed after the last final transcript; the legs' admissions are + // released with the state. + if let Some(meter) = state.read().await.leg_meter.clone() { + meter.finish(); + } + // Snapshot state before cleanup so we can drop the read lock before awaiting let (voice_manager, livekit_client, recording_egress_id, room_name) = { let state_guard = state.read().await; @@ -552,6 +586,18 @@ where } } +/// A close frame with a code, its reason cut to the 123 bytes a control frame allows. +fn coded_close(code: u16, reason: &str) -> Message { + let mut end = reason.len().min(123); + while !reason.is_char_boundary(end) { + end -= 1; + } + Message::Close(Some(axum::extract::ws::CloseFrame { + code, + reason: reason[..end].to_string().into(), + })) +} + async fn shutdown_voice_sender_task( shutdown_tx: tokio::sync::oneshot::Sender<()>, sender_task: &mut tokio::task::JoinHandle<()>, diff --git a/gateway/src/handlers/ws/messages.rs b/gateway/src/handlers/ws/messages.rs index ea42a022..ae0616d7 100644 --- a/gateway/src/handlers/ws/messages.rs +++ b/gateway/src/handlers/ws/messages.rs @@ -421,6 +421,12 @@ pub enum MessageRoute { Outgoing(OutgoingMessage), Binary(Bytes), Close, + /// Close with a code and reason (FRD-023 RT6: 1013 when a leg's deployment is at capacity, + /// 1008 when the session's caller lost access). + CloseWith { + code: u16, + reason: String, + }, } /// Prometheus counter: outbound WS messages dropped by the per-class send diff --git a/gateway/src/handlers/ws/mod.rs b/gateway/src/handlers/ws/mod.rs index 8e5c6b93..cdf23acd 100644 --- a/gateway/src/handlers/ws/mod.rs +++ b/gateway/src/handlers/ws/mod.rs @@ -417,6 +417,7 @@ //! All errors are sent back to the client as JSON messages with `type: "error"`. pub mod audio_handler; +pub mod bud_legs; pub mod command_handler; pub mod config; pub mod config_handler; diff --git a/gateway/src/handlers/ws/processor.rs b/gateway/src/handlers/ws/processor.rs index 6a31ed08..31410329 100644 --- a/gateway/src/handlers/ws/processor.rs +++ b/gateway/src/handlers/ws/processor.rs @@ -226,12 +226,31 @@ async fn handle_auth_message( app_state.config.has_jwt_auth(), ); + if path == WsAuthPath::Bud && !state.read().await.auth.is_pending() { + return refresh_bud_credential(token, state, message_tx, app_state).await; + } + let resolved: Option = match path { WsAuthPath::Bud => { // Unwrap is safe: `ws_auth_path` returns Bud only when bud_mode is Some. let bud = app_state.bud_mode.as_ref().expect("bud mode present"); match bud.authenticate(&token).await { - Ok(auth) => auth.id, + Ok(auth) => { + // FRD-023 RT6: the session acts as this caller from here on. + let bearer = crate::handlers::openai_realtime::handshake::Credential::new( + token.clone(), + crate::handlers::openai_realtime::handshake::CredentialSource::Bearer, + ); + let check = + crate::handlers::openai_realtime::session::authenticate(app_state, &bearer) + .await + .ok() + .map(|caller| caller.check); + let mut guard = state.write().await; + guard.credential = Some(crate::auth::SessionCredential::new(token.clone())); + guard.caller_check = check; + auth.id + } Err(e) => { warn!(error = ?e, "first-message bud authentication failed"); None @@ -302,6 +321,65 @@ async fn handle_auth_message( } } +/// An `auth` message on an authenticated Bud-mode session: a credential REFRESH (FRD-023 WP-RT6.2, +/// TC-WS-08). A Keycloak token lives minutes and a call longer, so a JWT caller keeps its voice +/// agent's LLM leg alive by sending a fresh token; the leg reads it on its next call. +/// +/// The new credential must identify the SAME principal — the session's legs are admitted, billed +/// and authorized against it. A refresh that fails leaves the session and its current credential +/// as they were: the credential failing is what `auth_expired` reports, and revalidation is what +/// ends a session whose caller lost access. +async fn refresh_bud_credential( + token: String, + state: &Arc>, + message_tx: &mpsc::Sender, + app_state: &Arc, +) -> bool { + let bearer = crate::handlers::openai_realtime::handshake::Credential::new( + token.clone(), + crate::handlers::openai_realtime::handshake::CredentialSource::Bearer, + ); + let caller = + match crate::handlers::openai_realtime::session::authenticate(app_state, &bearer).await { + Ok(caller) => caller, + Err(e) => { + warn!(code = e.code, "Bud-mode /ws credential refresh refused"); + send_error( + message_tx, + format!( + "auth_refresh_failed: {}. The session keeps its current credential.", + e.message + ), + ) + .await; + return true; + } + }; + let guard = state.read().await; + let same = guard.caller_check.as_ref() == Some(&caller.check); + let Some(credential) = guard.credential.clone().filter(|_| same) else { + drop(guard); + warn!("Bud-mode /ws credential refresh carried a different identity; refused"); + send_error( + message_tx, + "auth_refresh_refused: a refresh must renew this session's own credential (the same \ + API key or the same user); open a new connection to act as someone else.", + ) + .await; + return true; + }; + credential.replace(token); + let id = guard.auth.id.clone(); + drop(guard); + info!("Bud-mode /ws credential refreshed"); + send_critical( + message_tx, + MessageRoute::Outgoing(OutgoingMessage::Authenticated { id }), + ) + .await; + true +} + /// Handle custom plugin message /// /// Dispatches the message to all registered handlers for this message type. @@ -416,7 +494,7 @@ async fn handle_custom_message( #[cfg(test)] mod ws_auth_path_tests { - use super::{WsAuthPath, ws_auth_path}; + use super::*; #[test] fn bud_mode_is_consulted_whenever_it_is_configured() { @@ -464,4 +542,124 @@ mod ws_auth_path_tests { // "everything passes". assert_eq!(ws_auth_path(false, false, false), WsAuthPath::Unconfigured); } + + // ----------------------------------------------------------------------------------------- + // FRD-023 WP-RT6.2: `auth` on an authenticated Bud-mode session is a credential refresh. + // ----------------------------------------------------------------------------------------- + + const KEY: &str = "bud_ws_refresh_key"; + const OTHER_KEY: &str = "bud_ws_refresh_other_key"; + + async fn refresh_plane() -> Arc { + let blob = |project: &str| { + serde_json::json!({"__metadata__": {"api_key_id": "k", "user_id": "u", + "api_key_project_id": project}}) + .to_string() + }; + let a = ( + format!("api_key:{}", bud_auth::hash_api_key(KEY)), + blob("p1"), + ); + let b = ( + format!("api_key:{}", bud_auth::hash_api_key(OTHER_KEY)), + blob("p2"), + ); + crate::test_support::bud_state_with_credentials(&[ + (a.0.as_str(), a.1.as_str()), + (b.0.as_str(), b.1.as_str()), + ]) + .await + .0 + } + + /// A session authenticated as KEY, as the upgrade leaves it. + async fn authenticated(app_state: &Arc) -> Arc> { + let state = Arc::new(RwLock::new(ConnectionState::with_auth(Auth::new("p1")))); + let bearer = crate::handlers::openai_realtime::handshake::Credential::new( + KEY, + crate::handlers::openai_realtime::handshake::CredentialSource::Bearer, + ); + let check = crate::handlers::openai_realtime::session::authenticate(app_state, &bearer) + .await + .unwrap() + .check; + let mut guard = state.write().await; + guard.credential = Some(crate::auth::SessionCredential::new(KEY)); + guard.caller_check = Some(check); + drop(guard); + state + } + + fn first_error(rx: &mut mpsc::Receiver) -> Option { + while let Ok(route) = rx.try_recv() { + if let MessageRoute::Outgoing(OutgoingMessage::Error { message }) = route { + return Some(message); + } + } + None + } + + /// TC-WS-08 🔒 — a refresh with the same identity replaces the session's credential. + #[tokio::test] + async fn tc_ws_08_a_refresh_renews_the_credential() { + let app_state = refresh_plane().await; + let state = authenticated(&app_state).await; + let credential = state.read().await.credential.clone().unwrap(); + let (tx, mut rx) = mpsc::channel(8); + + // The same key, re-sent (a JWT caller sends a fresh token for the same subject). + assert!(handle_auth_message(KEY.to_string(), &state, &tx, &app_state).await); + assert!(matches!( + rx.try_recv(), + Ok(MessageRoute::Outgoing( + OutgoingMessage::Authenticated { .. } + )) + )); + assert_eq!(credential.current(), KEY); + } + + /// TC-WS-08 🔒 — a refresh can never change who the session is. + #[tokio::test] + async fn tc_ws_08_a_refresh_to_another_identity_is_refused() { + let app_state = refresh_plane().await; + let state = authenticated(&app_state).await; + let credential = state.read().await.credential.clone().unwrap(); + let (tx, mut rx) = mpsc::channel(8); + + let keep = handle_auth_message(OTHER_KEY.to_string(), &state, &tx, &app_state).await; + assert!(keep, "the session continues on its own credential"); + let error = first_error(&mut rx).expect("an error frame"); + assert!(error.starts_with("auth_refresh_refused"), "{error}"); + assert_eq!(credential.current(), KEY, "unchanged"); + } + + /// A refresh that fails leaves the session and its credential as they were. + #[tokio::test] + async fn a_failed_refresh_keeps_the_session() { + let app_state = refresh_plane().await; + let state = authenticated(&app_state).await; + let credential = state.read().await.credential.clone().unwrap(); + let (tx, mut rx) = mpsc::channel(8); + + let keep = handle_auth_message("bud_nope".to_string(), &state, &tx, &app_state).await; + assert!(keep); + let error = first_error(&mut rx).expect("an error frame"); + assert!(error.starts_with("auth_refresh_failed"), "{error}"); + assert_eq!(credential.current(), KEY); + } + + /// First-message auth in Bud mode keeps the credential and fixes the identity. + #[tokio::test] + async fn first_message_auth_keeps_the_credential_for_the_session() { + let app_state = refresh_plane().await; + let state = Arc::new(RwLock::new(ConnectionState::with_auth(Auth::pending()))); + let (tx, _rx) = mpsc::channel(8); + assert!(handle_auth_message(KEY.to_string(), &state, &tx, &app_state).await); + let guard = state.read().await; + assert_eq!( + guard.credential.as_ref().map(|c| c.current()).as_deref(), + Some(KEY) + ); + assert!(guard.caller_check.is_some()); + } } diff --git a/gateway/src/handlers/ws/state.rs b/gateway/src/handlers/ws/state.rs index 79057009..7e76c23d 100644 --- a/gateway/src/handlers/ws/state.rs +++ b/gateway/src/handlers/ws/state.rs @@ -47,6 +47,15 @@ pub struct ConnectionState { pub recording_egress_id: Option, /// Auth context for this connection (used for room name normalization) pub auth: Auth, + /// FRD-023 RT6: the caller's credential (Bud mode), refreshed by `auth` messages. + pub credential: Option, + /// FRD-023 RT6: who the session's credential identifies, fixed when the session authenticates + /// so an `auth` refresh can renew the credential but never change the principal. + pub caller_check: Option, + /// FRD-023 RT6: per-leg metering for a session whose legs address Bud deployments. + pub leg_meter: Option>, + /// FRD-023 RT6: each leg deployment's admission, held for the session (FRD-022 §6.2). + pub leg_admissions: Vec, /// D8 uplink opus decoder (feature `opus-codec`): `Some` only when the session negotiated /// `stt_config.audio_in_codec = opus`. Each client WS binary frame is one opus packet decoded @@ -90,6 +99,10 @@ impl ConnectionState { livekit_local_identity: None, recording_egress_id: None, auth: Auth::empty(), + credential: None, + caller_check: None, + leg_meter: None, + leg_admissions: Vec::new(), #[cfg(feature = "opus-codec")] opus_decoder: None, #[cfg(feature = "dag-routing")] @@ -117,6 +130,10 @@ impl ConnectionState { livekit_local_identity: None, recording_egress_id: None, auth, + credential: None, + caller_check: None, + leg_meter: None, + leg_admissions: Vec::new(), #[cfg(feature = "opus-codec")] opus_decoder: None, #[cfg(feature = "dag-routing")] diff --git a/gateway/src/middleware/auth.rs b/gateway/src/middleware/auth.rs index a6ecd103..8e6f0e0b 100644 --- a/gateway/src/middleware/auth.rs +++ b/gateway/src/middleware/auth.rs @@ -162,6 +162,14 @@ async fn authenticate_request( auth_id = ?auth.id, "bud authentication successful" ); + // FRD-023 RT6: a /ws session acts as its caller for the session's life (its legs' + // deployments, the LLM leg through budgateway, revalidation), so it keeps the + // credential — redacted, in memory only. + if request_path == "/ws" { + request + .extensions_mut() + .insert(crate::auth::SessionCredential::new(token.clone())); + } request.extensions_mut().insert(auth); Ok(next.run(request).await) } diff --git a/gateway/src/test_support.rs b/gateway/src/test_support.rs index ed14c594..ed10daa1 100644 --- a/gateway/src/test_support.rs +++ b/gateway/src/test_support.rs @@ -92,3 +92,39 @@ pub(crate) async fn bud_state(keys: &[(&str, &str)]) -> Arc { .bud_mode = Some(crate::auth::bud_mode::BudMode::for_plane(plane).expect("bud mode")); state } + +/// [`bud_state`] whose plane opens `voice_table` credentials with bud-auth's fixture key (the +/// plaintext of `test_cred_encrypted.hex` is `dg_vendor_key_abc123`) and enforces deployment +/// policies (FRD-023 RT6 tests). Returns the store, to mutate the control plane mid-test. +pub(crate) async fn bud_state_with_credentials( + keys: &[(&str, &str)], +) -> (Arc, Arc) { + let store = Arc::new(bud_auth::MemoryStore::new()); + for (k, v) in keys { + store.set(k, v); + } + let pem = std::fs::read_to_string(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../bud-auth/tests/fixtures/test_cred_private.pem" + )) + .expect("bud-auth's fixture key (git-ignored *.pem) must be present locally"); + let plane = Arc::new(bud_auth::BudPlane::with_decryptor( + store.clone() as Arc, + None, + bud_auth::CredentialDecryptor::from_pem(&pem).expect("fixture key parses"), + )); + plane.boot().await.expect("plane boots"); + let mut state = AppState::new(minimal_config()).await; + { + let s = Arc::get_mut(&mut state).expect("the state is not shared yet"); + s.bud_mode = Some(crate::auth::bud_mode::BudMode::for_plane(plane).expect("bud mode")); + s.policies = Some(crate::core::deployment_policy::DeploymentPolicies::local()); + } + (state, store) +} + +/// bud-auth's fixture ciphertext, for `voice_table` entries in tests. +pub(crate) const TEST_CREDENTIAL: &str = + include_str!("../../bud-auth/tests/fixtures/test_cred_encrypted.hex"); +/// Its plaintext. +pub(crate) const TEST_CREDENTIAL_PLAIN: &str = "dg_vendor_key_abc123"; diff --git a/gateway/tests/conversation_loop.rs b/gateway/tests/conversation_loop.rs index b1a0b2c7..869c2722 100644 --- a/gateway/tests/conversation_loop.rs +++ b/gateway/tests/conversation_loop.rs @@ -218,6 +218,10 @@ struct LlmMockState { reply: Arc>, /// P1: when set, the endpoint returns HTTP 500 (a failing LLM tier). fail: Arc, + /// FRD-023 RT6: every request's `Authorization` header, in order. + authorizations: Arc>>, + /// FRD-023 RT6: when set, the endpoint refuses the credential with HTTP 401. + unauthorized: Arc, } struct TestServer { @@ -264,9 +268,23 @@ where async fn start_llm_mock(state: LlmMockState) -> (String, TestServer) { async fn chat( State(state): State, + headers: axum::http::HeaderMap, Json(req): Json, ) -> (axum::http::StatusCode, Json) { state.requests.lock().push(req.clone()); + state.authorizations.lock().push( + headers + .get("authorization") + .and_then(|v| v.to_str().ok()) + .unwrap_or_default() + .to_string(), + ); + if state.unauthorized.load(Ordering::SeqCst) { + return ( + axum::http::StatusCode::UNAUTHORIZED, + Json(json!({ "error": { "message": "token expired", "code": "invalid_api_key" } })), + ); + } let delay = state.delay_ms.load(Ordering::SeqCst); if delay > 0 { tokio::time::sleep(Duration::from_millis(delay as u64)).await; @@ -2566,3 +2584,311 @@ async fn a_conversation_turn_emits_the_llm_leg_on_its_voice_turn_span() { "turn index missing: two turns in one session would be indistinguishable" ); } + +// ────────────────────────────────────────────────────────────────────────────── +// FRD-023 RT6: the voice agent's LLM leg on a Bud deployment. +// ────────────────────────────────────────────────────────────────────────────── + +/// The config `/ws` builds for a Bud-mode session: the operator's gateway address (not +/// SSRF-checked), no configured key, the caller's live credential. +fn bud_conv_config( + base_url: String, + credential: waav_gateway::auth::SessionCredential, +) -> ConversationConfig { + ConversationConfig { + base_url, + model: "chat-deployment".to_string(), + api_key: None, + server_llm_endpoint: true, + credential: Some(credential), + provider_kind: Some(waav_gateway::core::llm::AdapterKind::OpenAi), + streaming: false, + allow_interruption: true, + ..Default::default() + } +} + +/// TC-WS-06 🔒 / TC-WS-08 🔒 — every LLM call carries the caller's credential, and an `auth` +/// refresh reaches the next one. +#[tokio::test] +#[serial_test::serial] +async fn tc_ws_06_08_the_llm_leg_sends_the_live_session_credential() { + register_mock_tts(); + reset_tts_stats(); + let llm_state = LlmMockState::default(); + *llm_state.reply.lock() = "Sure.".to_string(); + let (base_url, _server) = start_llm_mock(llm_state.clone()).await; + let vm = build_voice_manager(); + vm.start().await.expect("vm start"); + + let credential = waav_gateway::auth::SessionCredential::new("bud_first_token"); + let orchestrator = ConversationOrchestrator::new( + "session-bud-llm", + bud_conv_config(base_url, credential.clone()), + vm.clone(), + ) + .expect("a server endpoint is not SSRF-checked (it is in-cluster by design)"); + + orchestrator.on_stt_result(&final_result("first")).await; + credential.replace("bud_refreshed_token"); + orchestrator.on_stt_result(&final_result("second")).await; + + assert_eq!( + *llm_state.authorizations.lock(), + vec![ + "Bearer bud_first_token".to_string(), + "Bearer bud_refreshed_token".to_string() + ] + ); +} + +/// TC-WS-09 🔒 — an expired credential fails the call with `auth_expired`, not the session: once +/// refreshed, the next turn is answered. +#[tokio::test] +#[serial_test::serial] +async fn tc_ws_09_an_expired_credential_is_auth_expired_and_the_session_continues() { + register_mock_tts(); + reset_tts_stats(); + let llm_state = LlmMockState::default(); + *llm_state.reply.lock() = "Back again.".to_string(); + llm_state.unauthorized.store(true, Ordering::SeqCst); + let (base_url, _server) = start_llm_mock(llm_state.clone()).await; + let vm = build_voice_manager(); + vm.start().await.expect("vm start"); + + let credential = waav_gateway::auth::SessionCredential::new("bud_expired_token"); + let orchestrator = ConversationOrchestrator::new( + "session-bud-expired", + bud_conv_config(base_url, credential.clone()), + vm.clone(), + ) + .expect("orchestrator"); + let fatals: Arc>> = Arc::default(); + let seen = fatals.clone(); + orchestrator.set_fatal_handler(Arc::new(move |e: String| seen.lock().push(e))); + + orchestrator.on_stt_result(&final_result("hello?")).await; + { + let fatals = fatals.lock(); + assert_eq!(fatals.len(), 1, "{fatals:?}"); + assert!(fatals[0].starts_with("auth_expired"), "{}", fatals[0]); + } + + credential.replace("bud_fresh_token"); + llm_state.unauthorized.store(false, Ordering::SeqCst); + orchestrator + .on_stt_result(&final_result("hello again")) + .await; + assert_eq!( + llm_state.requests.lock().len(), + 2, + "the session kept listening" + ); + assert!( + TTS_STATS + .spoken + .lock() + .iter() + .any(|s| s.contains("Back again")), + "the refreshed turn was answered" + ); + assert_eq!(fatals.lock().len(), 1); +} + +/// TC-WS-09 — without a refreshable credential, an auth failure still stops the session (the +/// standalone gateway's behaviour is unchanged). +#[tokio::test] +#[serial_test::serial] +async fn a_configured_key_refused_still_stops_the_session() { + unsafe { + std::env::set_var("WAAV_ALLOW_LOOPBACK_ENDPOINTS", "1"); + } + register_mock_tts(); + reset_tts_stats(); + let llm_state = LlmMockState::default(); + llm_state.unauthorized.store(true, Ordering::SeqCst); + let (base_url, _server) = start_llm_mock(llm_state.clone()).await; + let vm = build_voice_manager(); + vm.start().await.expect("vm start"); + let orchestrator = + ConversationOrchestrator::new("session-byok", conv_config(base_url, false), vm.clone()) + .expect("orchestrator"); + let fatals: Arc>> = Arc::default(); + let seen = fatals.clone(); + orchestrator.set_fatal_handler(Arc::new(move |e: String| seen.lock().push(e))); + + orchestrator.on_stt_result(&final_result("one")).await; + orchestrator.on_stt_result(&final_result("two")).await; + assert_eq!( + llm_state.requests.lock().len(), + 1, + "stopped after the fatal" + ); + assert!(!fatals.lock()[0].starts_with("auth_expired")); +} + +/// A mock STT whose result callback the test drives (FRD-023 RT6 STT-leg metering). +static STT_CALLBACK: once_cell::sync::Lazy< + Mutex>, +> = once_cell::sync::Lazy::new(|| Mutex::new(None)); + +struct CallbackStt { + ready: bool, +} + +#[async_trait::async_trait] +impl waav_gateway::core::stt::BaseSTT for CallbackStt { + fn new(_config: STTConfig) -> Result { + Ok(Self { ready: false }) + } + async fn connect(&mut self) -> Result<(), waav_gateway::core::stt::STTError> { + self.ready = true; + Ok(()) + } + async fn disconnect(&mut self) -> Result<(), waav_gateway::core::stt::STTError> { + self.ready = false; + Ok(()) + } + fn is_ready(&self) -> bool { + self.ready + } + async fn send_audio( + &mut self, + _audio: bytes::Bytes, + ) -> Result<(), waav_gateway::core::stt::STTError> { + Ok(()) + } + async fn on_result( + &mut self, + cb: waav_gateway::core::stt::STTResultCallback, + ) -> Result<(), waav_gateway::core::stt::STTError> { + *STT_CALLBACK.lock() = Some(cb); + Ok(()) + } + async fn on_error( + &mut self, + _cb: waav_gateway::core::stt::STTErrorCallback, + ) -> Result<(), waav_gateway::core::stt::STTError> { + Ok(()) + } + fn get_config(&self) -> Option<&STTConfig> { + None + } + async fn update_config( + &mut self, + _config: STTConfig, + ) -> Result<(), waav_gateway::core::stt::STTError> { + Ok(()) + } + fn get_provider_info(&self) -> &'static str { + "mock-stt-cb" + } +} + +/// TC-WS-12 🔒 — the STT-final observer sees every vendor final, and survives the voice agent +/// replacing the result callback; interim results are not billed records. +#[tokio::test] +#[serial_test::serial] +async fn tc_ws_12_the_stt_final_observer_sees_every_final() { + register_mock_tts(); + global_registry().register_stt( + "mock-stt-cb", + Arc::new(|config: STTConfig| { + CallbackStt::new(config) + .map(|s| Box::new(s) as Box) + }), + ProviderMetadata::stt("mock-stt-cb", "Callback STT (RT6)"), + ); + let stt_config = STTConfig { + provider: "mock-stt-cb".to_string(), + api_key: "test".to_string(), + ..Default::default() + }; + let tts_config = TTSConfig { + provider: "mock-tts-conv".to_string(), + api_key: "test".to_string(), + ..Default::default() + }; + let vm = Arc::new( + VoiceManager::new(VoiceManagerConfig::new(stt_config, tts_config), None) + .expect("voice manager"), + ); + vm.start().await.expect("vm start"); + + let finals: Arc>> = Arc::default(); + let seen = finals.clone(); + vm.set_stt_final_observer(Arc::new(move |t: &str| seen.lock().push(t.to_string()))); + // The session's own forwarder, then the voice agent's replacement. + vm.on_stt_result(|_r| Box::pin(async {})).await.unwrap(); + vm.on_stt_result(|_r| Box::pin(async {})).await.unwrap(); + + let cb = STT_CALLBACK + .lock() + .clone() + .expect("the STT callback was registered"); + cb(STTResult::new("hel".to_string(), false, false, 0.5)).await; + cb(STTResult::new("hello there".to_string(), true, false, 0.9)).await; + cb(STTResult::new("and more".to_string(), true, true, 0.9)).await; + + assert_eq!( + *finals.lock(), + vec!["hello there".to_string(), "and more".to_string()] + ); +} + +/// FRD-023 RT6 (§5.9) — a Bud-mode agent's turn is attributed to its caller and the chat +/// deployment, and carries no cost of its own (the legs bill themselves). +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial_test::serial] +async fn a_bud_conversation_turn_is_attributed_to_its_caller() { + register_mock_tts(); + reset_tts_stats(); + let llm_state = LlmMockState::default(); + *llm_state.reply.lock() = "Paris.".to_string(); + let (base_url, _server) = start_llm_mock(llm_state.clone()).await; + let vm = build_voice_manager(); + vm.start().await.expect("vm start"); + + let mut config = bud_conv_config( + base_url, + waav_gateway::auth::SessionCredential::new("bud_k"), + ); + config.attribution = Some(waav_gateway::core::conversation::TurnAttribution { + project_id: Some("proj-1".into()), + api_key_id: Some("key-1".into()), + api_key_project_id: Some("proj-1".into()), + user_id: Some("user-1".into()), + endpoint_name: Some("chat-deployment".into()), + }); + let orchestrator = + ConversationOrchestrator::new("session-attr", config, vm.clone()).expect("orchestrator"); + + let fields: span_capture::Fields = Default::default(); + use tracing_subscriber::layer::SubscriberExt; + let dispatch = tracing::Dispatch::new( + tracing_subscriber::registry().with(span_capture::CaptureLayer(fields.clone())), + ); + use tracing::instrument::WithSubscriber; + orchestrator + .run_turn("capital of France?") + .with_subscriber(dispatch) + .await + .expect("turn"); + + let f = fields.lock().unwrap(); + assert_eq!( + f.get("bud.project_id").map(String::as_str), + Some("proj-1"), + "{f:?}" + ); + assert_eq!(f.get("bud.api_key_id").map(String::as_str), Some("key-1")); + assert_eq!(f.get("bud.user_id").map(String::as_str), Some("user-1")); + assert_eq!( + f.get("bud.voice.endpoint_name").map(String::as_str), + Some("chat-deployment") + ); + assert!( + !f.contains_key("bud.voice.cost"), + "the legs bill themselves: {f:?}" + ); +} From 0e632ca072da671d1f4516bfd6223a92c7f84788 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 05:29:25 +0530 Subject: [PATCH 07/17] fix(gateway): realtime model ids verbatim, merged session updates, current defaults (FRD-023 F-1, F-3, F-4) F-1: OpenAIRealtimeModel was a closed enum whose parser turned every id it did not list into `gpt-realtime`, so gpt-realtime-1.5, -2.1 and -2.1-mini silently became a model that shuts down on 2027-01-20. It is now a string carried verbatim; only an empty id takes the default (TC-XL-08). F-3: the native `/realtime` session update dropped tools and turn detection, and the scaffold's update_session replaced the whole config, so the key and the server-set endpoint override were lost and the next reconnect dialled without a key. The update now carries both and merges. F-4: Gemini Live defaulted to gemini-2.0-flash-live-001 (shut down 2025-12-09) and Nova Sonic to v1 (EOL 2026-09-14): now gemini-3.8-live and amazon.nova-2-sonic-v1:0. Co-Authored-By: Claude Opus 5.5 --- gateway/src/core/realtime/gemini/protocol.rs | 16 +- .../src/core/realtime/nova_sonic/protocol.rs | 11 +- gateway/src/core/realtime/openai/config.rs | 100 ++++----- gateway/src/core/realtime/openai/mod.rs | 8 +- gateway/src/core/realtime/openai/protocol.rs | 31 ++- gateway/src/core/realtime/scaffold/session.rs | 91 +++++++- gateway/src/handlers/realtime/handler.rs | 207 ++++++++++++++---- gateway/src/plugin/builtin/mod.rs | 4 +- gateway/tests/realtime_provider_matrix.rs | 7 +- 9 files changed, 367 insertions(+), 108 deletions(-) diff --git a/gateway/src/core/realtime/gemini/protocol.rs b/gateway/src/core/realtime/gemini/protocol.rs index 399f8fbc..c56e85e7 100644 --- a/gateway/src/core/realtime/gemini/protocol.rs +++ b/gateway/src/core/realtime/gemini/protocol.rs @@ -49,10 +49,11 @@ use crate::core::realtime::scaffold::{ /// query is appended in `connect_spec` (Gemini auths by QUERY param, not header). pub(crate) const GEMINI_LIVE_URL: &str = "wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent"; -/// Current Gemini Live model (the half-cascade audio model). Used when -/// `cfg.model` is empty. Pipecat's default is a `gemini-2.5-flash-native-audio-*` -/// preview; the broadly-available stable Live model is `gemini-2.0-flash-live-001`. -pub(crate) const GEMINI_LIVE_DEFAULT_MODEL: &str = "gemini-2.0-flash-live-001"; +/// The Gemini Live model used when `cfg.model` is empty (FRD-023 F-4). It was +/// `gemini-2.0-flash-live-001`, which Google shut down on 2025-12-09; `gemini-3.8-live` is on +/// Google's pricing page as of 2026-09-24. A Bud deployment always names its model +/// (`voice_table.model`), so this default serves only the native `/realtime` path. +pub(crate) const GEMINI_LIVE_DEFAULT_MODEL: &str = "gemini-3.8-live"; /// Gemini Live OUTPUT is 24 kHz mono 16-bit PCM ⇒ 2 B/sample × 24 samples/ms = /// 48 B/ms. (Pipecat: `self._sample_rate = 24000`, output mime @@ -561,6 +562,13 @@ mod tests { assert_eq!(p.model(), GEMINI_LIVE_DEFAULT_MODEL); } + /// F-4 — the default is a model Google still serves. `gemini-2.0-flash-live-001` was shut + /// down on 2025-12-09; a default that no longer exists fails every session that relies on it. + #[test] + fn f4_the_default_model_is_current() { + assert_eq!(GEMINI_LIVE_DEFAULT_MODEL, "gemini-3.8-live"); + } + /// connect_spec: the api key is in the `?key=` QUERY (NOT a header), on the /// exact BidiGenerateContent URL. #[test] diff --git a/gateway/src/core/realtime/nova_sonic/protocol.rs b/gateway/src/core/realtime/nova_sonic/protocol.rs index e626e3bd..3f6abaaa 100644 --- a/gateway/src/core/realtime/nova_sonic/protocol.rs +++ b/gateway/src/core/realtime/nova_sonic/protocol.rs @@ -55,8 +55,9 @@ use crate::core::realtime::scaffold::{ }; use std::sync::Arc; -/// Default Bedrock model id for Nova Sonic (speech-to-speech v1). -pub(crate) const DEFAULT_MODEL: &str = "amazon.nova-sonic-v1:0"; +/// Default Bedrock model id: Nova 2 Sonic (FRD-023 F-4). Nova Sonic v1 (`amazon.nova-sonic-v1:0`) +/// reached end of life on 2026-09-14. +pub(crate) const DEFAULT_MODEL: &str = "amazon.nova-2-sonic-v1:0"; /// Default Nova Sonic voice when none is configured (an English voice from the /// documented `voiceId` set: matthew | tiffany | amy | …). @@ -672,6 +673,12 @@ mod tests { } /// from_config defaults the model + voice when omitted. + /// F-4 — Nova Sonic v1 reached end of life on 2026-09-14; the default is Nova 2 Sonic. + #[test] + fn f4_the_default_model_is_nova_2_sonic() { + assert_eq!(DEFAULT_MODEL, "amazon.nova-2-sonic-v1:0"); + } + #[test] fn from_config_defaults_model_and_voice() { let cfg = RealtimeConfig { diff --git a/gateway/src/core/realtime/openai/config.rs b/gateway/src/core/realtime/openai/config.rs index 93e3507c..ea8e512d 100644 --- a/gateway/src/core/realtime/openai/config.rs +++ b/gateway/src/core/realtime/openai/config.rs @@ -6,6 +6,8 @@ //! - Audio format configuration //! - Turn detection settings +use std::borrow::Cow; + use serde::{Deserialize, Serialize}; /// OpenAI Realtime API WebSocket endpoint. @@ -18,71 +20,67 @@ pub const OPENAI_REALTIME_SAMPLE_RATE: u32 = 24000; // Models // ============================================================================= -/// Supported OpenAI Realtime models. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)] -pub enum OpenAIRealtimeModel { - /// gpt-realtime — the current GA realtime model (default). - #[default] - #[serde(rename = "gpt-realtime")] - GptRealtime, +/// An OpenAI Realtime model id, carried VERBATIM (FRD-023 F-1). +/// +/// This was a closed enum whose parser mapped every id it did not list to `gpt-realtime` — so +/// `gpt-realtime-1.5`, `gpt-realtime-2.1` and `gpt-realtime-2.1-mini` silently became a model +/// that shuts down on 2027-01-20. OpenAI ships realtime models faster than a gateway release, +/// and the id is the vendor's to define: an unknown one must reach the vendor, which can refuse +/// it by name. Only an EMPTY id takes the default. +/// +/// The associated constants keep the names of the enum variants this type replaced. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(transparent)] +pub struct OpenAIRealtimeModel(Cow<'static, str>); + +#[allow(non_upper_case_globals)] +impl OpenAIRealtimeModel { + /// gpt-realtime — the GA realtime model (default). + pub const GptRealtime: Self = Self(Cow::Borrowed("gpt-realtime")); /// gpt-realtime-2 — reasoning-capable realtime model (honors `reasoning.effort`). - #[serde(rename = "gpt-realtime-2")] - GptRealtime2, + pub const GptRealtime2: Self = Self(Cow::Borrowed("gpt-realtime-2")); /// gpt-realtime-mini — smaller / lower-latency realtime model. - #[serde(rename = "gpt-realtime-mini")] - GptRealtimeMini, + pub const GptRealtimeMini: Self = Self(Cow::Borrowed("gpt-realtime-mini")); /// GPT-4o Realtime Preview (DEPRECATED — retained for backward compatibility) - #[serde(rename = "gpt-4o-realtime-preview")] - Gpt4oRealtimePreview, + pub const Gpt4oRealtimePreview: Self = Self(Cow::Borrowed("gpt-4o-realtime-preview")); /// GPT-4o Realtime Preview 2024-10-01 - #[serde(rename = "gpt-4o-realtime-preview-2024-10-01")] - Gpt4oRealtimePreview20241001, + pub const Gpt4oRealtimePreview20241001: Self = + Self(Cow::Borrowed("gpt-4o-realtime-preview-2024-10-01")); /// GPT-4o Realtime Preview 2024-12-17 - #[serde(rename = "gpt-4o-realtime-preview-2024-12-17")] - Gpt4oRealtimePreview20241217, + pub const Gpt4oRealtimePreview20241217: Self = + Self(Cow::Borrowed("gpt-4o-realtime-preview-2024-12-17")); /// GPT-4o Mini Realtime Preview - #[serde(rename = "gpt-4o-mini-realtime-preview")] - Gpt4oMiniRealtimePreview, + pub const Gpt4oMiniRealtimePreview: Self = Self(Cow::Borrowed("gpt-4o-mini-realtime-preview")); /// GPT-4o Mini Realtime Preview 2024-12-17 - #[serde(rename = "gpt-4o-mini-realtime-preview-2024-12-17")] - Gpt4oMiniRealtimePreview20241217, -} + pub const Gpt4oMiniRealtimePreview20241217: Self = + Self(Cow::Borrowed("gpt-4o-mini-realtime-preview-2024-12-17")); -impl OpenAIRealtimeModel { - /// Convert to the API parameter value. + /// The API parameter value — exactly the id the model was built from. #[inline] - pub fn as_str(&self) -> &'static str { - match self { - Self::GptRealtime => "gpt-realtime", - Self::GptRealtime2 => "gpt-realtime-2", - Self::GptRealtimeMini => "gpt-realtime-mini", - Self::Gpt4oRealtimePreview => "gpt-4o-realtime-preview", - Self::Gpt4oRealtimePreview20241001 => "gpt-4o-realtime-preview-2024-10-01", - Self::Gpt4oRealtimePreview20241217 => "gpt-4o-realtime-preview-2024-12-17", - Self::Gpt4oMiniRealtimePreview => "gpt-4o-mini-realtime-preview", - Self::Gpt4oMiniRealtimePreview20241217 => "gpt-4o-mini-realtime-preview-2024-12-17", - } + pub fn as_str(&self) -> &str { + &self.0 } - /// Parse from string, with fallback to default. + /// The id as given (trimmed); the default only when it is empty. Never a substitution. pub fn from_str_or_default(s: &str) -> Self { - match s.to_lowercase().as_str() { - "gpt-realtime" => Self::GptRealtime, - "gpt-realtime-2" => Self::GptRealtime2, - "gpt-realtime-mini" => Self::GptRealtimeMini, - "gpt-4o-realtime-preview" => Self::Gpt4oRealtimePreview, - "gpt-4o-realtime-preview-2024-10-01" => Self::Gpt4oRealtimePreview20241001, - "gpt-4o-realtime-preview-2024-12-17" => Self::Gpt4oRealtimePreview20241217, - "gpt-4o-mini-realtime-preview" => Self::Gpt4oMiniRealtimePreview, - "gpt-4o-mini-realtime-preview-2024-12-17" => Self::Gpt4oMiniRealtimePreview20241217, - _ => Self::default(), + let id = s.trim(); + if id.is_empty() { + Self::default() + } else { + Self(Cow::Owned(id.to_string())) } } } +impl Default for OpenAIRealtimeModel { + fn default() -> Self { + Self::GptRealtime + } +} + impl std::fmt::Display for OpenAIRealtimeModel { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.as_str()) + f.write_str(self.as_str()) } } @@ -304,9 +302,13 @@ mod tests { OpenAIRealtimeModel::from_str_or_default("gpt-4o-realtime-preview"), OpenAIRealtimeModel::Gpt4oRealtimePreview ); - // Default is now the GA gpt-realtime (the deprecated preview is retired). + // F-1: an id the gateway does not list is kept verbatim; only an empty one defaults. + assert_eq!( + OpenAIRealtimeModel::from_str_or_default("unknown").as_str(), + "unknown" + ); assert_eq!( - OpenAIRealtimeModel::from_str_or_default("unknown"), + OpenAIRealtimeModel::from_str_or_default(" "), OpenAIRealtimeModel::GptRealtime ); } diff --git a/gateway/src/core/realtime/openai/mod.rs b/gateway/src/core/realtime/openai/mod.rs index b965d467..5f08e3ca 100644 --- a/gateway/src/core/realtime/openai/mod.rs +++ b/gateway/src/core/realtime/openai/mod.rs @@ -124,9 +124,13 @@ mod tests { OpenAIRealtimeModel::from_str_or_default("gpt-realtime"), OpenAIRealtimeModel::GptRealtime ); - // Default is now the GA gpt-realtime. + // F-1: an id the gateway does not list is kept verbatim; only an empty one defaults. assert_eq!( - OpenAIRealtimeModel::from_str_or_default("unknown"), + OpenAIRealtimeModel::from_str_or_default("unknown").as_str(), + "unknown" + ); + assert_eq!( + OpenAIRealtimeModel::from_str_or_default(" "), OpenAIRealtimeModel::GptRealtime ); } diff --git a/gateway/src/core/realtime/openai/protocol.rs b/gateway/src/core/realtime/openai/protocol.rs index 160935ea..4bc4a2dd 100644 --- a/gateway/src/core/realtime/openai/protocol.rs +++ b/gateway/src/core/realtime/openai/protocol.rs @@ -54,7 +54,7 @@ pub struct OpenAiProtocol { impl OpenAiProtocol { /// The configured model (for the newtype's inherent `model()` accessor). pub fn model(&self) -> OpenAIRealtimeModel { - self.model + self.model.clone() } /// The configured voice (for the newtype's inherent `voice()` accessor). @@ -564,6 +564,35 @@ mod tests { } } + /// TC-XL-08 (F-1) — a model the gateway has never heard of reaches the vendor VERBATIM. A + /// closed enum turned `gpt-realtime-1.5`, `-2.1` and `-2.1-mini` into `gpt-realtime` (shut + /// down 2027-01-20) without a word. + #[test] + fn tc_xl_08_an_unlisted_model_reaches_upstream_verbatim() { + for model in [ + "gpt-realtime-2.1", + "gpt-realtime-2.1-mini", + "gpt-realtime-1.5", + ] { + let cfg = RealtimeConfig { + model: model.into(), + ..base_cfg() + }; + let p = proto(&cfg); + assert_eq!(p.model().as_str(), model); + let ConnectSpec::WebSocket { url, .. } = p.connect_spec(&cfg).unwrap() else { + panic!("expected a WebSocket connect spec"); + }; + assert_eq!(url, format!("{OPENAI_REALTIME_URL}?model={model}")); + } + // Only an EMPTY model takes the default. + let cfg = RealtimeConfig { + model: String::new(), + ..base_cfg() + }; + assert_eq!(proto(&cfg).model().as_str(), "gpt-realtime"); + } + /// THE GOLDEN WIRE ORACLE: every outbound message the protocol serializes /// must be byte-equivalent to the GA wire the live-validated bespoke client /// produced. The expected substrings are ported verbatim from the existing diff --git a/gateway/src/core/realtime/scaffold/session.rs b/gateway/src/core/realtime/scaffold/session.rs index cc2fe124..2af87934 100644 --- a/gateway/src/core/realtime/scaffold/session.rs +++ b/gateway/src/core/realtime/scaffold/session.rs @@ -675,6 +675,48 @@ impl RealtimeSession

{ } } +/// Fold a session update into the session's config (F-3). +/// +/// An update names only what it changes, so a field it leaves unset keeps its value. The +/// server-held fields — the key, the endpoint and its server-config override, reconnection, +/// the trace parent — are never taken from an update unless it sets them: replacing the config +/// wholesale used to drop the key, and the next reconnect dialled without one. +fn merge_session_update(current: &mut RealtimeConfig, update: RealtimeConfig) { + fn keep(slot: &mut Option, new: Option) { + if new.is_some() { + *slot = new; + } + } + if !update.api_key.is_empty() { + current.api_key = update.api_key; + } + if !update.model.trim().is_empty() { + current.model = update.model; + } + keep(&mut current.voice, update.voice); + keep(&mut current.instructions, update.instructions); + keep(&mut current.temperature, update.temperature); + keep( + &mut current.max_response_output_tokens, + update.max_response_output_tokens, + ); + keep(&mut current.input_audio_format, update.input_audio_format); + keep(&mut current.output_audio_format, update.output_audio_format); + keep( + &mut current.input_audio_transcription, + update.input_audio_transcription, + ); + keep(&mut current.turn_detection, update.turn_detection); + keep(&mut current.tools, update.tools); + keep(&mut current.tool_choice, update.tool_choice); + keep(&mut current.modalities, update.modalities); + keep(&mut current.reasoning_effort, update.reasoning_effort); + keep( + &mut current.input_audio_noise_reduction, + update.input_audio_noise_reduction, + ); +} + #[async_trait] impl BaseRealtime for RealtimeSession

{ fn new(config: RealtimeConfig) -> RealtimeResult { @@ -817,7 +859,7 @@ impl BaseRealtime for RealtimeSession

{ } async fn update_session(&mut self, config: RealtimeConfig) -> RealtimeResult<()> { - self.config = config; + merge_session_update(&mut self.config, config); self.push_wires(self.protocol.build_session_config(&self.config, None)) .await } @@ -898,6 +940,53 @@ impl BaseRealtime for RealtimeSession

{ mod tests { use super::*; + /// F-3 — an update MERGES into the session's config. It used to REPLACE it, so the api key, + /// the server-set endpoint override and every field the update did not name were lost, and + /// the next reconnect dialled without a key. + #[tokio::test] + async fn f3_update_session_merges_and_keeps_the_server_held_fields() { + use crate::core::realtime::base::{FunctionDefinition, ToolDefinition}; + use crate::core::realtime::gemini::GeminiProtocol; + let cfg = RealtimeConfig { + provider: "gemini".into(), + api_key: "gkey".into(), + model: "m1".into(), + instructions: Some("be brief".into()), + realtime_endpoint_override: Some("ws://127.0.0.1:9/x".into()), + ..Default::default() + }; + let mut s = RealtimeSession::::new(cfg).unwrap(); + let update = RealtimeConfig { + voice: Some("Puck".into()), + tools: Some(vec![ToolDefinition { + tool_type: "function".into(), + function: FunctionDefinition { + name: "lookup".into(), + description: None, + parameters: None, + }, + }]), + turn_detection: Some(crate::core::realtime::base::TurnDetectionConfig::None), + ..Default::default() + }; + // Not connected: the wire send fails, the merge still happens. + let _ = s.update_session(update).await; + let c = s.config(); + assert_eq!(c.api_key, "gkey"); + assert_eq!(c.model, "m1"); + assert_eq!(c.instructions.as_deref(), Some("be brief")); + assert_eq!( + c.realtime_endpoint_override.as_deref(), + Some("ws://127.0.0.1:9/x") + ); + assert_eq!(c.voice.as_deref(), Some("Puck")); + assert_eq!(c.tools.as_ref().map(Vec::len), Some(1)); + assert!(matches!( + c.turn_detection, + Some(crate::core::realtime::base::TurnDetectionConfig::None) + )); + } + #[test] fn connection_state_helpers_recover_from_poisoned_rwlock() { let state = Arc::new(StdRwLock::new(ConnectionState::Connected)); diff --git a/gateway/src/handlers/realtime/handler.rs b/gateway/src/handlers/realtime/handler.rs index aefdedda..7d4383b3 100644 --- a/gateway/src/handlers/realtime/handler.rs +++ b/gateway/src/handlers/realtime/handler.rs @@ -1007,15 +1007,19 @@ async fn handle_session_update( return true; }; - // Build update config (reuse existing API key) + // Build the update. Only the fields the client named: the provider MERGES it into the + // session's config and keeps the key and everything else (F-3). Tools and turn detection + // used to be dropped here, so an update that changed them changed nothing. let update_config = RealtimeConfig { api_key: String::new(), // Provider should retain existing key - model: config.model.unwrap_or_default(), - voice: config.voice, - instructions: config.instructions, + model: config.model.clone().unwrap_or_default(), + voice: config.voice.clone(), + instructions: config.instructions.clone(), temperature: config.temperature, max_response_output_tokens: config.max_response_tokens, - modalities: config.modalities, + turn_detection: map_turn_detection(&config), + tools: map_tools(&config), + modalities: config.modalities.clone(), reasoning_effort: config.reasoning_effort, // S2S input_audio_noise_reduction: config.input_audio_noise_reduction.clone(), ..Default::default() @@ -1077,45 +1081,10 @@ fn canonical_realtime_provider(provider_name: &str) -> Option<&'static str> { /// field, and the upstream override is injected SEPARATELY by the handler from /// trusted server config only. Keep it that way (no SSRF via client input). pub fn build_realtime_config(api_key: String, config: &RealtimeSessionConfig) -> RealtimeConfig { - use crate::core::realtime::{InputTranscriptionConfig, TurnDetectionConfig}; - - let turn_detection = config.turn_detection.as_ref().map(|td| match td { - crate::handlers::realtime::messages::TurnDetectionConfig::ServerVad { - threshold, - silence_duration_ms, - prefix_padding_ms, - } => TurnDetectionConfig::ServerVad { - threshold: *threshold, - prefix_padding_ms: *prefix_padding_ms, - silence_duration_ms: *silence_duration_ms, - create_response: Some(true), - interrupt_response: Some(true), - }, - crate::handlers::realtime::messages::TurnDetectionConfig::Semantic { eagerness } => { - TurnDetectionConfig::SemanticVad { - eagerness: eagerness.clone(), - create_response: Some(true), - interrupt_response: Some(true), - } - } - crate::handlers::realtime::messages::TurnDetectionConfig::Manual => { - TurnDetectionConfig::None - } - }); + use crate::core::realtime::InputTranscriptionConfig; - let tools = config.tools.as_ref().map(|tools| { - tools - .iter() - .map(|t| crate::core::realtime::ToolDefinition { - tool_type: t.tool_type.clone(), - function: crate::core::realtime::FunctionDefinition { - name: t.function.name.clone(), - description: t.function.description.clone(), - parameters: t.function.parameters.clone(), - }, - }) - .collect() - }); + let turn_detection = map_turn_detection(config); + let tools = map_tools(config); let input_audio_transcription = if config.transcribe_input.unwrap_or(true) { Some(InputTranscriptionConfig { @@ -1153,6 +1122,53 @@ pub fn build_realtime_config(api_key: String, config: &RealtimeSessionConfig) -> } } +/// The client's turn detection in the provider vocabulary (shared by config and update, F-3). +fn map_turn_detection( + config: &RealtimeSessionConfig, +) -> Option { + use crate::core::realtime::TurnDetectionConfig; + config.turn_detection.as_ref().map(|td| match td { + crate::handlers::realtime::messages::TurnDetectionConfig::ServerVad { + threshold, + silence_duration_ms, + prefix_padding_ms, + } => TurnDetectionConfig::ServerVad { + threshold: *threshold, + prefix_padding_ms: *prefix_padding_ms, + silence_duration_ms: *silence_duration_ms, + create_response: Some(true), + interrupt_response: Some(true), + }, + crate::handlers::realtime::messages::TurnDetectionConfig::Semantic { eagerness } => { + TurnDetectionConfig::SemanticVad { + eagerness: eagerness.clone(), + create_response: Some(true), + interrupt_response: Some(true), + } + } + crate::handlers::realtime::messages::TurnDetectionConfig::Manual => { + TurnDetectionConfig::None + } + }) +} + +/// The client's tools in the provider vocabulary (shared by config and update, F-3). +fn map_tools(config: &RealtimeSessionConfig) -> Option> { + config.tools.as_ref().map(|tools| { + tools + .iter() + .map(|t| crate::core::realtime::ToolDefinition { + tool_type: t.tool_type.clone(), + function: crate::core::realtime::FunctionDefinition { + name: t.function.name.clone(), + description: t.function.description.clone(), + parameters: t.function.parameters.clone(), + }, + }) + .collect() + }) +} + #[cfg(test)] mod tests { use super::*; @@ -1574,4 +1590,109 @@ mod frd023_native_tests { assert!(!disconnected.load(Ordering::SeqCst), "and still connected"); assert_eq!(session_id.as_deref(), Some("sess-1")); } + + /// Records the config each `update_session` receives. + struct RecordingRt(Arc>>); + + #[async_trait::async_trait] + impl BaseRealtime for RecordingRt { + fn new(_c: RealtimeConfig) -> RealtimeResult { + unreachable!() + } + async fn connect(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn disconnect(&mut self) -> RealtimeResult<()> { + Ok(()) + } + fn is_ready(&self) -> bool { + true + } + fn get_connection_state(&self) -> ConnectionState { + ConnectionState::Connected + } + async fn send_audio(&mut self, _a: bytes::Bytes) -> RealtimeResult<()> { + Ok(()) + } + async fn send_text(&mut self, _t: &str) -> RealtimeResult<()> { + Ok(()) + } + async fn create_response(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn cancel_response(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn commit_audio_buffer(&mut self) -> RealtimeResult<()> { + Ok(()) + } + async fn clear_audio_buffer(&mut self) -> RealtimeResult<()> { + Ok(()) + } + fn on_transcript(&mut self, _c: TranscriptCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_audio(&mut self, _c: AudioOutputCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_error(&mut self, _c: RealtimeErrorCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_function_call(&mut self, _c: FunctionCallCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_speech_event(&mut self, _c: SpeechEventCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_response_done(&mut self, _c: ResponseDoneCallback) -> RealtimeResult<()> { + Ok(()) + } + fn on_reconnection(&mut self, _c: ReconnectionCallback) -> RealtimeResult<()> { + Ok(()) + } + async fn update_session(&mut self, c: RealtimeConfig) -> RealtimeResult<()> { + self.0.lock().unwrap().push(c); + Ok(()) + } + async fn submit_function_result(&mut self, _id: &str, _r: &str) -> RealtimeResult<()> { + Ok(()) + } + fn get_provider_info(&self) -> serde_json::Value { + serde_json::json!({}) + } + } + + /// F-3 — a native `update_session` carries the tools and the turn detection it was sent; + /// before, both were dropped and the vendor kept the old ones without a word. + #[tokio::test] + async fn f3_update_session_carries_tools_and_turn_detection() { + let seen = Arc::new(std::sync::Mutex::new(Vec::new())); + let mut provider: Option> = + Some(Box::new(RecordingRt(Arc::clone(&seen)))); + let (tx, _rx) = mpsc::channel(8); + let update: RealtimeSessionConfig = serde_json::from_value(serde_json::json!({ + "voice": "marin", + "turn_detection": {"mode": "semantic", "eagerness": "low"}, + "tools": [{"type": "function", "function": {"name": "lookup", "description": "d", + "parameters": {"type": "object"}}}] + })) + .unwrap(); + + handle_session_update(update, &mut provider, &tx).await; + + let seen = seen.lock().unwrap(); + let cfg = seen.last().expect("update_session was called"); + assert_eq!(cfg.voice.as_deref(), Some("marin")); + let tools = cfg.tools.as_ref().expect("tools carried"); + assert_eq!(tools[0].function.name, "lookup"); + assert!( + matches!( + cfg.turn_detection, + Some(crate::core::realtime::TurnDetectionConfig::SemanticVad { ref eagerness, .. }) + if eagerness.as_deref() == Some("low") + ), + "turn detection carried: {:?}", + cfg.turn_detection + ); + } } diff --git a/gateway/src/plugin/builtin/mod.rs b/gateway/src/plugin/builtin/mod.rs index 363e0778..1ee371b7 100644 --- a/gateway/src/plugin/builtin/mod.rs +++ b/gateway/src/plugin/builtin/mod.rs @@ -1276,7 +1276,7 @@ fn gemini_realtime_metadata() -> ProviderMetadata { .with_description( "Google Gemini Live (BidiGenerateContent) — speech-to-speech (base64+JSON, 16k in / 24k out; MULTI-FRAME serverContent; session resumption; server VAD)", ) - .with_models(["gemini-2.0-flash-live-001"]) + .with_models(["gemini-3.8-live"]) .with_aliases(["gemini-live", "google"]) .with_features([ "full-duplex", @@ -1302,7 +1302,7 @@ fn nova_sonic_realtime_metadata() -> ProviderMetadata { .with_description( "AWS Nova Sonic — Amazon's speech-to-speech model (base64-PCM + JSON events, 16k in / 24k out; BedrockBidi: an Amazon Bedrock InvokeModelWithBidirectionalStream HTTP/2 event stream; AWS SigV4 via aws-config, NO api-key; server VAD)", ) - .with_models(["amazon.nova-sonic-v1:0"]) + .with_models(["amazon.nova-2-sonic-v1:0"]) .with_aliases(["nova-sonic", "aws"]) .with_features(["full-duplex", "function-calling", "turn-detection", "barge-in"]) } diff --git a/gateway/tests/realtime_provider_matrix.rs b/gateway/tests/realtime_provider_matrix.rs index 0c4e3afb..b87b503f 100644 --- a/gateway/tests/realtime_provider_matrix.rs +++ b/gateway/tests/realtime_provider_matrix.rs @@ -60,10 +60,9 @@ const KEYLESS_PROVIDERS: [&str; 1] = ["nova_sonic"]; /// - **azure** (`azure/protocol.rs`): `from_config` REQUIRES both `endpoint` /// (the Azure resource) AND a non-empty `model` (the deployment name) — either /// missing ⇒ `InvalidConfiguration`. Both are set here; -/// - **openai** (`openai/protocol.rs`): `from_config` parses `model` into the -/// `OpenAIRealtimeModel` enum via `from_str_or_default`, which DEFAULTS an -/// unknown string to `gpt-realtime` rather than erroring — so any model -/// string constructs. We pass the real GA default `"gpt-realtime"` explicitly; +/// - **openai** (`openai/protocol.rs`): `from_config` carries `model` VERBATIM +/// (`OpenAIRealtimeModel`, FRD-023 F-1) — any non-empty id constructs and reaches +/// the vendor as given. We pass the real GA default `"gpt-realtime"` explicitly; /// - **grok / inworld / deepgram**: take a raw/lenient model string (inworld and /// deepgram don't even require it at construction — only `connect_spec`/the /// Settings wire consult it). A descriptive model is supplied anyway; From 51f9f4c140a9fb4201a30b9f26f5a52dddc60af2 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 05:41:57 +0530 Subject: [PATCH 08/17] feat(gateway): xAI realtime deployments over the GA relay (FRD-023 RT7.3) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `grok` joins the relay vendors: upstream wss://api.x.ai/v1/realtime?model= with the deployment's key as Bearer; an api_base is SSRF-validated like OpenAI's. xAI opens its session with `conversation.created` and never sends `session.created`, so that event now starts the session (the deployment defaults are applied once) and the client is given a GA `session.created`. Cumulative `…input_audio_transcription.updated` events are forwarded verbatim, and nothing waits on `rate_limits.updated` (TC-XL-06). Co-Authored-By: Claude Opus 5.5 --- .../src/handlers/openai_realtime/policy.rs | 46 +++++++- .../src/handlers/openai_realtime/session.rs | 56 ++++++--- .../src/handlers/openai_realtime/upstream.rs | 46 ++++++-- gateway/tests/openai_realtime_relay.rs | 109 +++++++++++++++++- 4 files changed, 230 insertions(+), 27 deletions(-) diff --git a/gateway/src/handlers/openai_realtime/policy.rs b/gateway/src/handlers/openai_realtime/policy.rs index 0c5e540f..4f5bf423 100644 --- a/gateway/src/handlers/openai_realtime/policy.rs +++ b/gateway/src/handlers/openai_realtime/policy.rs @@ -297,8 +297,16 @@ pub fn client_event<'a>(raw: &'a str, rules: &ClientRules) -> ClientOutcome<'a> #[derive(Debug, PartialEq)] pub enum Tap { None, - SessionCreated { vendor_session_id: Option }, - SessionUpdated { event_id: Option }, + SessionCreated { + vendor_session_id: Option, + }, + /// xAI opens a session with `conversation.created` and no `session.created` (WP-RT7.3). + ConversationCreated { + vendor_session_id: Option, + }, + SessionUpdated { + event_id: Option, + }, ResponseDone(Value), TranscriptionCompleted(Value), Error(Value), @@ -350,6 +358,18 @@ pub fn vendor_event<'a>(raw: &'a str, deployment: &str) -> VendorOutcome<'a> { }; VendorOutcome::Forward(Cow::Owned(event.to_string()), tap) } + "conversation.created" => { + let vendor_session_id = serde_json::from_str::(raw).ok().and_then(|v| { + v.get("conversation") + .and_then(|c| c.get("id")) + .and_then(Value::as_str) + .map(str::to_string) + }); + VendorOutcome::Forward( + Cow::Borrowed(raw), + Tap::ConversationCreated { vendor_session_id }, + ) + } "response.done" => match serde_json::from_str::(raw) { Ok(v) => VendorOutcome::Forward(Cow::Borrowed(raw), Tap::ResponseDone(v)), Err(_) => VendorOutcome::Invalid, @@ -460,6 +480,28 @@ pub fn defaults_update( ) } +/// A GA `session.created` for a vendor that bootstraps without one (xAI's +/// `conversation.created`). Only what the gateway knows: the deployment name as the model (as +/// every relayed `session.created` carries it) and the vendor's session id. +pub fn synthesized_session_created( + event_id: &str, + deployment: &str, + session_type: &str, + vendor_session_id: Option<&str>, +) -> String { + serde_json::json!({ + "type": "session.created", + "event_id": event_id, + "session": { + "object": "realtime.session", + "type": session_type, + "id": vendor_session_id, + "model": deployment, + } + }) + .to_string() +} + /// An OpenAI `error` event from the gateway. pub fn error_event( event_id: &str, diff --git a/gateway/src/handlers/openai_realtime/session.rs b/gateway/src/handlers/openai_realtime/session.rs index 64b3cdc8..08bdb8ba 100644 --- a/gateway/src/handlers/openai_realtime/session.rs +++ b/gateway/src/handlers/openai_realtime/session.rs @@ -561,6 +561,8 @@ struct Relay<'a> { /// Client frames wait here until the vendor has applied the deployment defaults (§5.6). held: VecDeque<(String, bool)>, ready: bool, + /// The vendor's session exists (`session.created`, or xAI's `conversation.created`). + started: bool, /// The `event_id` of the defaults update in flight, and its deadline. awaiting_defaults: Option, hold_deadline: Option, @@ -712,6 +714,29 @@ impl Relay<'_> { } } + /// The vendor's session exists: apply the deployment defaults (§5.6), once. + async fn session_started(&mut self, vendor_session_id: Option) -> Result<(), End> { + if self.started { + return Ok(()); + } + self.started = true; + self.meter.set_vendor_session_id(vendor_session_id); + let event_id = format!("evt_bud_defaults_{}", uuid::Uuid::new_v4().simple()); + match policy::defaults_update( + self.p.settings.as_ref(), + self.p.endpoint.model.as_deref(), + &event_id, + ) { + Some(update) => { + self.send_to_vendor(update).await?; + self.awaiting_defaults = Some(event_id); + self.hold_deadline = Some(Instant::now() + self.timings.hold); + Ok(()) + } + None => self.release_held().await, + } + } + async fn on_vendor_text(&mut self, raw: &str) -> Result<(), End> { self.last_activity = Instant::now(); match policy::vendor_event(raw, &self.p.endpoint_name) { @@ -721,23 +746,23 @@ impl Relay<'_> { let text = text.into_owned(); match tap { Tap::SessionCreated { vendor_session_id } => { - self.meter.set_vendor_session_id(vendor_session_id); self.to_client(text).await?; - let event_id = - format!("evt_bud_defaults_{}", uuid::Uuid::new_v4().simple()); - match policy::defaults_update( - self.p.settings.as_ref(), - self.p.endpoint.model.as_deref(), - &event_id, - ) { - Some(update) => { - self.send_to_vendor(update).await?; - self.awaiting_defaults = Some(event_id); - self.hold_deadline = Some(Instant::now() + self.timings.hold); - Ok(()) - } - None => self.release_held().await, + self.session_started(vendor_session_id).await + } + Tap::ConversationCreated { vendor_session_id } => { + self.to_client(text).await?; + if self.started { + return Ok(()); } + // A GA client waits for `session.created`; xAI never sends one. + self.to_client(policy::synthesized_session_created( + &next_event_id(), + &self.p.endpoint_name, + &self.p.rules.session_type, + vendor_session_id.as_deref(), + )) + .await?; + self.session_started(vendor_session_id).await } Tap::SessionUpdated { .. } => { self.to_client(text).await?; @@ -863,6 +888,7 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< timings: timings.clone(), held: VecDeque::new(), ready: false, + started: false, awaiting_defaults: None, // The vendor must say `session.created` within the hold window as well. hold_deadline: Some(now + timings.connect), diff --git a/gateway/src/handlers/openai_realtime/upstream.rs b/gateway/src/handlers/openai_realtime/upstream.rs index ff52a84d..ff5f376d 100644 --- a/gateway/src/handlers/openai_realtime/upstream.rs +++ b/gateway/src/handlers/openai_realtime/upstream.rs @@ -12,10 +12,14 @@ use tokio::net::TcpStream; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use tokio_tungstenite::{MaybeTlsStream, WebSocketStream}; -/// The relay-capable vendors (D-2). Everything else needs the translate engine (RT7). -pub const RELAY_VENDORS: &[&str] = &["openai", "azure_openai"]; +/// The relay-capable vendors (D-2): they speak OpenAI Realtime GA themselves. xAI joined in RT7 +/// (WP-RT7.3); the vendors without a GA surface are served by the translate engine +/// ([`super::facade`]). +pub const RELAY_VENDORS: &[&str] = &["openai", "azure_openai", "grok"]; const OPENAI_DEFAULT_BASE: &str = "https://api.openai.com/v1"; +/// xAI's GA-compatible realtime surface (CONTRACTS C7): `wss://api.x.ai/v1/realtime?model=…`. +const XAI_DEFAULT_BASE: &str = "https://api.x.ai/v1"; pub type UpstreamSocket = WebSocketStream>; @@ -41,10 +45,9 @@ pub enum UpstreamError { impl std::fmt::Display for UpstreamError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Self::UnsupportedVendor(v) => write!( - f, - "realtime sessions are not yet served for vendor '{v}' (relay vendors: openai, azure_openai)" - ), + Self::UnsupportedVendor(v) => { + write!(f, "realtime sessions are not served for vendor '{v}'") + } Self::MissingCredential => write!(f, "the deployment has no usable vendor credential"), Self::MissingModel => write!(f, "the deployment names no vendor model"), Self::MissingApiBase => write!( @@ -134,12 +137,19 @@ pub fn build( }; match vendor.as_str() { - "openai" => { + // xAI speaks GA on its own host (WP-RT7.3): Bearer, the vendor model in the query. Its + // quirks are on the event stream, not the handshake (`policy::vendor_event`). + "openai" | "grok" => { + let default_base = if vendor == "grok" { + XAI_DEFAULT_BASE + } else { + OPENAI_DEFAULT_BASE + }; let api_base = endpoint .api_base .as_deref() .filter(|b| !b.trim().is_empty()); - let base = to_ws_base(api_base.unwrap_or(OPENAI_DEFAULT_BASE))?; + let base = to_ws_base(api_base.unwrap_or(default_base))?; Ok(UpstreamRequest { url: format!("{base}/realtime?{}", query(model)?), headers: vec![("authorization", format!("Bearer {credential}"))], @@ -296,6 +306,26 @@ mod tests { assert!(req.needs_ssrf_check); } + /// TC-XL-06 (unit half) — xAI is relayed: its GA endpoint, the vendor model verbatim, Bearer, + /// and no SSRF check for the vendor constant; an `api_base` is validated like OpenAI's. + #[test] + fn tc_xl_06_xai_url_and_bearer() { + let req = build(&endpoint("grok", None, Some("grok-voice-2")), false).unwrap(); + assert_eq!(req.url, "wss://api.x.ai/v1/realtime?model=grok-voice-2"); + assert_eq!( + req.headers, + vec![("authorization", "Bearer sk-vendor".to_string())] + ); + assert!(!req.needs_ssrf_check); + let req = build( + &endpoint("grok", Some("https://xai-proxy.example/v1"), Some("m")), + false, + ) + .unwrap(); + assert_eq!(req.url, "wss://xai-proxy.example/v1/realtime?model=m"); + assert!(req.needs_ssrf_check); + } + /// TC-UP-03 — an `api_base` has its scheme converted and keeps its path. #[test] fn tc_up_03_api_base_scheme_conversion() { diff --git a/gateway/tests/openai_realtime_relay.rs b/gateway/tests/openai_realtime_relay.rs index f6b7bc93..f14bdc75 100644 --- a/gateway/tests/openai_realtime_relay.rs +++ b/gateway/tests/openai_realtime_relay.rs @@ -218,6 +218,10 @@ struct Behaviour { /// Answer `input_audio_buffer.commit` with a completed transcription carrying this usage. transcription_usage: Option, usage: Json, + /// Speak like xAI (TC-XL-06): bootstrap with `conversation.created` instead of + /// `session.created`, answer a commit with CUMULATIVE + /// `conversation.item.input_audio_transcription.updated` events, never `rate_limits.updated`. + xai: bool, } impl Default for Behaviour { @@ -231,6 +235,7 @@ impl Default for Behaviour { go_silent_after: None, transcription_usage: None, usage: documented_usage(), + xai: false, } } } @@ -303,8 +308,13 @@ impl MockVendor { .last() .and_then(|(pq, _)| pq.split("model=").nth(1).map(str::to_string)) .unwrap_or_default(); - let created = json!({"type": "session.created", "event_id": "evt_v0", - "session": {"id": "sess_vendor_1", "object": "realtime.session", "type": "realtime", "model": model}}); + let created = if b.xai { + json!({"type": "conversation.created", "event_id": "evt_x0", + "conversation": {"id": "conv_xai_1", "object": "realtime.conversation"}}) + } else { + json!({"type": "session.created", "event_id": "evt_v0", + "session": {"id": "sess_vendor_1", "object": "realtime.session", "type": "realtime", "model": model}}) + }; if ws .send(Message::Text(created.to_string().into())) .await @@ -348,6 +358,14 @@ impl MockVendor { "response": {"id": rid, "status": "completed", "usage": b.usage, "output": [{"type": "message", "content": [{"type": "output_audio", "transcript": "hello there"}]}]}})); } + "input_audio_buffer.commit" if b.xai => { + for partial in ["hel", "hello", "hello world"] { + out.push(json!({"type": "conversation.item.input_audio_transcription.updated", + "item_id": "item_x", "content_index": 0, "transcript": partial})); + } + out.push(json!({"type": "conversation.item.input_audio_transcription.completed", + "item_id": "item_x", "content_index": 0, "transcript": "hello world"})); + } "input_audio_buffer.commit" => { if let Some(u) = &b.transcription_usage { out.push(json!({"type": "conversation.item.input_audio_transcription.completed", @@ -1751,6 +1769,93 @@ async fn duration_priced_sessions_bill_segments() { } } +// ============================================================================================= +// xAI — TC-XL-06 (relayed, not translated) +// ============================================================================================= + +/// TC-XL-06 — an xAI deployment is RELAYED: its upstream is the xAI GA URL with the vendor model +/// and Bearer; xAI's `conversation.created` bootstrap starts the session (the client gets a +/// GA `session.created`, the deployment defaults are applied); cumulative +/// `…input_audio_transcription.updated` events reach the client verbatim and in order; and no +/// `rate_limits.updated` is needed for anything. +#[tokio::test] +async fn tc_xl_06_xai_relay_quirks() { + let cap = Capture::install(); + let vendor = MockVendor::start(Behaviour { + xai: true, + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "grok-rt", + "a1a1a1a1-0000-4000-8000-0000000000a6", + rt_entry( + &vendor, + json!({"vendor": "grok", "model": "grok-voice-2", + "config": {"realtime": {"defaults": {"voice": "ara"}}}}), + ), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "grok-rt").await; + let created = until_type(&mut c, "session.created").await; + assert_eq!(created["session"]["model"], "grok-rt"); + // The defaults went to xAI once its session existed. + for _ in 0..50 { + if !vendor.frames_of("session.update").is_empty() { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + let defaults = &vendor.frames_of("session.update")[0]; + assert_eq!(defaults["session"]["audio"]["output"]["voice"], "ara"); + + send(&mut c, json!({"type": "input_audio_buffer.commit"})).await; + let mut partials = Vec::new(); + loop { + let v = next_json(&mut c).await; + match v["type"].as_str() { + Some("conversation.item.input_audio_transcription.updated") => { + partials.push(v["transcript"].as_str().unwrap().to_string()) + } + Some("conversation.item.input_audio_transcription.completed") => break, + _ => {} + } + } + assert_eq!(partials, vec!["hel", "hello", "hello world"]); + + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; + // No `rate_limits.updated` ever arrives, and the session is fine without it. + keep_alive(&mut c, Duration::from_millis(700)).await; + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; + + let (pq, headers) = vendor.upgrades()[0].clone(); + assert_eq!(pq, "/v1/realtime?model=grok-voice-2"); + assert_eq!( + headers.get("authorization").map(String::as_str), + Some(&*format!("Bearer {VENDOR_KEY}")) + ); + // Two responses (and the unpriced transcription turn xAI reported without usage). + let turns = cap.wait_for("voice.turn", 3).await; + let responses: Vec<&SpanData> = turns + .iter() + .filter(|t| text(t, "bud.voice.rt.component").as_deref() == Some("response")) + .collect(); + assert_eq!(responses.len(), 2); + for t in responses { + assert_eq!(text(t, "bud.voice.rt.vendor").as_deref(), Some("grok")); + assert_eq!( + text(t, "bud.voice.vendor_session_id").as_deref(), + Some("conv_xai_1") + ); + assert!(number(t, "bud.voice.cost").is_some()); + } +} + // ============================================================================================= // Credentials never logged — TC-SEC-07 // ============================================================================================= From 548618c759973aaf475871de926490b64865b1b8 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 06:13:28 +0530 Subject: [PATCH 09/17] fix(gateway): a refused /v1/realtime upgrade speaks OpenAI's error envelope MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The per-IP and global connection limits answered a plain-text 429/503 with no Retry-After. On /v1/realtime the client is an OpenAI SDK, which reads error.code and backs off on the header, so those refusals now use the relay's own envelope (rate_limit_exceeded / server_at_capacity, Retry-After: 1); native routes keep their text body and gain the header. Found by the TC-COMPAT run under concurrent load (FRD-023 §5.2). Co-Authored-By: Claude Opus 5.5 --- gateway/src/middleware/connection_limit.rs | 107 ++++++++++++++++++++- 1 file changed, 103 insertions(+), 4 deletions(-) diff --git a/gateway/src/middleware/connection_limit.rs b/gateway/src/middleware/connection_limit.rs index 440765d9..de42f526 100644 --- a/gateway/src/middleware/connection_limit.rs +++ b/gateway/src/middleware/connection_limit.rs @@ -95,6 +95,7 @@ pub async fn connection_limit_middleware( } let client_ip = addr.ip(); + let openai_path = crate::middleware::auth::is_openai_compatible_path(request.uri().path()); // Try to acquire a connection slot match state.try_acquire_connection(client_ip) { @@ -115,26 +116,48 @@ pub async fn connection_limit_middleware( ip = %client_ip, "Rejecting connection: global limit reached" ); - ( + refusal( + openai_path, StatusCode::SERVICE_UNAVAILABLE, + "server_at_capacity", "Server at capacity. Please try again later.", ) - .into_response() } Err(ConnectionLimitError::PerIpLimitReached) => { tracing::warn!( ip = %client_ip, "Rejecting connection: per-IP limit reached" ); - ( + refusal( + openai_path, StatusCode::TOO_MANY_REQUESTS, + "rate_limit_exceeded", "Too many connections from your IP address.", ) - .into_response() } } } +/// A refused upgrade, with `Retry-After: 1` either way. On the OpenAI-compatible routes +/// (`/v1/realtime`) the client is an OpenAI SDK, which reads `error.code` and backs off on the +/// header, so it gets the relay's own error envelope (FRD-023 §5.2); the native routes keep their +/// text body. +fn refusal(openai_path: bool, status: StatusCode, code: &'static str, message: &str) -> Response { + if openai_path { + return crate::handlers::openai_realtime::handshake::HandshakeError::new( + status, code, message, + ) + .retry_after(1) + .into_response(); + } + let mut response = (status, message.to_string()).into_response(); + response.headers_mut().insert( + axum::http::header::RETRY_AFTER, + axum::http::HeaderValue::from_static("1"), + ); + response +} + #[cfg(test)] mod tests { use super::*; @@ -300,6 +323,82 @@ mod tests { ); } + /// FRD-023 §5.2: a refusal on `/v1/realtime` reaches an OpenAI SDK, which reads `error.code` + /// and backs off on `Retry-After` — a plain-text 429 gives it neither. `/ws` keeps its native + /// body (its clients read text) and gains the hint too. + #[tokio::test] + async fn a_refusal_on_the_openai_path_uses_the_openai_envelope_and_retry_after() { + use axum::extract::connect_info::MockConnectInfo; + use tower::ServiceExt; + for (max_ws, per_ip, want, code) in [ + ( + Some(1000), + 1, + StatusCode::TOO_MANY_REQUESTS, + "rate_limit_exceeded", + ), + ( + Some(1), + 100, + StatusCode::SERVICE_UNAVAILABLE, + "server_at_capacity", + ), + ] { + let state = AppState::new(limit_test_config(max_ws, per_ip)).await; + let held: Arc>> = Default::default(); + let h = held.clone(); + let handler = axum::routing::get( + move |axum::Extension(slot): axum::Extension| { + h.lock().unwrap().push(slot); + async { StatusCode::SWITCHING_PROTOCOLS } + }, + ); + let app = axum::Router::new() + .route("/v1/realtime", handler.clone()) + .route("/ws", handler) + .layer(axum::middleware::from_fn_with_state( + state.clone(), + connection_limit_middleware, + )) + .layer(MockConnectInfo(SocketAddr::from(([10, 0, 0, 9], 4242)))); + let upgrade = |path: &str| { + Request::builder() + .uri(path) + .header("upgrade", "websocket") + .body(Body::empty()) + .unwrap() + }; + let first = app.clone().oneshot(upgrade("/v1/realtime")).await.unwrap(); + assert_eq!(first.status(), StatusCode::SWITCHING_PROTOCOLS); + + let refused = app.clone().oneshot(upgrade("/v1/realtime")).await.unwrap(); + assert_eq!(refused.status(), want); + assert_eq!(refused.headers()["retry-after"], "1"); + let body = axum::body::to_bytes(refused.into_body(), 4096) + .await + .unwrap(); + let json: serde_json::Value = + serde_json::from_slice(&body).expect("an OpenAI error envelope"); + assert_eq!(json["error"]["code"], code, "{json}"); + assert!( + json["error"]["message"] + .as_str() + .is_some_and(|m| !m.is_empty()) + ); + + let native = app.clone().oneshot(upgrade("/ws")).await.unwrap(); + assert_eq!(native.status(), want); + assert_eq!(native.headers()["retry-after"], "1"); + let body = axum::body::to_bytes(native.into_body(), 4096) + .await + .unwrap(); + assert!( + serde_json::from_slice::(&body).is_err(), + "native text body" + ); + } + } + /// A live session keeps its slot until it ends. #[tokio::test] async fn a_live_session_holds_its_slot() { From cdf025ced218e0320b8dee596bf9b7a3d6c2689c Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 06:29:36 +0530 Subject: [PATCH 10/17] feat(gateway): realtime providers report usage and goAway, reconnect on plan, sign Nova with deployment keys (FRD-023 RT7.0-RT7.2) The S2S scaffold gains what the translate engine needs: - S2sEvent::Usage (per-modality tokens or seconds; cumulative or delta), ItemAdded, ItemDone and GoAway, and a raw event tap (BaseRealtime::on_event) that sees every normalized event in wire order. - Planned reconnects: a vendor goAway, or a connection cap (Nova Sonic's 8 minutes, replaced 30 s early), replaces the connection at the next turn boundary with no backoff and without counting as a failure; the resumption handle is carried. The reconnect is announced only once the session is ready again. - Per-protocol input sample rates (Gemini, Nova, ElevenLabs 16 kHz; Hume its configured rate). - Gemini: usageMetadata becomes one cumulative Usage, first in its message's events; goAway carries its deadline. - Nova 2 Sonic: usageEvent becomes one delta Usage (a totals-only report is billed as the difference, per connection); the assistant audio block is an item; history is replayed as text blocks on a new stream. - The Bedrock factory can sign with a deployment's static key pair and nothing else (no environment, shared config or instance identity), and in Bud mode refuses to dial without one. The stream now opens in the background: the SDK's send() returns only at the first output event, which Nova sends only after input, so awaiting it in connect deadlocked until the dial timeout. - bud-auth keeps provider_params.region for vendor nova_sonic. Co-Authored-By: Claude Opus 5.5 --- bud-auth/src/credentials.rs | 30 +- gateway/src/core/realtime/base.rs | 54 +++ gateway/src/core/realtime/deepgram/mod.rs | 11 + .../src/core/realtime/deepgram/protocol.rs | 5 + gateway/src/core/realtime/elevenlabs/mod.rs | 11 + .../src/core/realtime/elevenlabs/protocol.rs | 5 + gateway/src/core/realtime/gemini/mod.rs | 11 + gateway/src/core/realtime/gemini/protocol.rs | 240 +++++++++--- gateway/src/core/realtime/hume/client.rs | 11 + gateway/src/core/realtime/hume/protocol.rs | 4 + gateway/src/core/realtime/mod.rs | 14 +- gateway/src/core/realtime/nova_sonic/mod.rs | 29 ++ .../src/core/realtime/nova_sonic/protocol.rs | 259 ++++++++++++- gateway/src/core/realtime/scaffold/event.rs | 30 ++ gateway/src/core/realtime/scaffold/mock.rs | 130 +++++++ gateway/src/core/realtime/scaffold/mod.rs | 4 +- .../src/core/realtime/scaffold/protocol.rs | 13 + gateway/src/core/realtime/scaffold/session.rs | 122 +++++- .../src/core/realtime/scaffold/transport.rs | 357 +++++++++++++++--- gateway/tests/realtime_full_integration.rs | 2 + 20 files changed, 1220 insertions(+), 122 deletions(-) diff --git a/bud-auth/src/credentials.rs b/bud-auth/src/credentials.rs index fdd8dceb..62d590bc 100644 --- a/bud-auth/src/credentials.rs +++ b/bud-auth/src/credentials.rs @@ -427,7 +427,8 @@ pub fn allowed_provider_params(vendor: &str) -> &'static [&'static str] { .replace('-', "_") .as_str() { - "aws_polly" | "aws_transcribe" => &["region"], + // Nova 2 Sonic (realtime, FRD-023 RT7.2) signs Bedrock requests in this region. + "aws_polly" | "aws_transcribe" | "nova_sonic" => &["region"], "google" => &["project_id", "location"], "azure_openai" => &["api_version"], _ => &[], @@ -1247,6 +1248,33 @@ mod tests { assert!(allowed_provider_params("deepgram").is_empty()); } + /// FRD-023 RT7.2 (CONTRACTS C7) — a Nova 2 Sonic realtime entry carries its REQUIRED region + /// in `provider_params`; dropping it at parse would leave the session unable to sign. + #[test] + fn nova_sonic_keeps_its_region_and_its_key_pair() { + assert_eq!(allowed_provider_params("nova_sonic"), &["region"]); + assert_eq!(allowed_provider_params("nova-sonic"), &["region"]); + let json = serde_json::json!({ "ep-nova": { + "vendor": "nova_sonic", + "credential": encrypt_like_budapp( + r#"{"access_key_id":"AKIDNOVA","secret_access_key":"s3cr3t"}"# + ), + "endpoints": ["realtime_session"], + "model": "amazon.nova-2-sonic-v1:0", + "provider_params": { "region": "us-east-1", "endpoint_override": "https://evil" }, + }}) + .to_string(); + let map = parse_voice_blob(&json, &decryptor()).unwrap(); + let ep = map.get("ep-nova").unwrap(); + assert_eq!(ep.provider_param("region"), Some("us-east-1")); + assert_eq!(ep.provider_param("endpoint_override"), None); + let parts = ep.credential_parts.as_ref().expect("the AWS pair splits"); + assert_eq!( + parts.get("access_key_id").map(String::as_str), + Some("AKIDNOVA") + ); + } + #[test] fn aws_regions_are_checked_against_the_published_shape() { for ok in [ diff --git a/gateway/src/core/realtime/base.rs b/gateway/src/core/realtime/base.rs index d2dee0bc..ba962b8e 100644 --- a/gateway/src/core/realtime/base.rs +++ b/gateway/src/core/realtime/base.rs @@ -295,6 +295,35 @@ pub struct RealtimeConfig { /// engine opens a fresh root span). Server-set only; other providers ignore it. #[serde(default, skip_serializing_if = "Option::is_none")] pub trace: Option, + + /// SERVER-SET: reconnect proactively after a connection has lived this long (a vendor's + /// connection cap — Nova Sonic's 8 minutes). `None` ⇒ the protocol's own cap + /// ([`RealtimeProtocol::max_connection`](crate::core::realtime::scaffold::RealtimeProtocol::max_connection)). + /// Never read from a client message (`serde(skip)`). + #[serde(skip)] + pub max_connection: Option, +} + +/// A static AWS key pair from a Bud deployment's credential (FRD-023 RT7.2: Nova Sonic signs with +/// the DEPLOYMENT's keys, never the gateway's own AWS identity). `Debug` is redacted. +#[derive(Clone, PartialEq, Eq)] +pub struct AwsStaticCredentials { + pub access_key_id: String, + pub secret_access_key: String, + pub session_token: Option, +} + +impl fmt::Debug for AwsStaticCredentials { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AwsStaticCredentials") + .field("access_key_id", &"[redacted]") + .field("secret_access_key", &"[redacted]") + .field( + "session_token", + &self.session_token.as_ref().map(|_| "[redacted]"), + ) + .finish() + } } /// Configuration for input audio transcription. @@ -580,6 +609,13 @@ pub struct ReconnectionEvent { pub type ReconnectionCallback = Arc Pin + Send>> + Send + Sync>; +/// Receives every normalized provider event, in order (FRD-023 RT7, the translate engine). +pub type S2sEventCallback = Arc< + dyn Fn(crate::core::realtime::scaffold::S2sEvent) -> Pin + Send>> + + Send + + Sync, +>; + // ============================================================================= // Base Trait // ============================================================================= @@ -755,6 +791,24 @@ pub trait BaseRealtime: Send + Sync { /// Default is a no-op so providers that don't (yet) consume the handles compile unchanged. fn set_resilience(&mut self, _resilience: crate::core::resilience::ResilienceHandles) {} + /// Every normalized event the provider produces, in wire order, BEFORE the callbacks above + /// see it — including those no callback carries (`Usage`, `ItemAdded`, `InterruptedByServer`, + /// …). The translate engine (FRD-023 §5.7) builds the OpenAI GA event stream from it. The + /// callback must not block: the provider's receive loop awaits it. + /// + /// Default: refused, so a provider that is not on the S2S scaffold cannot be translated by + /// accident with half its events missing. + fn on_event(&mut self, _callback: S2sEventCallback) -> RealtimeResult<()> { + Err(RealtimeError::InvalidConfiguration( + "this realtime provider does not expose its event stream".to_string(), + )) + } + + /// The PCM16 sample rates the provider speaks: `(input, output)`, when known. + fn audio_rates(&self) -> Option<(u32, u32)> { + None + } + // ── B-G2: S2S-as-a-service surface (defaults are no-ops so every // provider compiles; OpenAI Realtime implements them fully) ── diff --git a/gateway/src/core/realtime/deepgram/mod.rs b/gateway/src/core/realtime/deepgram/mod.rs index 05c49a36..dafe6ca3 100644 --- a/gateway/src/core/realtime/deepgram/mod.rs +++ b/gateway/src/core/realtime/deepgram/mod.rs @@ -168,6 +168,17 @@ impl BaseRealtime for DeepgramRealtime { self.0.set_resilience(resilience) } + fn on_event( + &mut self, + callback: crate::core::realtime::base::S2sEventCallback, + ) -> RealtimeResult<()> { + self.0.on_event(callback) + } + + fn audio_rates(&self) -> Option<(u32, u32)> { + self.0.audio_rates() + } + fn emits_user_turn_frames(&self) -> bool { self.0.emits_user_turn_frames() } diff --git a/gateway/src/core/realtime/deepgram/protocol.rs b/gateway/src/core/realtime/deepgram/protocol.rs index 94279218..e413c230 100644 --- a/gateway/src/core/realtime/deepgram/protocol.rs +++ b/gateway/src/core/realtime/deepgram/protocol.rs @@ -145,6 +145,11 @@ impl RealtimeProtocol for DeepgramProtocol { "deepgram" } + fn input_sample_rate(&self) -> u32 { + // One rate both ways (the `Settings` audio block). + self.output_sample_rate + } + fn caps(&self) -> ProtocolCaps { ProtocolCaps { // Deepgram runs server-side VAD + turn-taking, so it owns turns and diff --git a/gateway/src/core/realtime/elevenlabs/mod.rs b/gateway/src/core/realtime/elevenlabs/mod.rs index 2dfb7d47..4ea6e401 100644 --- a/gateway/src/core/realtime/elevenlabs/mod.rs +++ b/gateway/src/core/realtime/elevenlabs/mod.rs @@ -171,6 +171,17 @@ impl BaseRealtime for ElevenLabsRealtime { self.0.set_resilience(resilience) } + fn on_event( + &mut self, + callback: crate::core::realtime::base::S2sEventCallback, + ) -> RealtimeResult<()> { + self.0.on_event(callback) + } + + fn audio_rates(&self) -> Option<(u32, u32)> { + self.0.audio_rates() + } + fn emits_user_turn_frames(&self) -> bool { self.0.emits_user_turn_frames() } diff --git a/gateway/src/core/realtime/elevenlabs/protocol.rs b/gateway/src/core/realtime/elevenlabs/protocol.rs index 0f32976c..37585eb4 100644 --- a/gateway/src/core/realtime/elevenlabs/protocol.rs +++ b/gateway/src/core/realtime/elevenlabs/protocol.rs @@ -90,6 +90,11 @@ impl RealtimeProtocol for ElevenLabsProtocol { "elevenlabs" } + fn input_sample_rate(&self) -> u32 { + // `pcm_16000` both directions. + OUTPUT_SAMPLE_RATE + } + fn caps(&self) -> ProtocolCaps { ProtocolCaps { // ConvAI runs server-side VAD + turn-taking, so it owns turns and diff --git a/gateway/src/core/realtime/gemini/mod.rs b/gateway/src/core/realtime/gemini/mod.rs index f483eb1f..12a67f4c 100644 --- a/gateway/src/core/realtime/gemini/mod.rs +++ b/gateway/src/core/realtime/gemini/mod.rs @@ -177,6 +177,17 @@ impl BaseRealtime for GeminiRealtime { self.0.set_resilience(resilience) } + fn on_event( + &mut self, + callback: crate::core::realtime::base::S2sEventCallback, + ) -> RealtimeResult<()> { + self.0.on_event(callback) + } + + fn audio_rates(&self) -> Option<(u32, u32)> { + self.0.audio_rates() + } + fn emits_user_turn_frames(&self) -> bool { self.0.emits_user_turn_frames() } diff --git a/gateway/src/core/realtime/gemini/protocol.rs b/gateway/src/core/realtime/gemini/protocol.rs index c56e85e7..98dbffa6 100644 --- a/gateway/src/core/realtime/gemini/protocol.rs +++ b/gateway/src/core/realtime/gemini/protocol.rs @@ -42,9 +42,15 @@ use crate::core::realtime::base::{ TranscriptRole, }; use crate::core::realtime::scaffold::{ - ConnectSpec, Inbound, OutFrame, ProtocolCaps, RealtimeProtocol, S2sEvent, + ConnectSpec, Inbound, OutFrame, ProtocolCaps, RealtimeProtocol, S2sEvent, UsageReport, }; +/// A protobuf `Duration` in its JSON form (`"12.5s"`). +fn parse_proto_duration(s: &str) -> Option { + let secs: f64 = s.trim().strip_suffix('s')?.trim().parse().ok()?; + (secs.is_finite() && secs >= 0.0).then(|| std::time::Duration::from_secs_f64(secs)) +} + /// Gemini Live BidiGenerateContent WebSocket endpoint. The `?key=` /// query is appended in `connect_spec` (Gemini auths by QUERY param, not header). pub(crate) const GEMINI_LIVE_URL: &str = "wss://generativelanguage.googleapis.com/ws/google.ai.generativelanguage.v1beta.GenerativeService.BidiGenerateContent"; @@ -240,6 +246,121 @@ impl GeminiProtocol { out } + /// `usageMetadata` → one cumulative [`UsageReport`]. Per modality from the `*TokensDetails` + /// lists (`TEXT`, `AUDIO`, `IMAGE`; `VIDEO` frames bill as image); a vendor that sends only + /// the totals has them attributed to text rather than dropped. `cachedContentTokenCount` is a + /// SUBSET of `promptTokenCount` (Google's contract), which is what the cost formula expects; + /// thinking tokens bill at the output-text rate. + fn map_usage(u: &Value) -> S2sEvent { + fn n(v: Option<&Value>) -> u64 { + v.and_then(Value::as_u64).unwrap_or(0) + } + /// `(text, audio, image)` from a `[{modality, tokenCount}]` list, or `None` if absent. + fn split(list: Option<&Value>) -> Option<(u64, u64, u64)> { + let list = list?.as_array()?; + let mut out = (0, 0, 0); + for d in list { + let count = n(d.get("tokenCount")); + match d.get("modality").and_then(Value::as_str) { + Some("AUDIO") => out.1 += count, + Some("IMAGE") | Some("VIDEO") => out.2 += count, + _ => out.0 += count, + } + } + Some(out) + } + let prompt = + split(u.get("promptTokensDetails")).unwrap_or((n(u.get("promptTokenCount")), 0, 0)); + let tool = split(u.get("toolUsePromptTokensDetails")).unwrap_or(( + n(u.get("toolUsePromptTokenCount")), + 0, + 0, + )); + let cached = split(u.get("cacheTokensDetails")).unwrap_or(( + n(u.get("cachedContentTokenCount")), + 0, + 0, + )); + let response = + split(u.get("responseTokensDetails")).unwrap_or((n(u.get("responseTokenCount")), 0, 0)); + let mut tokens = crate::core::realtime_cost::RealtimeUsage { + input_text: prompt.0 + tool.0, + input_audio: prompt.1 + tool.1, + input_image: prompt.2 + tool.2, + cached_text: cached.0, + cached_audio: cached.1, + cached_image: cached.2, + output_text: response.0 + n(u.get("thoughtsTokenCount")), + output_audio: response.1, + }; + // A cached count above its class is a vendor inconsistency; it must never make the + // uncached remainder negative (that would REDUCE the bill). + tokens.cached_text = tokens.cached_text.min(tokens.input_text); + tokens.cached_audio = tokens.cached_audio.min(tokens.input_audio); + tokens.cached_image = tokens.cached_image.min(tokens.input_image); + S2sEvent::Usage(UsageReport { + tokens, + seconds: None, + cumulative: true, + }) + } + + /// Everything but `usageMetadata`, which rides on any message. + fn map_message(value: &Value) -> Vec { + // goAway: the connection closes after `timeLeft`; the driver reconnects with the + // resumption handle at the next turn boundary (FRD-023 RT7.1, TC-XL-04). + if let Some(ga) = value.get("goAway") { + let time_left = ga + .get("timeLeft") + .and_then(Value::as_str) + .and_then(parse_proto_duration); + return vec![S2sEvent::GoAway { time_left }]; + } + + // serverContent — the MULTI-FRAME case (audio + text parts + turn flags). + if let Some(sc) = value.get("serverContent") { + return Self::map_server_content(sc); + } + + // toolCall ⇒ one FunctionCall per functionCalls[] entry. + if let Some(tc) = value.get("toolCall") { + return Self::map_tool_call(tc); + } + + // sessionResumptionUpdate ⇒ ResumptionHandle (driver stores it, feeds it + // back into build_session_config on reconnect). Only emit when the server + // marks it resumable AND provides a non-empty handle. + if let Some(update) = value.get("sessionResumptionUpdate") { + let resumable = update + .get("resumable") + .and_then(Value::as_bool) + .unwrap_or(false); + if let Some(handle) = update.get("newHandle").and_then(Value::as_str) + && resumable + && !handle.is_empty() + { + return vec![S2sEvent::ResumptionHandle(handle.to_string())]; + } + return vec![S2sEvent::Ignore]; + } + + // toolCallCancellation: the server cancelled pending tool calls; nothing + // for the gateway to forward (the call ids would already be in flight) ⇒ Ignore. + vec![S2sEvent::Ignore] + } + + fn encode_user_audio_chunk(&self, pcm: &[u8]) -> Value { + // base64-in-JSON: realtimeInput.mediaChunks[0] with the 16 kHz input mime. + json!({ + "realtimeInput": { + "mediaChunks": [{ + "mimeType": format!("audio/pcm;rate={INPUT_SAMPLE_RATE}"), + "data": BASE64_STANDARD.encode(pcm), + }] + } + }) + } + /// Lower a `toolCall` object into one `FunctionCall` per `functionCalls[]` /// entry. Each call carries `id` (may be absent on Vertex), `name`, and /// `args` (a JSON object stringified into `arguments`). @@ -302,6 +423,10 @@ impl RealtimeProtocol for GeminiProtocol { "gemini" } + fn input_sample_rate(&self) -> u32 { + INPUT_SAMPLE_RATE + } + fn caps(&self) -> ProtocolCaps { ProtocolCaps { // Gemini runs server-side VAD + turn-taking, so it owns turns and @@ -416,50 +541,18 @@ impl RealtimeProtocol for GeminiProtocol { return vec![S2sEvent::SessionReady { session_id: None }]; } - // serverContent — the MULTI-FRAME case (audio + text parts + turn flags). - if let Some(sc) = value.get("serverContent") { - return Self::map_server_content(sc); - } - - // toolCall ⇒ one FunctionCall per functionCalls[] entry. - if let Some(tc) = value.get("toolCall") { - return Self::map_tool_call(tc); + // usageMetadata rides on any server message (usually the turn's last). It goes FIRST, + // so it lands on the response the same message completes (FRD-023 §5.7, TC-XL-03). + let mut events = Self::map_message(&value); + if let Some(u) = value.get("usageMetadata").map(Self::map_usage) { + events.retain(|e| !matches!(e, S2sEvent::Ignore)); + events.insert(0, u); } - - // sessionResumptionUpdate ⇒ ResumptionHandle (driver stores it, feeds it - // back into build_session_config on reconnect). Only emit when the server - // marks it resumable AND provides a non-empty handle. - if let Some(update) = value.get("sessionResumptionUpdate") { - let resumable = update - .get("resumable") - .and_then(Value::as_bool) - .unwrap_or(false); - if let Some(handle) = update.get("newHandle").and_then(Value::as_str) - && resumable - && !handle.is_empty() - { - return vec![S2sEvent::ResumptionHandle(handle.to_string())]; - } - return vec![S2sEvent::Ignore]; - } - - // toolCallCancellation: the server cancelled pending tool calls; nothing - // for the gateway to forward (the call ids would already be in flight). - // goAway: the server warns of an imminent disconnect — the scaffold's - // reconnect (with the stored resumption handle) handles it. Both ⇒ Ignore. - vec![S2sEvent::Ignore] + events } fn encode_user_audio(&self, pcm: &[u8]) -> Self::Wire { - // base64-in-JSON: realtimeInput.mediaChunks[0] with the 16 kHz input mime. - json!({ - "realtimeInput": { - "mediaChunks": [{ - "mimeType": format!("audio/pcm;rate={INPUT_SAMPLE_RATE}"), - "data": BASE64_STANDARD.encode(pcm), - }] - } - }) + self.encode_user_audio_chunk(pcm) } fn send_text(&self, text: &str) -> Vec { @@ -821,6 +914,67 @@ mod tests { } } + /// TC-XL-03 (protocol half) — `usageMetadata` becomes ONE cumulative `Usage`, FIRST in the + /// message's events (so it lands on the response the same message completes), split by + /// modality; cached tokens stay a subset of their class; thinking bills as output text. + #[test] + fn tc_xl_03_usage_metadata_becomes_one_cumulative_usage_first() { + let p = proto(&base_cfg()); + let raw = json!({ + "serverContent": {"turnComplete": true}, + "usageMetadata": { + "promptTokenCount": 120, "cachedContentTokenCount": 20, + "responseTokenCount": 80, "thoughtsTokenCount": 5, "totalTokenCount": 205, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 30}, + {"modality": "AUDIO", "tokenCount": 90}], + "cacheTokensDetails": [{"modality": "TEXT", "tokenCount": 20}], + "responseTokensDetails": [{"modality": "AUDIO", "tokenCount": 70}, + {"modality": "TEXT", "tokenCount": 10}] + } + }) + .to_string(); + let evs = p.map_server_event(Inbound::Text(&raw)); + let S2sEvent::Usage(u) = &evs[0] else { + panic!("usage first, got {evs:?}"); + }; + assert!( + u.cumulative, + "a Gemini report is the response's running total" + ); + let t = u.tokens; + assert_eq!( + (t.input_text, t.input_audio, t.cached_text, t.cached_audio), + (30, 90, 20, 0) + ); + assert_eq!((t.output_text, t.output_audio), (15, 70)); + assert!(matches!(evs.last(), Some(S2sEvent::ResponseDone { .. }))); + + // Totals only: attributed to text, never dropped. + let raw = json!({"usageMetadata": {"promptTokenCount": 50, "responseTokenCount": 25}}) + .to_string(); + let evs = p.map_server_event(Inbound::Text(&raw)); + let [S2sEvent::Usage(u)] = evs.as_slice() else { + panic!("{evs:?}"); + }; + assert_eq!((u.tokens.input_text, u.tokens.output_text), (50, 25)); + } + + /// TC-XL-04 (protocol half) — `goAway` carries its deadline to the driver. + #[test] + fn tc_xl_04_go_away_carries_time_left() { + let p = proto(&base_cfg()); + let evs = p.map_server_event(Inbound::Text(r#"{"goAway":{"timeLeft":"12.5s"}}"#)); + assert!( + matches!(evs.as_slice(), [S2sEvent::GoAway { time_left: Some(d) }] if *d == std::time::Duration::from_millis(12_500)), + "{evs:?}" + ); + let evs = p.map_server_event(Inbound::Text(r#"{"goAway":{}}"#)); + assert!( + matches!(evs.as_slice(), [S2sEvent::GoAway { time_left: None }]), + "{evs:?}" + ); + } + /// setupComplete ⇒ SessionReady. #[test] fn setup_complete_maps_to_session_ready() { @@ -878,12 +1032,12 @@ mod tests { } } - /// goAway / toolCallCancellation / unknown / non-JSON / binary ⇒ Ignore. + /// toolCallCancellation / unknown / non-JSON / binary ⇒ Ignore. (`goAway` is acted on: + /// `tc_xl_04_go_away_carries_time_left`.) #[test] fn go_away_cancellation_unknown_and_binary_ignore() { let p = proto(&base_cfg()); for raw in [ - r#"{"goAway":{"timeLeft":"5s"}}"#, r#"{"toolCallCancellation":{"ids":["call_1"]}}"#, r#"{"somethingNew":{}}"#, "not json", diff --git a/gateway/src/core/realtime/hume/client.rs b/gateway/src/core/realtime/hume/client.rs index 55722435..22c6a950 100644 --- a/gateway/src/core/realtime/hume/client.rs +++ b/gateway/src/core/realtime/hume/client.rs @@ -222,6 +222,17 @@ impl BaseRealtime for HumeEVI { self.0.set_resilience(resilience) } + fn on_event( + &mut self, + callback: crate::core::realtime::base::S2sEventCallback, + ) -> RealtimeResult<()> { + self.0.on_event(callback) + } + + fn audio_rates(&self) -> Option<(u32, u32)> { + self.0.audio_rates() + } + fn emits_user_turn_frames(&self) -> bool { self.0.emits_user_turn_frames() } diff --git a/gateway/src/core/realtime/hume/protocol.rs b/gateway/src/core/realtime/hume/protocol.rs index 3b62d3d5..103a97ac 100644 --- a/gateway/src/core/realtime/hume/protocol.rs +++ b/gateway/src/core/realtime/hume/protocol.rs @@ -110,6 +110,10 @@ impl RealtimeProtocol for HumeProtocol { "hume" } + fn input_sample_rate(&self) -> u32 { + self.config.sample_rate + } + fn caps(&self) -> ProtocolCaps { ProtocolCaps { // Hume EVI runs SERVER-side VAD + turn-taking (it auto-generates diff --git a/gateway/src/core/realtime/mod.rs b/gateway/src/core/realtime/mod.rs index 19be8254..5f4a885d 100644 --- a/gateway/src/core/realtime/mod.rs +++ b/gateway/src/core/realtime/mod.rs @@ -66,13 +66,13 @@ pub mod yandex; pub use azure::{AzureProtocol, AzureRealtime}; pub use base::{ - AudioOutputCallback, BaseRealtime, BoxedRealtime, ConnectionState, FunctionCallCallback, - FunctionCallRequest, FunctionDefinition, InputTranscriptionConfig, RealtimeAudioData, - RealtimeConfig, RealtimeError, RealtimeErrorCallback, RealtimeFactory, - RealtimeResponseOverride, RealtimeResult, ReconnectionCallback, ReconnectionEvent, - ReplayConversationItem, ResponseDoneCallback, SpeechEvent, SpeechEventCallback, ToolDefinition, - TranscriptCallback, TranscriptResult, TranscriptRole, TurnDetectionConfig, clamp_truncate_ms, - run_barge_in_sequence, + AudioOutputCallback, AwsStaticCredentials, BaseRealtime, BoxedRealtime, ConnectionState, + FunctionCallCallback, FunctionCallRequest, FunctionDefinition, InputTranscriptionConfig, + RealtimeAudioData, RealtimeConfig, RealtimeError, RealtimeErrorCallback, RealtimeFactory, + RealtimeResponseOverride, RealtimeResult, ReconnectionCallback, ReconnectionConfig, + ReconnectionEvent, ReplayConversationItem, ResponseDoneCallback, S2sEventCallback, SpeechEvent, + SpeechEventCallback, ToolDefinition, TranscriptCallback, TranscriptResult, TranscriptRole, + TurnDetectionConfig, clamp_truncate_ms, run_barge_in_sequence, }; pub use deepgram::{DeepgramProtocol, DeepgramRealtime}; pub use elevenlabs::{ElevenLabsProtocol, ElevenLabsRealtime}; diff --git a/gateway/src/core/realtime/nova_sonic/mod.rs b/gateway/src/core/realtime/nova_sonic/mod.rs index 2acb5ccb..48ceebf4 100644 --- a/gateway/src/core/realtime/nova_sonic/mod.rs +++ b/gateway/src/core/realtime/nova_sonic/mod.rs @@ -60,6 +60,24 @@ impl NovaSonicRealtime { self.0.resilience_breaker() } + /// A session signed with a Bud deployment's own AWS keys (FRD-023 RT7.2): it dials through + /// `factory` (built with [`BedrockBidiTransportFactory::with_credentials`]) instead of the + /// `aws-config` default chain. + /// + /// [`BedrockBidiTransportFactory::with_credentials`]: crate::core::realtime::scaffold::BedrockBidiTransportFactory::with_credentials + pub fn with_transport( + config: RealtimeConfig, + factory: crate::core::realtime::scaffold::BedrockBidiTransportFactory, + ) -> RealtimeResult { + use crate::core::realtime::scaffold::RealtimeProtocol; + let protocol = NovaSonicProtocol::from_config(&config)?; + Ok(Self(RealtimeSession::with_transport( + protocol, + config, + std::sync::Arc::new(factory), + )?)) + } + /// Get the session ID if connected (Nova Sonic surfaces it on `completionStart` /// / `completionEnd`; the driver tracks none over this seam ⇒ `None`). pub async fn session_id(&self) -> Option { @@ -181,6 +199,17 @@ impl BaseRealtime for NovaSonicRealtime { self.0.set_resilience(resilience) } + fn on_event( + &mut self, + callback: crate::core::realtime::base::S2sEventCallback, + ) -> RealtimeResult<()> { + self.0.on_event(callback) + } + + fn audio_rates(&self) -> Option<(u32, u32)> { + self.0.audio_rates() + } + fn emits_user_turn_frames(&self) -> bool { self.0.emits_user_turn_frames() } diff --git a/gateway/src/core/realtime/nova_sonic/protocol.rs b/gateway/src/core/realtime/nova_sonic/protocol.rs index 3f6abaaa..7c76ce82 100644 --- a/gateway/src/core/realtime/nova_sonic/protocol.rs +++ b/gateway/src/core/realtime/nova_sonic/protocol.rs @@ -45,13 +45,14 @@ use bytes::Bytes; use serde_json::{Value, json}; use std::sync::Mutex; +use crate::core::realtime::base::ReplayConversationItem; use crate::core::realtime::base::{ FunctionCallRequest, RealtimeConfig, RealtimeError, RealtimeResponseOverride, RealtimeResult, TranscriptRole, }; use crate::core::realtime::scaffold::{ BedrockBidiTransportFactory, ConnectSpec, Inbound, OutFrame, ProtocolCaps, RealtimeProtocol, - RealtimeTransportFactory, S2sEvent, + RealtimeTransportFactory, S2sEvent, UsageReport, }; use std::sync::Arc; @@ -76,6 +77,9 @@ const OUTPUT_BYTES_PER_MS: u64 = 48; /// 8/16/24 kHz). const INPUT_SAMPLE_RATE: u32 = 16_000; +/// A Bedrock bidirectional stream lives at most 8 minutes; replace it 30 s before that. +const MAX_CONNECTION: std::time::Duration = std::time::Duration::from_secs(8 * 60 - 30); + /// The current text-content role/stage tracked across a `contentStart`→`textOutput` /// pair. Nova Sonic puts the role (USER ASR vs ASSISTANT) + `generationStage` /// (FINAL vs SPECULATIVE) on the `contentStart`, while the following `textOutput` @@ -114,6 +118,11 @@ pub struct NovaSonicProtocol { /// on the following `textOutput`. Single-writer (the driver's recv loop); /// `Mutex` only to satisfy `Send + Sync` on the `&self` trait. current_text: Mutex>, + /// The assistant AUDIO content block in flight (its `contentId`), for `ItemAdded`/`ItemDone`. + current_audio: Mutex>, + /// The running `usageEvent` totals on THIS connection, for a report that carries totals + /// but no delta. Reset at every (re)connect: a new stream counts from zero. + usage_totals: Mutex, } impl NovaSonicProtocol { @@ -197,6 +206,58 @@ impl NovaSonicProtocol { Some(json!({ "tools": specs })) } + /// One `usageEvent` → one delta [`UsageReport`] (see the `usageEvent` arm). + fn map_usage(&self, body: &Value) -> Vec { + fn n(v: Option<&Value>) -> u64 { + v.and_then(Value::as_u64).unwrap_or(0) + } + fn read(block: &Value) -> crate::core::realtime_cost::RealtimeUsage { + let input = block.get("input"); + let output = block.get("output"); + crate::core::realtime_cost::RealtimeUsage { + input_audio: n(input.and_then(|i| i.get("speechTokens"))), + input_text: n(input.and_then(|i| i.get("textTokens"))), + output_audio: n(output.and_then(|o| o.get("speechTokens"))), + output_text: n(output.and_then(|o| o.get("textTokens"))), + ..Default::default() + } + } + let details = body.get("details"); + let tokens = if let Some(delta) = details.and_then(|d| d.get("delta")) { + let delta = read(delta); + if let Some(total) = details.and_then(|d| d.get("total")) + && let Ok(mut t) = self.usage_totals.lock() + { + *t = read(total); + } + delta + } else if let Some(total) = details.and_then(|d| d.get("total")) { + let total = read(total); + let Ok(mut last) = self.usage_totals.lock() else { + return vec![S2sEvent::Ignore]; + }; + let delta = crate::core::realtime_cost::RealtimeUsage { + input_audio: total.input_audio.saturating_sub(last.input_audio), + input_text: total.input_text.saturating_sub(last.input_text), + output_audio: total.output_audio.saturating_sub(last.output_audio), + output_text: total.output_text.saturating_sub(last.output_text), + ..Default::default() + }; + *last = total; + delta + } else { + return vec![S2sEvent::Ignore]; + }; + if tokens.is_empty() { + return vec![S2sEvent::Ignore]; + } + vec![S2sEvent::Usage(UsageReport { + tokens, + seconds: None, + cumulative: false, + })] + } + /// `{"event": {"audioInput": {promptName, contentName, content}}}` — one /// base64-PCM user-audio chunk into the single audio content container. fn audio_input_event(&self, pcm: &[u8]) -> Value { @@ -240,6 +301,8 @@ impl RealtimeProtocol for NovaSonicProtocol { prompt_name: uuid::Uuid::new_v4().to_string(), audio_content_name: uuid::Uuid::new_v4().to_string(), current_text: Mutex::new(None), + current_audio: Mutex::new(None), + usage_totals: Mutex::new(Default::default()), }) } @@ -247,13 +310,23 @@ impl RealtimeProtocol for NovaSonicProtocol { /// stream, so it opts into the Bedrock-bidi transport factory instead of the /// default plain WebSocket. fn transport_factory(&self) -> Arc { - Arc::new(BedrockBidiTransportFactory) + Arc::new(BedrockBidiTransportFactory::default_chain()) } fn provider_id(&self) -> &'static str { "nova_sonic" } + fn input_sample_rate(&self) -> u32 { + INPUT_SAMPLE_RATE + } + + /// Nova Sonic ends a bidirectional stream at 8 minutes; the driver replaces it 30 s before, + /// at a turn boundary when there is one (FRD-023 RT7.2). + fn max_connection(&self) -> Option { + Some(MAX_CONNECTION) + } + fn caps(&self) -> ProtocolCaps { ProtocolCaps { // Nova Sonic runs server-side VAD + turn-taking + barge-in, so it owns @@ -284,6 +357,13 @@ impl RealtimeProtocol for NovaSonicProtocol { cfg: &RealtimeConfig, _resumption: Option<&str>, ) -> Vec { + // A (re)connect is a new stream: its usage totals count from zero. + if let Ok(mut t) = self.usage_totals.lock() { + *t = Default::default(); + } + if let Ok(mut a) = self.current_audio.lock() { + *a = None; + } // The Nova Sonic bootstrap sequence (each a `{"event": {…}}` JSON), in the // documented order, leaving the session ready for `audioInput`: // sessionStart → promptStart → (SYSTEM contentStart/textInput/contentEnd) @@ -403,6 +483,19 @@ impl RealtimeProtocol for NovaSonicProtocol { // textOutput (and the type for audio). We stash the TEXT role/stage so // the subsequent textOutput maps to the right transcript role/finality. "contentStart" => { + // The assistant's spoken output is a conversation item of its own. + if body.get("type").and_then(Value::as_str) == Some("AUDIO") + && body.get("role").and_then(Value::as_str) == Some("ASSISTANT") + && let Some(id) = body.get("contentId").and_then(Value::as_str) + { + if let Ok(mut g) = self.current_audio.lock() { + *g = Some(id.to_string()); + } + return vec![S2sEvent::ItemAdded { + item_id: id.to_string(), + role: TranscriptRole::Assistant, + }]; + } if body.get("type").and_then(Value::as_str) == Some("TEXT") { let role = match body.get("role").and_then(Value::as_str) { Some("USER") => TranscriptRole::User, @@ -506,11 +599,24 @@ impl RealtimeProtocol for NovaSonicProtocol { // (the driver clears local playback). Other stopReasons are not turn- // complete (completionEnd is the response-done signal), so Ignore. "contentEnd" => { + let mut out = Vec::new(); + let id = body.get("contentId").and_then(Value::as_str); + if let Ok(mut g) = self.current_audio.lock() + && id.is_some() + && g.as_deref() == id + { + *g = None; + out.push(S2sEvent::ItemDone { + item_id: id.unwrap_or_default().to_string(), + }); + } if body.get("stopReason").and_then(Value::as_str) == Some("INTERRUPTED") { - vec![S2sEvent::InterruptedByServer] - } else { - vec![S2sEvent::Ignore] + out.push(S2sEvent::InterruptedByServer); + } + if out.is_empty() { + out.push(S2sEvent::Ignore); } + out } // completionEnd: the model finished this response generation. Drives // on_response_done (the response/completion id is not separately needed). @@ -521,7 +627,12 @@ impl RealtimeProtocol for NovaSonicProtocol { .unwrap_or("") .to_string(), }], - // completionStart / usageEvent / unknown ⇒ nothing actionable. + // usageEvent: the tokens since the last report (`details.delta`), metered once each + // (FRD-023 §5.10, TC-XL-05). Speech bills at the audio rates, text at the text + // rates. The running `details.total` is used only when a report has no delta — and + // then as a difference, so it is never billed twice. + "usageEvent" => self.map_usage(body), + // completionStart / unknown ⇒ nothing actionable. _ => vec![S2sEvent::Ignore], } } @@ -586,6 +697,48 @@ impl RealtimeProtocol for NovaSonicProtocol { Vec::new() } + /// Conversation history after a reconnect (the 8-minute cap, FRD-023 RT7.2): a new stream + /// starts a new Nova session, so each logged turn is sent back as a non-interactive TEXT + /// block with its role — the context survives the connection. + fn replay_item(&self, item: &ReplayConversationItem) -> Vec { + let content_name = uuid::Uuid::new_v4().to_string(); + let role = match item.role { + TranscriptRole::User => "USER", + TranscriptRole::Assistant => "ASSISTANT", + }; + vec![ + json!({ + "event": { + "contentStart": { + "promptName": self.prompt_name, + "contentName": content_name, + "type": "TEXT", + "interactive": false, + "role": role, + "textInputConfiguration": { "mediaType": "text/plain" }, + } + } + }), + json!({ + "event": { + "textInput": { + "promptName": self.prompt_name, + "contentName": content_name, + "content": item.text, + } + } + }), + json!({ + "event": { + "contentEnd": { + "promptName": self.prompt_name, + "contentName": content_name, + } + } + }), + ] + } + fn format_tool_result(&self, call_id: &str, result: &str) -> Vec { // A tool result is a TOOL content block: contentStart(TOOL, role TOOL, // toolResultInputConfiguration referencing the toolUseId) / toolResult / @@ -673,6 +826,90 @@ mod tests { } /// from_config defaults the model + voice when omitted. + /// TC-XL-05 (protocol half) — a `usageEvent` is ONE delta report: speech → audio, text → + /// text, from `details.delta` (never the running total, which would bill twice). + #[test] + fn tc_xl_05_usage_event_is_one_delta_report() { + let p = NovaSonicProtocol::from_config(&base_cfg()).unwrap(); + let raw = json!({"event": {"usageEvent": { + "completionId": "c1", "totalTokens": 999, + "details": { + "delta": {"input": {"speechTokens": 40, "textTokens": 3}, + "output": {"speechTokens": 50, "textTokens": 7}}, + "total": {"input": {"speechTokens": 400, "textTokens": 30}, + "output": {"speechTokens": 500, "textTokens": 70}} + } + }}}) + .to_string(); + let evs = p.map_server_event(Inbound::Text(&raw)); + let [S2sEvent::Usage(u)] = evs.as_slice() else { + panic!("{evs:?}"); + }; + assert!(!u.cumulative); + let t = u.tokens; + assert_eq!( + (t.input_audio, t.input_text, t.output_audio, t.output_text), + (40, 3, 50, 7) + ); + } + + /// A report with totals and no delta is billed as the DIFFERENCE from the last total on + /// the same connection; a reconnect starts the count again. + #[test] + fn usage_totals_without_a_delta_bill_the_difference_per_connection() { + let p = NovaSonicProtocol::from_config(&base_cfg()).unwrap(); + let total = |speech_in: u64| { + json!({"event": {"usageEvent": {"details": {"total": { + "input": {"speechTokens": speech_in, "textTokens": 0}, + "output": {"speechTokens": 0, "textTokens": 0}}}}}}) + .to_string() + }; + let billed = |raw: String| match p.map_server_event(Inbound::Text(&raw)).as_slice() { + [S2sEvent::Usage(u)] => u.tokens.input_audio, + _ => 0, + }; + assert_eq!(billed(total(100)), 100); + assert_eq!(billed(total(150)), 50); + let _ = p.build_session_config(&base_cfg(), None); // a reconnect + assert_eq!(billed(total(30)), 30); + } + + /// The assistant's audio block is an item: `ItemAdded` at its start, `ItemDone` at its end. + #[test] + fn assistant_audio_blocks_are_items() { + let p = NovaSonicProtocol::from_config(&base_cfg()).unwrap(); + let start = json!({"event": {"contentStart": {"type": "AUDIO", "role": "ASSISTANT", + "contentId": "ct-1"}}}) + .to_string(); + assert!(matches!( + p.map_server_event(Inbound::Text(&start)).as_slice(), + [S2sEvent::ItemAdded { item_id, role: TranscriptRole::Assistant }] if item_id == "ct-1" + )); + let end = json!({"event": {"contentEnd": {"contentId": "ct-1", "stopReason": "END_TURN"}}}) + .to_string(); + assert!(matches!( + p.map_server_event(Inbound::Text(&end)).as_slice(), + [S2sEvent::ItemDone { item_id }] if item_id == "ct-1" + )); + } + + /// History survives a reconnect: each logged turn goes back as a non-interactive TEXT + /// block with its role. + #[test] + fn replay_sends_history_as_text_blocks() { + let p = NovaSonicProtocol::from_config(&base_cfg()).unwrap(); + let wires = p.replay_item(&ReplayConversationItem { + role: TranscriptRole::Assistant, + text: "hi there".into(), + }); + assert_eq!(wires.len(), 3); + let cs = &wires[0]["event"]["contentStart"]; + assert_eq!(cs["role"], "ASSISTANT"); + assert_eq!(cs["type"], "TEXT"); + assert_eq!(cs["interactive"], false); + assert_eq!(wires[1]["event"]["textInput"]["content"], "hi there"); + } + /// F-4 — Nova Sonic v1 reached end of life on 2026-09-14; the default is Nova 2 Sonic. #[test] fn f4_the_default_model_is_nova_2_sonic() { @@ -1095,7 +1332,8 @@ mod tests { } /// Turn/response controls are empty (server VAD owns them); defaults (truncate, - /// input-buffer clear, replay) are empty too. + /// input-buffer clear) are empty too. (Replay is not: the 8-minute reconnect needs the + /// history — `replay_sends_history_as_text_blocks`.) #[test] fn turn_and_default_controls_are_empty() { let p = proto(&base_cfg()); @@ -1104,12 +1342,5 @@ mod tests { assert!(p.cancel_response().is_empty()); assert!(p.truncate("x", 1).is_empty()); assert!(p.clear_input_buffer().is_empty()); - assert!( - p.replay_item(&ReplayConversationItem { - role: TranscriptRole::User, - text: "x".into(), - }) - .is_empty() - ); } } diff --git a/gateway/src/core/realtime/scaffold/event.rs b/gateway/src/core/realtime/scaffold/event.rs index c1d36ca1..25a1692d 100644 --- a/gateway/src/core/realtime/scaffold/event.rs +++ b/gateway/src/core/realtime/scaffold/event.rs @@ -260,6 +260,19 @@ mod endpoint_override_tests { } } +/// One vendor usage report ([`S2sEvent::Usage`]). +#[derive(Debug, Clone, Copy, PartialEq, Default)] +pub struct UsageReport { + /// Per-modality token counts; cached counts are a SUBSET of their input class. + pub tokens: crate::core::realtime_cost::RealtimeUsage, + /// Billed seconds, for a vendor that reports duration instead of tokens. + pub seconds: Option, + /// `true`: a running total for the response in progress — a later report for the same + /// response REPLACES it (Gemini's `usageMetadata`). `false`: a delta — reports ADD (Nova + /// Sonic's `usageEvent.details.delta`). + pub cumulative: bool, +} + /// The normalized realtime server event. Every provider's `map_server_event` /// lowers raw wire payloads to a `Vec` of these (a Vec because some providers — /// e.g. Gemini — bundle several logical frames into one wire message). The @@ -299,6 +312,23 @@ pub enum S2sEvent { InterruptedByServer, /// A resumption handle to carry across reconnects (Gemini session-resumption). ResumptionHandle(String), + /// The vendor's usage report (FRD-023 §5.7, §5.10): per-modality token counts, or billed + /// seconds. Gemini `usageMetadata`, Nova Sonic `usageEvent`. Each report is metered exactly + /// once by whoever consumes it; the driver itself only forwards it. + Usage(UsageReport), + /// A conversation item began (the vendor's own item id, where it has one). + ItemAdded { + item_id: String, + role: TranscriptRole, + }, + /// A conversation item is complete. + ItemDone { item_id: String }, + /// The vendor will close this connection soon (Gemini `goAway`). The driver reconnects at + /// the next turn boundary — or when `time_left` runs out — carrying the resumption handle, + /// with no backoff and without counting it as a failure. + GoAway { + time_left: Option, + }, /// An outbound frame the driver must send in response to an inbound event /// (e.g. ElevenLabs ConvAI ping→pong); the dispatcher routes it to the /// transport instead of a callback. diff --git a/gateway/src/core/realtime/scaffold/mock.rs b/gateway/src/core/realtime/scaffold/mock.rs index 3a24e0b2..57ca5eb4 100644 --- a/gateway/src/core/realtime/scaffold/mock.rs +++ b/gateway/src/core/realtime/scaffold/mock.rs @@ -106,6 +106,7 @@ pub struct MockProtocol { caps: ProtocolCaps, scripts: Arc>>>, obs: MockObserver, + max_connection: Option, } impl MockProtocol { @@ -116,11 +117,18 @@ impl MockProtocol { caps, scripts: Arc::new(Mutex::new(scripts.into())), obs: obs.clone(), + max_connection: None, }, obs, ) } + /// A connection cap, as Nova Sonic has. + pub fn with_max_connection(mut self, d: std::time::Duration) -> Self { + self.max_connection = Some(d); + self + } + fn map_one(cmd: &str) -> S2sEvent { let parts: Vec<&str> = cmd.splitn(4, ':').collect(); match parts.as_slice() { @@ -169,6 +177,17 @@ impl MockProtocol { // `inbound_command_triggers_outbound_send_frame` test below. ["send", text] => S2sEvent::SendFrame(OutFrame::Text(text.to_string())), ["resume", h] => S2sEvent::ResumptionHandle(h.to_string()), + ["goaway", ms] => S2sEvent::GoAway { + time_left: ms.parse().ok().map(std::time::Duration::from_millis), + }, + ["usage", n] => S2sEvent::Usage(super::event::UsageReport { + tokens: crate::core::realtime_cost::RealtimeUsage { + input_audio: n.parse().unwrap_or(0), + ..Default::default() + }, + seconds: None, + cumulative: false, + }), ["err", msg] => S2sEvent::Error(RealtimeError::ProviderError(msg.to_string())), _ => S2sEvent::Ignore, } @@ -195,6 +214,9 @@ impl RealtimeProtocol for MockProtocol { fn caps(&self) -> ProtocolCaps { self.caps } + fn max_connection(&self) -> Option { + self.max_connection + } fn connect_spec(&self, _cfg: &RealtimeConfig) -> RealtimeResult { Ok(ConnectSpec::WebSocket { url: "wss://mock/realtime".into(), @@ -443,6 +465,114 @@ mod tests { s.disconnect().await.unwrap(); } + /// FRD-023 RT7.0 — the event tap sees every normalized event in wire order, including the + /// ones no native callback carries (`Usage`). + #[tokio::test] + async fn the_event_tap_sees_every_event_in_order() { + let script = vec![Step::Text("multi:usage:7;;audio:10;;done:r1".into())]; + let (proto, _o) = MockProtocol::new(ProtocolCaps::default(), vec![script]); + let mut s = RealtimeSession::from_parts(proto, cfg()).unwrap(); + let seen: Arc>> = Arc::new(Mutex::new(Vec::new())); + let c = seen.clone(); + s.on_event(Arc::new(move |ev| { + let label = match ev { + S2sEvent::Usage(u) => format!("usage:{}", u.tokens.input_audio), + S2sEvent::Audio { data, .. } => format!("audio:{}", data.len()), + S2sEvent::ResponseDone { response_id } => format!("done:{response_id}"), + other => format!("{other:?}"), + }; + c.lock().unwrap().push(label); + Box::pin(async {}) + })) + .unwrap(); + s.connect().await.unwrap(); + until(|| seen.lock().unwrap().len() >= 3).await; + assert_eq!( + seen.lock().unwrap().as_slice(), + &["usage:7", "audio:10", "done:r1"] + ); + s.disconnect().await.unwrap(); + } + + /// TC-XL-04 (driver half) — a `goAway` replaces the connection AT ONCE (no backoff: the + /// default first delay is ~1 s) carrying the resumption handle, and is not a failure. + #[tokio::test] + async fn go_away_reconnects_at_once_with_the_resumption_handle() { + let scripts = vec![ + vec![ + Step::Text("resume:h1".into()), + Step::Text("goaway:0".into()), + ], + vec![], + ]; + let (proto, obs) = MockProtocol::new(ProtocolCaps::default(), scripts); + let mut s = RealtimeSession::from_parts(proto, cfg()).unwrap(); + let reconnects = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let r = reconnects.clone(); + s.on_reconnection(Arc::new(move |ev| { + assert!(ev.success); + r.fetch_add(1, Ordering::SeqCst); + Box::pin(async {}) + })) + .unwrap(); + s.connect().await.unwrap(); + until(|| obs.connect_count() == 2).await; + until(|| { + obs.sent_text() + .contains(&"session_update:resume=h1".to_string()) + }) + .await; + until(|| reconnects.load(Ordering::SeqCst) == 1).await; + assert!(s.is_ready()); + s.disconnect().await.unwrap(); + } + + /// A `goAway` during a response waits for the response to finish (its deadline is far). + #[tokio::test] + async fn go_away_waits_for_the_response_in_flight() { + let scripts = vec![ + vec![ + Step::Text("audio:10".into()), + Step::Text("goaway:60000".into()), + ], + vec![], + ]; + let (proto, obs) = MockProtocol::new(ProtocolCaps::default(), scripts); + let mut s = RealtimeSession::from_parts(proto, cfg()).unwrap(); + s.connect().await.unwrap(); + tokio::time::sleep(Duration::from_millis(150)).await; + assert_eq!(obs.connect_count(), 1, "cut the response off"); + s.disconnect().await.unwrap(); + + let scripts = vec![ + vec![ + Step::Text("audio:10".into()), + Step::Text("goaway:60000".into()), + Step::Text("done:r1".into()), + ], + vec![], + ]; + let (proto, obs) = MockProtocol::new(ProtocolCaps::default(), scripts); + let mut s = RealtimeSession::from_parts(proto, cfg()).unwrap(); + s.connect().await.unwrap(); + until(|| obs.connect_count() == 2).await; + s.disconnect().await.unwrap(); + } + + /// TC-XL-05 (driver half) — a connection cap replaces the connection before the vendor + /// does, over and over, without ever counting as a quick failure (3 of which stop the + /// session). + #[tokio::test] + async fn a_connection_cap_reconnects_proactively_and_is_not_a_failure() { + let (proto, obs) = MockProtocol::new(ProtocolCaps::default(), vec![]); + let proto = proto.with_max_connection(Duration::from_millis(30)); + let mut s = RealtimeSession::from_parts(proto, cfg()).unwrap(); + s.connect().await.unwrap(); + until(|| obs.connect_count() >= 5).await; + assert_ne!(s.get_connection_state(), ConnectionState::Failed); + s.disconnect().await.unwrap(); + } + #[tokio::test] async fn quick_failure_cutoff_stops_the_storm() { // Every connect immediately errors → quick failure. After 3, give up. diff --git a/gateway/src/core/realtime/scaffold/mod.rs b/gateway/src/core/realtime/scaffold/mod.rs index 6715ff49..97e832f9 100644 --- a/gateway/src/core/realtime/scaffold/mod.rs +++ b/gateway/src/core/realtime/scaffold/mod.rs @@ -18,7 +18,9 @@ mod protocol; mod session; mod transport; -pub use event::{ConnectSpec, Inbound, OutFrame, ProtocolCaps, S2sEvent, apply_endpoint_override}; +pub use event::{ + ConnectSpec, Inbound, OutFrame, ProtocolCaps, S2sEvent, UsageReport, apply_endpoint_override, +}; pub use protocol::RealtimeProtocol; pub use session::RealtimeSession; pub use transport::{ diff --git a/gateway/src/core/realtime/scaffold/protocol.rs b/gateway/src/core/realtime/scaffold/protocol.rs index 56f023c6..5738478b 100644 --- a/gateway/src/core/realtime/scaffold/protocol.rs +++ b/gateway/src/core/realtime/scaffold/protocol.rs @@ -42,6 +42,19 @@ pub trait RealtimeProtocol: Send + Sync + Sized + 'static { /// support). Drives the driver's mode-aware barge-in + audio stamping. fn caps(&self) -> ProtocolCaps; + /// The PCM16 sample rate the vendor expects on INPUT (FRD-023 §5.7: a GA client sends + /// 24 kHz; the translate engine resamples to this). Default: the GA rate. + fn input_sample_rate(&self) -> u32 { + 24_000 + } + + /// The longest a single vendor connection may live (Nova Sonic: 8 minutes). The driver + /// reconnects proactively before it — at a turn boundary where it can — so the session + /// outlives the connection. `None`: no cap. + fn max_connection(&self) -> Option { + None + } + /// Build the connect target (URL + headers) from config. Sync: any async /// handshake (Ultravox REST-create, AWS SigV4) is the transport factory's job. fn connect_spec(&self, cfg: &RealtimeConfig) -> RealtimeResult; diff --git a/gateway/src/core/realtime/scaffold/session.rs b/gateway/src/core/realtime/scaffold/session.rs index 2af87934..c2a7f84f 100644 --- a/gateway/src/core/realtime/scaffold/session.rs +++ b/gateway/src/core/realtime/scaffold/session.rs @@ -20,8 +20,8 @@ use super::super::base::{ AudioOutputCallback, BaseRealtime, ConnectionState, FunctionCallCallback, RealtimeAudioData, RealtimeConfig, RealtimeError, RealtimeErrorCallback, RealtimeResponseOverride, RealtimeResult, ReconnectionCallback, ReconnectionConfig, ReconnectionEvent, ReplayConversationItem, - ResponseDoneCallback, SpeechEventCallback, TranscriptCallback, TranscriptResult, - TranscriptRole, clamp_truncate_ms, + ResponseDoneCallback, S2sEventCallback, SpeechEventCallback, TranscriptCallback, + TranscriptResult, TranscriptRole, clamp_truncate_ms, }; use super::event::{OutFrame, ProtocolCaps, S2sEvent}; use super::protocol::RealtimeProtocol; @@ -41,6 +41,20 @@ const PREROLL_CAP_BYTES: usize = 32_000; /// stops the storm instead of hammering. const MIN_STABLE_CONNECTION: Duration = Duration::from_secs(5); const MAX_CONSECUTIVE_QUICK_FAILURES: u32 = 3; +/// A planned reconnect (a vendor's `goAway`, a connection cap) waits for the response in +/// flight to finish — but only until this much before the vendor's own deadline. +const GO_AWAY_MARGIN: Duration = Duration::from_millis(500); + +/// Why a live connection is being replaced on purpose (not a failure, no backoff). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Planned { + /// Not planned: keep the connection. + No, + /// Replace it at the next turn boundary, or at `deadline` at the latest. + WhenIdle, + /// Replace it now. + Now, +} /// Playback position of the currently-streaming assistant item (for truncate). /// `first_delta` anchors WALL-CLOCK elapsed so a barge-in truncates to what the @@ -92,6 +106,8 @@ struct DispatchCtx { playback: Arc>>, conversation_log: Arc>>, resumption: Arc>>, + /// The raw event tap (`BaseRealtime::on_event`). + event_cb: Arc>>, caps: ProtocolCaps, } @@ -221,6 +237,12 @@ impl DispatchCtx { S2sEvent::ResumptionHandle(h) => { *self.resumption.write().await = Some(h); } + // Carried by the event tap only: no callback of the native surface takes them. A + // `GoAway` is acted on by the supervisor before dispatch. + S2sEvent::Usage(_) + | S2sEvent::ItemAdded { .. } + | S2sEvent::ItemDone { .. } + | S2sEvent::GoAway { .. } => {} // Routed to the outbound path by the supervisor BEFORE dispatch (it // owns the out channel); it is not a user-facing callback event, so // reaching the dispatcher is a no-op (defensive — the supervisor's @@ -260,8 +282,18 @@ impl RealtimeSession

{ /// Build a session from a protocol + config (the provider newtype calls this /// from `BaseRealtime::new`). pub fn from_parts(protocol: P, config: RealtimeConfig) -> RealtimeResult { - let caps = protocol.caps(); let factory = protocol.transport_factory(); + Self::with_transport(protocol, config, factory) + } + + /// Build a session that dials through `factory` instead of the protocol's own — e.g. a + /// Bedrock factory holding a deployment's static AWS keys (FRD-023 RT7.2). + pub fn with_transport( + protocol: P, + config: RealtimeConfig, + factory: Arc, + ) -> RealtimeResult { + let caps = protocol.caps(); let reconnect_cfg = config.reconnection.clone().unwrap_or_default(); Ok(Self { protocol: Arc::new(protocol), @@ -289,6 +321,7 @@ impl RealtimeSession

{ playback: Arc::new(StdMutex::new(None)), conversation_log: Arc::new(RwLock::new(Vec::new())), resumption: Arc::new(RwLock::new(None)), + event_cb: Arc::new(Mutex::new(None)), caps, }, reconnection_cb: Arc::new(Mutex::new(None)), @@ -493,18 +526,49 @@ impl RealtimeSession

{ if let Some(r) = &resilience { r.breaker.record_success(); } - Self::notify_reconnect(&reconnection_cb, attempt, true, None).await; } crate::core::metrics::bridge::record_reconnect(provider_id, "connected"); + // Connected BEFORE the reconnect is announced: a listener that resumes sending on + // the announcement must find the session ready. connected.store(true, Ordering::SeqCst); write_connection_state(&state, ConnectionState::Connected); + if is_reconnect { + Self::notify_reconnect(&reconnection_cb, attempt, true, None).await; + } // First connect succeeded → unblock connect(). Self::signal_ready(&mut ready_tx, Ok(())); attempt = 0; + // A connection cap (Nova Sonic: 8 min) or a vendor `goAway` replaces the connection + // on purpose, at a turn boundary when there is one (FRD-023 RT7.1, RT7.2). + let cap = config.max_connection.or_else(|| protocol.max_connection()); + let cap_at = cap.map(|d| tokio::time::Instant::now() + d); + // After the cap, a response in flight gets this long to finish. + let cap_grace = cap.map(|d| (d / 16).max(Duration::from_millis(1))); + let mut planned = Planned::No; + let mut deadline: Option = None; + let mut in_response = false; + // ── inner loop: pump outbound + dispatch inbound ── loop { + if planned == Planned::WhenIdle && !in_response { + planned = Planned::Now; + } + if planned == Planned::Now { + break; + } tokio::select! { + _ = tokio::time::sleep_until(cap_at.unwrap_or_else(tokio::time::Instant::now)), + if cap_at.is_some() && planned == Planned::No => { + planned = Planned::WhenIdle; + let grace_end = cap_at.unwrap_or_else(tokio::time::Instant::now) + + cap_grace.unwrap_or_default(); + deadline = Some(deadline.map_or(grace_end, |d| d.min(grace_end))); + } + _ = tokio::time::sleep_until(deadline.unwrap_or_else(tokio::time::Instant::now)), + if deadline.is_some() => { + planned = Planned::Now; + } out = out_rx.recv() => match out { Some(frame) => { if transport.send(frame).await.is_err() { @@ -519,7 +583,33 @@ impl RealtimeSession

{ }, inbound = transport.recv() => match inbound { Some(Ok(frame)) => { + let tap = ctx.event_cb.lock().await.clone(); for ev in protocol.map_server_event(frame.as_inbound()) { + if let Some(tap) = &tap + && !matches!(ev, S2sEvent::SendFrame(_) | S2sEvent::Ignore) + { + tap(ev.clone()).await; + } + match &ev { + S2sEvent::Audio { .. } + | S2sEvent::FunctionCall(_) + | S2sEvent::Transcript { + role: TranscriptRole::Assistant, + .. + } => in_response = true, + S2sEvent::ResponseDone { .. } + | S2sEvent::InterruptedByServer => in_response = false, + S2sEvent::GoAway { time_left } => { + let by = tokio::time::Instant::now() + + time_left + .unwrap_or_default() + .saturating_sub(GO_AWAY_MARGIN); + deadline = Some(deadline.map_or(by, |d| d.min(by))); + planned = Planned::WhenIdle; + continue; + } + _ => {} + } // An inbound-triggered outbound frame (e.g. ConvAI // ping→pong): send it on the transport DIRECTLY, exactly // like the InterruptedByServer cancel/truncate sends just @@ -581,6 +671,16 @@ impl RealtimeSession

{ connected.store(false, Ordering::SeqCst); transport.close().await; + if planned == Planned::Now && !intentional_disconnect.load(Ordering::SeqCst) { + // Replaced on purpose: no failure signal, no backoff, the resumption handle + // (if any) is carried by the next dial. + tracing::info!(provider = provider_id, "planned reconnect"); + crate::core::metrics::bridge::record_reconnect(provider_id, "planned"); + write_connection_state(&state, ConnectionState::Reconnecting); + is_reconnect = true; + continue 'outer; + } + // D-G2: feed the SHARED, per-provider circuit breaker this // connection's measured lifetime, so a bad-credential // "handshake-OK-then-server-immediately-drops" signature fast-trips @@ -881,6 +981,20 @@ impl BaseRealtime for RealtimeSession

{ self.resilience = Some(resilience); } + fn on_event(&mut self, callback: S2sEventCallback) -> RealtimeResult<()> { + if let Ok(mut g) = self.cb.event_cb.try_lock() { + *g = Some(callback); + } + Ok(()) + } + + fn audio_rates(&self) -> Option<(u32, u32)> { + Some(( + self.protocol.input_sample_rate(), + self.caps.output_sample_rate, + )) + } + fn emits_user_turn_frames(&self) -> bool { self.caps.emits_user_turn_frames } diff --git a/gateway/src/core/realtime/scaffold/transport.rs b/gateway/src/core/realtime/scaffold/transport.rs index 6dfabd69..bfb03638 100644 --- a/gateway/src/core/realtime/scaffold/transport.rs +++ b/gateway/src/core/realtime/scaffold/transport.rs @@ -552,18 +552,41 @@ fn unwrap_output_payload(event: &BidiOutputEvent) -> Option { /// drained by [`recv`](RealtimeTransport::recv). Dropping the transport drops /// `input_tx` → the input stream ends → the HTTP/2 request finalizes (the same /// channel-close finalize `AwsTranscribeTransport` relies on). +type BedrockOutput = aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver< + BidiOutputEvent, + aws_sdk_bedrockruntime::types::error::InvokeModelWithBidirectionalStreamOutputError, +>; + +/// The output half of a Bedrock stream: still opening, or open. +/// +/// The SDK's `send()` for `InvokeModelWithBidirectionalStream` does not return at the response +/// headers — it waits for the stream's FIRST output event (`try_recv_initial_response`). Nova +/// Sonic sends nothing until it has input, and the input (the session configuration, then the +/// caller's audio) is written only after the transport exists. Awaiting `send()` inside +/// `connect` therefore deadlocks until the dial timeout (found by the in-process Bedrock mock, +/// FRD-023 TC-XL-05). The open runs in its own task instead: `connect` returns at once, the +/// input channel buffers what the driver writes, and the first `recv` waits for the open. +enum BedrockOutputState { + Opening(tokio::task::JoinHandle>), + Open(Box), + Closed, +} + pub struct BedrockBidiTransport { /// Outbound input events → the SDK's input event-stream sender (via the /// channel-fed `async_stream` installed at connect). Cloned `send()` surface. input_tx: mpsc::Sender, /// This connection's OWN output event receiver (owned outright, dropped with - /// the transport) — the bidi stream's server→client half. (`event_receiver` - /// is a private module; the public re-export is via `primitives::event_stream`, - /// exactly as `AwsTranscribeTransport` references its result stream.) - output_rx: aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver< - BidiOutputEvent, - aws_sdk_bedrockruntime::types::error::InvokeModelWithBidirectionalStreamOutputError, - >, + /// the transport) — the bidi stream's server→client half — once the open completes. + output: BedrockOutputState, +} + +impl Drop for BedrockBidiTransport { + fn drop(&mut self) { + if let BedrockOutputState::Opening(handle) = &self.output { + handle.abort(); + } + } } #[async_trait] @@ -595,7 +618,29 @@ impl RealtimeTransport for BedrockBidiTransport { // decoded JSON event as a Text frame, map a stream error to Some(Err), // and a clean end (`Ok(None)`) to None (no reconnect). loop { - match self.output_rx.recv().await { + let output = match &mut self.output { + BedrockOutputState::Open(output) => output, + BedrockOutputState::Closed => return None, + BedrockOutputState::Opening(handle) => { + let opened = match handle.await { + Ok(result) => result, + Err(e) => Err(RealtimeError::ConnectionFailed(format!( + "Bedrock stream open task failed: {e}" + ))), + }; + match opened { + Ok(output) => { + self.output = BedrockOutputState::Open(Box::new(output)); + continue; + } + Err(e) => { + self.output = BedrockOutputState::Closed; + return Some(Err(e)); + } + } + } + }; + match output.recv().await { Ok(Some(event)) => { if let Some(json) = unwrap_output_payload(&event) { return Some(Ok(OutFrame::Text(json))); @@ -622,6 +667,65 @@ impl RealtimeTransport for BedrockBidiTransport { } } +/// May this dial proceed? In Bud mode a Bedrock stream is signed with the DEPLOYMENT's static +/// keys or not at all — the gateway's own AWS identity (env, shared config, instance role) is +/// never lent to a tenant session (FRD-023 D-5, RT0). Pure, so the rule is testable without +/// flipping the process-wide Bud-mode flag. +pub(crate) fn bedrock_dial_allowed( + credentials: Option<&crate::core::realtime::base::AwsStaticCredentials>, + in_bud_mode: bool, +) -> RealtimeResult<()> { + if credentials.is_none() && in_bud_mode { + return Err(RealtimeError::AuthenticationFailed( + "a Bud deployment's Nova Sonic session is signed with the deployment's own AWS keys; \ + the gateway's AWS identity is never used" + .to_string(), + )); + } + Ok(()) +} + +/// The Bedrock client for a dial. With static credentials the config is built from them and +/// the spec ALONE — no environment, no shared config, no instance metadata: in Bud mode not even +/// an `AWS_ENDPOINT_URL` may redirect a request signed with a tenant's keys. Without, the +/// `aws-config` default chain (the native path, as before). +async fn bedrock_client( + factory: &BedrockBidiTransportFactory, + region: Option, +) -> BedrockClient { + match &factory.credentials { + Some(c) => { + let mut b = aws_sdk_bedrockruntime::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .credentials_provider(aws_credential_types::Credentials::new( + c.access_key_id.clone(), + c.secret_access_key.clone(), + c.session_token.clone(), + None, + "bud-voice-table", + )) + .region(region.map(aws_config::Region::new)); + if let Some(url) = &factory.endpoint_url { + b = b.endpoint_url(url.clone()); + } + if let Some(h) = &factory.http_client { + b = b.http_client(h.clone()); + } + BedrockClient::from_conf(b.build()) + } + None => { + let mut loader = aws_config::defaults(BehaviorVersion::latest()); + if let Some(r) = region { + loader = loader.region(aws_config::Region::new(r)); + } + if let Some(h) = &factory.http_client { + loader = loader.http_client(h.clone()); + } + BedrockClient::new(&loader.load().await) + } + } +} + /// Factory for the AWS NOVA SONIC pattern: opens an Amazon Bedrock /// `InvokeModelWithBidirectionalStream` HTTP/2 bidi event stream and returns a /// [`BedrockBidiTransport`] over it. @@ -636,8 +740,56 @@ impl RealtimeTransport for BedrockBidiTransport { /// chain (region from the spec, else the environment), install a channel-fed /// `async_stream` as the input half, `send()` the request, and hand back the /// output [`EventReceiver`](aws_sdk_bedrockruntime::primitives::event_stream::EventReceiver) -/// half. Credentials are AWS SigV4 via the default chain — NO api-key. -pub struct BedrockBidiTransportFactory; +/// half. Credentials are AWS SigV4 — via the default chain, or (a Bud deployment, FRD-023 +/// RT7.2) the deployment's static key pair ([`Self::with_credentials`]). NO api-key. +#[derive(Clone, Default)] +pub struct BedrockBidiTransportFactory { + credentials: Option, + endpoint_url: Option, + http_client: Option, +} + +impl std::fmt::Debug for BedrockBidiTransportFactory { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BedrockBidiTransportFactory") + .field("static_credentials", &self.credentials.is_some()) + .field("endpoint_url", &self.endpoint_url) + .finish() + } +} + +impl BedrockBidiTransportFactory { + /// The `aws-config` default credential chain (the native path). + pub fn default_chain() -> Self { + Self::default() + } + + /// Sign with this key pair and nothing else. + pub fn with_credentials( + credentials: crate::core::realtime::base::AwsStaticCredentials, + ) -> Self { + Self { + credentials: Some(credentials), + ..Self::default() + } + } + + /// A Bedrock endpoint other than the region's (already SSRF-validated by the caller). + pub fn endpoint_url(mut self, url: Option) -> Self { + self.endpoint_url = url; + self + } + + /// The HTTP client the SDK dials with. `None`: the SDK's own. Tests pass an in-process + /// connector that speaks the Bedrock event stream. + pub fn http_client( + mut self, + client: Option, + ) -> Self { + self.http_client = client; + self + } +} #[async_trait] impl RealtimeTransportFactory for BedrockBidiTransportFactory { @@ -647,58 +799,53 @@ impl RealtimeTransportFactory for BedrockBidiTransportFactory { "BedrockBidiTransportFactory only supports ConnectSpec::BedrockBidi".to_string(), )); }; - - // Bound the WHOLE dial (config load + stream open) so a misconfigured - // deployment fails fast with a clear error instead of stalling the client - // across slow IMDS-timeout backoff retries (see BEDROCK_CONNECT_TIMEOUT). - tokio::time::timeout(BEDROCK_CONNECT_TIMEOUT, async move { - // Build the AWS config + Bedrock client from the DEFAULT credential chain - // (env / shared config / IAM role), exactly like AwsTranscribeTransport. - // Region: the spec's, else resolved by the loader from the environment. - let mut loader = aws_config::defaults(BehaviorVersion::latest()); - if let Some(r) = region { - loader = loader.region(aws_config::Region::new(r)); + bedrock_dial_allowed( + self.credentials.as_ref(), + crate::auth::bud_mode::process_in_bud_mode(), + )?; + + // Bound the client build (a default-chain config load can stall on IMDS) so a + // misconfigured deployment fails fast instead of stalling the client. + let client = tokio::time::timeout(BEDROCK_CONNECT_TIMEOUT, bedrock_client(self, region)) + .await + .map_err(|_| { + RealtimeError::ConnectionFailed( + "Bedrock connect timed out (check AWS credentials/region)".to_string(), + ) + })?; + + // The input half: a bounded channel whose receiver drives an async stream + // of union events; the sender is the transport's `send()` surface. (Same + // channel-fed `async_stream` Transcribe attaches as its audio input — here + // it is the outbound-frame path.) + let (input_tx, mut input_rx) = mpsc::channel::(BEDROCK_INPUT_CHANNEL_DEPTH); + let input_stream = async_stream::stream! { + while let Some(event) = input_rx.recv().await { + yield Ok::(event); } - let aws_config = loader.load().await; - let client = BedrockClient::new(&aws_config); - - // The input half: a bounded channel whose receiver drives an async stream - // of union events; the sender is the transport's `send()` surface. (Same - // channel-fed `async_stream` Transcribe attaches as its audio input — here - // it is the outbound-frame path.) - let (input_tx, mut input_rx) = - mpsc::channel::(BEDROCK_INPUT_CHANNEL_DEPTH); - let input_stream = async_stream::stream! { - while let Some(event) = input_rx.recv().await { - yield Ok::(event); - } - }; + }; - // Open the bidi stream. `.body(stream.into())` + `.send()` mirrors - // Transcribe's `.audio_stream(stream.into()).send_with(&client)`. - let output = client + // Open the bidi stream in its own task (see `BedrockOutputState`): the SDK returns + // from `send()` only at the first output event, which Nova sends only after input. + let open = tokio::spawn(async move { + client .invoke_model_with_bidirectional_stream() .model_id(model_id) .body(input_stream.into()) .send() .await + .map(|output| output.body) .map_err(|e| { RealtimeError::ConnectionFailed(format!( "failed to open Bedrock bidirectional stream: {e}" )) - })?; + }) + }); - Ok(Box::new(BedrockBidiTransport { - input_tx, - output_rx: output.body, - }) as Box) - }) - .await - .map_err(|_| { - RealtimeError::ConnectionFailed( - "Bedrock connect timed out (check AWS credentials/region)".to_string(), - ) - })? + Ok(Box::new(BedrockBidiTransport { + input_tx, + output: BedrockOutputState::Opening(open), + }) as Box) } } @@ -711,6 +858,108 @@ mod tests { use super::*; use serde_json::json; + fn aws_keys() -> crate::core::realtime::base::AwsStaticCredentials { + crate::core::realtime::base::AwsStaticCredentials { + access_key_id: "AKIDBUDDEPLOYMENT1".into(), + secret_access_key: "deployment-secret".into(), + session_token: None, + } + } + + /// FRD-023 RT7.2 🔒 — in Bud mode a Bedrock dial without the deployment's keys is refused; + /// the default chain (the gateway's own AWS identity) is for the native path only. + #[test] + fn rt7_2_bud_mode_never_dials_bedrock_with_the_gateway_identity() { + assert!(matches!( + bedrock_dial_allowed(None, true), + Err(RealtimeError::AuthenticationFailed(_)) + )); + assert!(bedrock_dial_allowed(Some(&aws_keys()), true).is_ok()); + assert!(bedrock_dial_allowed(None, false).is_ok()); + } + + /// The key pair never reaches a log line through `Debug`. + #[test] + fn aws_static_credentials_debug_is_redacted() { + let printed = format!("{:?}", aws_keys()); + assert!(!printed.contains("AKIDBUDDEPLOYMENT1"), "{printed}"); + assert!(!printed.contains("deployment-secret"), "{printed}"); + let printed = format!( + "{:?}", + BedrockBidiTransportFactory::with_credentials(aws_keys()) + ); + assert!(!printed.contains("deployment-secret"), "{printed}"); + } + + /// FRD-023 RT7.2 🔒 — a factory holding a deployment's keys SIGNS WITH THEM (SigV4, the + /// `bedrock` service, the deployment's region), whatever the process environment holds. + #[tokio::test] + async fn rt7_2_static_credentials_sign_the_bedrock_request() { + use aws_smithy_runtime_api::client::http::{ + HttpClient, HttpConnector, HttpConnectorFuture, HttpConnectorSettings, + SharedHttpClient, SharedHttpConnector, + }; + use aws_smithy_runtime_api::client::orchestrator::HttpRequest; + use aws_smithy_runtime_api::client::runtime_components::RuntimeComponents; + use aws_smithy_runtime_api::http::{Response, StatusCode}; + use aws_smithy_types::body::SdkBody; + use std::sync::{Arc, Mutex}; + + #[derive(Debug, Clone, Default)] + struct Capture(Arc>>); + impl HttpConnector for Capture { + fn call(&self, req: HttpRequest) -> HttpConnectorFuture { + let auth = req + .headers() + .get("authorization") + .unwrap_or_default() + .to_string(); + self.0.lock().unwrap().push((req.uri().to_string(), auth)); + HttpConnectorFuture::new(async move { + let mut resp = + Response::new(StatusCode::try_from(200u16).unwrap(), SdkBody::empty()); + resp.headers_mut() + .insert("content-type", "application/vnd.amazon.eventstream"); + Ok(resp) + }) + } + } + impl HttpClient for Capture { + fn http_connector( + &self, + _s: &HttpConnectorSettings, + _c: &RuntimeComponents, + ) -> SharedHttpConnector { + SharedHttpConnector::new(self.clone()) + } + } + + let capture = Capture::default(); + let factory = BedrockBidiTransportFactory::with_credentials(aws_keys()) + .http_client(Some(SharedHttpClient::new(capture.clone()))); + let mut transport = factory + .connect(ConnectSpec::BedrockBidi { + model_id: "amazon.nova-2-sonic-v1:0".into(), + region: Some("eu-north-1".into()), + }) + .await + .expect("connect returns at once; the stream opens in the background"); + // The first read waits for the open (here: an empty stream, so it ends). + let _ = tokio::time::timeout(std::time::Duration::from_secs(10), transport.recv()).await; + let seen = capture.0.lock().unwrap().clone(); + let (uri, auth) = seen.first().expect("the SDK sent the request"); + assert!( + uri.starts_with("https://bedrock-runtime.eu-north-1.amazonaws.com/"), + "{uri}" + ); + assert!(auth.starts_with("AWS4-HMAC-SHA256 "), "{auth}"); + assert!( + auth.contains("Credential=AKIDBUDDEPLOYMENT1/") + && auth.contains("/eu-north-1/bedrock/aws4_request"), + "signed with the deployment's key in its region: {auth}" + ); + } + /// The join-url extraction the REST-handshake factory does: present string /// field at the pointer ⇒ that url. #[test] @@ -871,7 +1120,9 @@ mod tests { headers: vec![], }; assert!(matches!( - BedrockBidiTransportFactory.connect(spec).await, + BedrockBidiTransportFactory::default_chain() + .connect(spec) + .await, Err(RealtimeError::ConnectionFailed(_)) )); } @@ -1041,7 +1292,9 @@ mod tests { path: "/tmp/x.sock".to_string(), }; assert!(matches!( - BedrockBidiTransportFactory.connect(spec).await, + BedrockBidiTransportFactory::default_chain() + .connect(spec) + .await, Err(RealtimeError::ConnectionFailed(_)) )); } diff --git a/gateway/tests/realtime_full_integration.rs b/gateway/tests/realtime_full_integration.rs index 0c27055a..0ae244e0 100644 --- a/gateway/tests/realtime_full_integration.rs +++ b/gateway/tests/realtime_full_integration.rs @@ -147,6 +147,8 @@ fn rich_config_for(name: &str) -> RealtimeConfig { realtime_endpoint_override: None, reconnection: None, trace: None, + // SERVER-SET connection cap (FRD-023 RT7); the protocol's own when unset. + max_connection: None, }; match name { From e2e9a391862260757e919b8c5049320369a93aa9 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 06:29:47 +0530 Subject: [PATCH 11/17] feat(gateway): GA translate engine on /v1/realtime for Gemini Live, Nova 2 Sonic and per-minute agents (FRD-023 RT7.0, RT7.1, RT7.2, RT7.4) A deployment whose vendor has no OpenAI Realtime GA surface (CONTRACTS C7: gemini, nova_sonic, deepgram_voice_agent, elevenlabs_convai, hume_evi) is served by WaaV's native provider behind a GA facade (handlers/openai_realtime/facade.rs). Everything around the session stays the relay's: authentication and ek_bud_ secrets, resolution, the admission held for the session, revalidation, idle and maximum-length limits, pings, the 1012 drain, the session span and the client-event policy. - The vendor leg comes from voice_table only: the credential (Nova: the key pair in credential_parts and provider_params.region, refused without either), the model, and an api_base that is converted to ws(s) and SSRF-validated. A provider that cannot be built from the entry is refused before the upgrade (502 deployment_misconfigured). The vendors serve speech-to-speech sessions only. - The vendor setup is deferred to the client's first session.update, so the one opening message these vendors accept carries the client's voice, instructions, tools and turn detection over the deployment's defaults; afterwards a change to any of them is refused with event_not_allowed naming the field (the same value again is accepted). - GA in: session.update, input_audio_buffer.append (24 kHz resampled to the vendor's rate)/commit/clear, conversation.item.create (user text, function_call_output), response.create/cancel. Refused by name: conversation.item.truncate/retrieve/delete, output_audio_buffer.clear, audio and image content, MCP tools, stored prompts, out-of-band responses, unknown events. GA out: response.created ... audio and transcript deltas (resampled to 24 kHz; WAV chunks unwrapped) ... function calls ... response.done with a usage the gateway computed. - Metering: a vendor usage report is the usage of the response it belongs to and exactly one voice.turn; a report outside a response is billed alone; a response cut off by close keeps its usage. Per-minute vendors bill duration segments from the moment the vendor connection opens, through the relay's segment clock (now a shared SegmentClock). - A vendor that reconnects (goAway, the Nova cap) is invisible to the client: vendor-bound work waits out the gap. Tests (in-process mocks of each wire protocol, including a Bedrock event-stream connector): TC-XL-01...05 and 07, plus xAI priced per minute. Co-Authored-By: Claude Opus 5.5 --- gateway/Cargo.lock | 2 + gateway/Cargo.toml | 5 + .../src/handlers/openai_realtime/facade.rs | 1821 +++++++++++++++++ .../handlers/openai_realtime/facade/tests.rs | 651 ++++++ .../src/handlers/openai_realtime/metering.rs | 160 +- gateway/src/handlers/openai_realtime/mod.rs | 3 + .../src/handlers/openai_realtime/session.rs | 177 +- .../src/handlers/openai_realtime/upstream.rs | 6 +- gateway/tests/openai_realtime_relay.rs | 59 + gateway/tests/realtime_translate.rs | 1440 +++++++++++++ 10 files changed, 4259 insertions(+), 65 deletions(-) create mode 100644 gateway/src/handlers/openai_realtime/facade.rs create mode 100644 gateway/src/handlers/openai_realtime/facade/tests.rs create mode 100644 gateway/tests/realtime_translate.rs diff --git a/gateway/Cargo.lock b/gateway/Cargo.lock index 2da4e825..a5a10ac2 100644 --- a/gateway/Cargo.lock +++ b/gateway/Cargo.lock @@ -7794,6 +7794,7 @@ dependencies = [ "hmac 0.12.1", "hound", "http 1.4.1", + "http-body 1.0.1", "http-body-util", "inventory", "jsonwebtoken 10.4.0", @@ -7830,6 +7831,7 @@ dependencies = [ "reqwest 0.12.28", "resil", "rhai", + "rsa", "rtrb", "rubato 0.16.2", "rustls 0.23.45", diff --git a/gateway/Cargo.toml b/gateway/Cargo.toml index 990aafea..966242af 100644 --- a/gateway/Cargo.toml +++ b/gateway/Cargo.toml @@ -331,6 +331,11 @@ pulp = "0.18" [dev-dependencies] aws-smithy-eventstream = "0.60" +# FRD-023 RT7: the in-process Bedrock event-stream mock (Nova Sonic) and test credentials +# encrypted the way budapp encrypts them. All three are already in the lock (bud-auth, hyper). +http-body = "1" +http-body-util = "0.1" +rsa = { version = "0.9", features = ["sha2"] } tower = { version = "0.5.2", features = ["util"] } tempfile = "3.8" serial_test = "3.2" diff --git a/gateway/src/handlers/openai_realtime/facade.rs b/gateway/src/handlers/openai_realtime/facade.rs new file mode 100644 index 00000000..f82726c8 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/facade.rs @@ -0,0 +1,1821 @@ +//! The translate engine: OpenAI Realtime GA over WaaV's native realtime providers (FRD-023 §5.7, +//! RT7). +//! +//! A deployment whose vendor has no GA surface — Gemini Live, Nova 2 Sonic, and the per-minute +//! voice agents (Deepgram Voice Agent, ElevenLabs Agents, Hume EVI) — is served by the vendor's +//! [`BaseRealtime`] provider (the S2S scaffold) behind this facade. The client still speaks GA; +//! everything around the session is the relay's: authentication, `ek_bud_` secrets, resolution, +//! the admission held for the session, revalidation, idle and maximum-length limits, pings, the +//! drain close, the session span and the policy on client events. +//! +//! **What is translated** (§5.7): `session.update` (voice, instructions, turn detection, tools, +//! input transcription), `input_audio_buffer.{append,commit,clear}` (24 kHz PCM resampled to the +//! vendor's rate), `conversation.item.create` (user text, `function_call_output`), +//! `response.{create,cancel}`. The vendor's events come back as the GA server events a client +//! plays: `response.created` … `response.output_audio.delta` … `response.done`, with a `usage` +//! the gateway computed from the vendor's own report. +//! +//! **What is refused** with `event_not_allowed`, naming the event or field: what the vendor cannot +//! do (`conversation.item.truncate`/`retrieve`/`delete`, `output_audio_buffer.clear`, audio or +//! image content, MCP tools, stored prompts, out-of-band responses) and, once the vendor's session +//! is set up, any CHANGE to a setup-time field (voice, instructions, tools, turn detection): these +//! vendors take them only in their opening message (Gemini's `setup`), so a later change would +//! otherwise be accepted and silently ignored. Re-sending the same value is fine — SDKs resend the +//! whole session on every update. +//! +//! **Setup is deferred** to the client's first `session.update` (or first audio, text or +//! response), so the one opening message the vendor accepts carries the client's configuration as +//! well as the deployment's. The client gets `session.created` at once; `session.updated` when +//! the vendor has accepted the setup. +//! +//! **Metering** (§5.10): a token vendor's usage report (Gemini `usageMetadata`, Nova `usageEvent`) +//! becomes the `usage` of the `response.done` it belongs to and exactly one `voice.turn`; a report +//! outside any response is metered on its own. A per-minute vendor bills 60 s duration segments +//! from the moment its connection opens. + +use std::collections::VecDeque; +use std::sync::Arc; +use std::time::Duration; + +use axum::extract::ws::{Message, WebSocket}; +use base64::Engine as _; +use base64::prelude::BASE64_STANDARD; +use bytes::Bytes; +use futures_util::StreamExt; +use serde_json::{Map, Value, json}; +use tokio::sync::mpsc; +use tokio::time::Instant; +use tracing::{debug, info, warn}; + +use bud_auth::{RealtimeSettings, VoiceEndpoint}; + +use crate::core::audio::resampler::{StreamResampler, flush_pcm16, resample_pcm16}; +use crate::core::realtime::scaffold::{BedrockBidiTransportFactory, S2sEvent, UsageReport}; +use crate::core::realtime::{ + AwsStaticCredentials, BaseRealtime, FunctionDefinition, InputTranscriptionConfig, + RealtimeConfig, ReconnectionConfig, SpeechEvent, ToolDefinition, TranscriptRole, + TurnDetectionConfig, +}; +use crate::core::realtime_cost::RealtimeUsage; +use crate::middleware::connection_limit::ConnectionSlot; +use crate::state::AppState; + +use super::metering::{SegmentClock, SessionMeter, ga_usage}; +use super::policy::{self, ClientOutcome, ClientRules}; +use super::session::{ + self, CLIENT_QUEUE, End, Engine, Outbound, Prepared, SessionLimits, Timings, gateway_error, + next_event_id, +}; +use super::upstream::{self, UpstreamError}; + +/// The client's audio format (GA `audio/pcm` is 24 kHz PCM16 mono only). +const GA_RATE: u32 = 24_000; +/// Vendor-bound work queued while the vendor connection is being replaced (a `goAway`, the Nova +/// cap). Past this the vendor is not coming back in useful time. +const MAX_PENDING: usize = 2_048; + +/// A vendor served by the translate engine (CONTRACTS C7). +#[derive(Debug)] +pub struct TranslateVendor { + /// `voice_table.vendor`. + pub vendor: &'static str, + /// The provider's name in WaaV's realtime registry. + provider: &'static str, + /// Billed by duration, not tokens (its usage is not reported per response). + pub per_minute: bool, + /// Announces its session (`S2sEvent::SessionReady`); until it does, nothing else is sent. + awaits_ready: bool, +} + +const VENDORS: &[TranslateVendor] = &[ + TranslateVendor { + vendor: "gemini", + provider: "gemini", + per_minute: false, + awaits_ready: true, + }, + TranslateVendor { + vendor: "nova_sonic", + provider: "nova_sonic", + per_minute: false, + awaits_ready: false, + }, + TranslateVendor { + vendor: "deepgram_voice_agent", + provider: "deepgram", + per_minute: true, + awaits_ready: true, + }, + TranslateVendor { + vendor: "elevenlabs_convai", + provider: "elevenlabs", + per_minute: true, + awaits_ready: true, + }, + TranslateVendor { + vendor: "hume_evi", + provider: "hume", + per_minute: true, + awaits_ready: true, + }, +]; + +fn vendor_info(vendor: &str) -> Option<&'static TranslateVendor> { + let v = vendor.trim().to_ascii_lowercase(); + VENDORS.iter().find(|t| t.vendor == v) +} + +/// Is this `voice_table.vendor` served by the translate engine? +pub fn is_translate_vendor(vendor: &str) -> bool { + vendor_info(vendor).is_some() +} + +// ============================================================================================= +// The plan: everything decided before the upgrade +// ============================================================================================= + +/// How to reach a translated vendor, decided before the upgrade from the deployment alone. +pub struct TranslatePlan { + pub info: &'static TranslateVendor, + /// The provider config: the deployment's credential, model, address and defaults. + base: RealtimeConfig, + /// Nova Sonic: the deployment's key pair (never the gateway's AWS identity). + aws: Option, + /// Nova Sonic: a Bedrock endpoint other than the region's (`api_base`). + bedrock_endpoint: Option, + /// An address from `voice_table` that must clear the SSRF validator. + ssrf: Option<(String, &'static [&'static str])>, +} + +impl std::fmt::Debug for TranslatePlan { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Never the config: it carries the vendor key. + f.debug_struct("TranslatePlan") + .field("vendor", &self.info.vendor) + .field("model", &self.base.model) + .finish() + } +} + +fn nonempty(v: Option<&str>) -> Option { + v.map(str::trim) + .filter(|s| !s.is_empty()) + .map(str::to_string) +} + +/// GA `turn_detection` (an object, or `null` for none) in the provider vocabulary. +fn turn_detection_from_ga(v: &Value) -> Option { + if v.is_null() { + return Some(TurnDetectionConfig::None); + } + serde_json::from_value(v.clone()).ok() +} + +/// GA function tools (`{type, name, description, parameters}`) in the provider vocabulary. +fn tools_from_ga(tools: &[Value]) -> Vec { + tools + .iter() + .filter(|t| t.get("type").and_then(Value::as_str) == Some("function")) + .map(|t| ToolDefinition { + tool_type: "function".into(), + function: FunctionDefinition { + name: t + .get("name") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(), + description: nonempty(t.get("description").and_then(Value::as_str)), + parameters: t.get("parameters").cloned(), + }, + }) + .collect() +} + +fn tools_to_ga(tools: Option<&Vec>) -> Value { + Value::Array( + tools + .into_iter() + .flatten() + .map(|t| { + json!({ + "type": "function", + "name": t.function.name, + "description": t.function.description, + "parameters": t.function.parameters, + }) + }) + .collect(), + ) +} + +impl TranslatePlan { + /// Build the vendor leg from the deployment (FRD-023 §5.3, CONTRACTS C7). The credential, + /// model and address come from `voice_table` and nowhere else (D-5). + pub fn build( + endpoint: &VoiceEndpoint, + settings: Option<&RealtimeSettings>, + ) -> Result { + let info = vendor_info(&endpoint.vendor) + .ok_or_else(|| UpstreamError::UnsupportedVendor(endpoint.vendor.clone()))?; + if settings.is_some_and(RealtimeSettings::is_transcription) { + return Err(UpstreamError::Misconfigured(format!( + "vendor '{}' serves speech-to-speech sessions only, not transcription", + info.vendor + ))); + } + + let mut base = RealtimeConfig { + provider: info.provider.to_string(), + model: endpoint + .model + .as_deref() + .unwrap_or_default() + .trim() + .to_string(), + // Transparent reconnects (a `goAway`, the Nova cap) need the driver's supervisor; + // a vendor that stays away past these attempts ends the session with 1011. + reconnection: Some(ReconnectionConfig::default()), + ..Default::default() + }; + let mut aws = None; + let mut bedrock_endpoint = None; + let mut ssrf = None; + + if info.vendor == "nova_sonic" { + // SigV4 with the deployment's own key pair (budapp packs the AWS names without the + // `aws_` prefix) in the deployment's region — refused without either, never falling + // back to the gateway's AWS identity (RT0, D-5). + let part = |name: &str| { + endpoint + .credential_parts + .as_ref() + .and_then(|p| p.get(name)) + .map(|v| v.trim().to_string()) + .filter(|v| !v.is_empty()) + }; + let (Some(access_key_id), Some(secret_access_key)) = + (part("access_key_id"), part("secret_access_key")) + else { + return Err(UpstreamError::Misconfigured( + "the credential must be an AWS access key pair; without it the session would \ + authenticate as the gateway's own AWS identity" + .into(), + )); + }; + aws = Some(AwsStaticCredentials { + access_key_id, + secret_access_key, + session_token: part("session_token"), + }); + let region = endpoint.provider_param("region").ok_or_else(|| { + UpstreamError::Misconfigured( + "no AWS region is configured for the deployment".into(), + ) + })?; + // The Nova protocol reads its region from the generic `endpoint` slot. + base.endpoint = Some(region.to_string()); + if let Some(api_base) = nonempty(endpoint.api_base.as_deref()) { + ssrf = Some((api_base.clone(), &["https", "http"][..])); + bedrock_endpoint = Some(api_base); + } + } else { + base.api_key = + nonempty(endpoint.credential.as_deref()).ok_or(UpstreamError::MissingCredential)?; + if let Some(api_base) = nonempty(endpoint.api_base.as_deref()) { + let ws = upstream::to_ws_base(&api_base)?; + ssrf = Some((ws.clone(), &["ws", "wss"][..])); + // Server-config-only override, now from the deployment (F-5: an `https://` base + // is converted rather than ignored). + base.realtime_endpoint_override = Some(ws); + } + } + + // The deployment's defaults (§5.6). The client may override them in its first + // `session.update`, which is folded into the same setup. + if let Some(d) = settings.and_then(|s| s.defaults.as_ref()) { + base.voice = nonempty(d.voice.as_deref()); + base.instructions = nonempty(d.instructions.as_deref()); + base.modalities = d.output_modalities.clone(); + base.turn_detection = d.turn_detection.as_ref().and_then(turn_detection_from_ga); + base.input_audio_transcription = d + .input_transcription + .as_ref() + .and_then(|t| t.model.clone()) + .map(|model| InputTranscriptionConfig { model }); + base.input_audio_noise_reduction = d.noise_reduction.clone(); + base.max_response_output_tokens = d.max_output_tokens.map(|m| m as i32); + } + + let plan = Self { + info, + base, + aws, + bedrock_endpoint, + ssrf, + }; + // Construct (never connect) once: a deployment the provider cannot be built for — an + // ElevenLabs entry with no agent id — is refused before it takes a slot. + drop(plan.provider(plan.base.clone(), None).map_err(|e| { + UpstreamError::Misconfigured(format!("the vendor provider refused the entry: {e}")) + })?); + Ok(plan) + } + + /// SSRF-validate the one address that is not a vendor constant (blocking DNS, so off the + /// async workers). + pub async fn validate(&self) -> Result<(), UpstreamError> { + let Some((url, schemes)) = self.ssrf.clone() else { + return Ok(()); + }; + tokio::task::spawn_blocking(move || crate::core::net::validate_url_for_ssrf(&url, schemes)) + .await + .map_err(|e| UpstreamError::InvalidApiBase(format!("validation task failed: {e}")))? + .map_err(UpstreamError::InvalidApiBase) + } + + /// Build (not connect) the provider for `config`. + fn provider( + &self, + config: RealtimeConfig, + bedrock_http: Option, + ) -> Result, crate::core::realtime::RealtimeError> { + if self.info.vendor == "nova_sonic" { + let Some(keys) = self.aws.clone() else { + return Err(crate::core::realtime::RealtimeError::AuthenticationFailed( + "no deployment key pair".into(), + )); + }; + let factory = BedrockBidiTransportFactory::with_credentials(keys) + .endpoint_url(self.bedrock_endpoint.clone()) + .http_client(bedrock_http); + return Ok(Box::new( + crate::core::realtime::NovaSonicRealtime::with_transport(config, factory)?, + )); + } + crate::core::realtime::create_realtime_provider(self.info.provider, config) + } +} + +// ============================================================================================= +// The translator: GA events ↔ provider calls and events (pure, no I/O) +// ============================================================================================= + +/// One thing the session shell must do, in order. +#[derive(Debug, PartialEq)] +pub(super) enum Act { + /// A GA server event for the client. + Client(String), + /// Build the provider from [`Translator::config`] and connect it (the deferred setup). + Connect, + /// Revalidate the caller before a `response.create` (D-17). + Revalidate, + /// Client audio for the vendor, at the vendor's rate. + Audio(Bytes), + Text(String), + CreateResponse, + Cancel, + Commit, + Clear, + ToolResult { + call_id: String, + output: String, + }, + /// One billed response. + Meter { + response_id: String, + status: &'static str, + usage: RealtimeUsage, + transcript: Option, + }, +} + +/// The response in flight. +#[derive(Debug)] +struct Response { + id: String, + /// The assistant message item (audio + transcript), once output began. + item_id: Option, + output_index: u64, + /// Output items completed so far (function calls, then the message at the end). + output: Vec, + transcript: String, + /// Assistant interim text since its last final (see `Translator::assistant_text`). + interim_since_final: String, + usage: RealtimeUsage, + usage_seen: bool, + /// The client cancelled it: its remaining output is dropped. + cancelled: bool, +} + +/// The user's turn in flight (their speech and its transcription). +#[derive(Debug, Default)] +struct UserTurn { + item_id: Option, + transcript: String, + speaking: bool, +} + +pub(super) struct Translator { + vendor: &'static TranslateVendor, + /// The deployment name the client connected with (reported as the session's model). + deployment: String, + rules: ClientRules, + /// What the vendor is (or will be) set up with. + config: RealtimeConfig, + session_id: String, + vendor_session_id: Option, + /// The provider was asked to connect. + setup: bool, + /// The vendor accepted the setup. + ready: bool, + /// `session.updated` owed once the vendor is ready (the client's event ids). + owed_updates: Vec>, + /// Vendor-bound acts held until the vendor is ready. + held: VecDeque, + response: Option, + last_response_id: Option, + last_item_id: Option, + user: UserTurn, + input_rate: u32, + output_rate: u32, + to_vendor: StreamResampler, + to_client: StreamResampler, +} + +fn new_id(prefix: &str) -> String { + format!("{prefix}_bud_{}", uuid::Uuid::new_v4().simple()) +} + +/// A GA server event: `type`, a fresh `event_id`, and the body's fields. +fn event(kind: &str, body: Value) -> String { + let mut m = match body { + Value::Object(m) => m, + _ => Map::new(), + }; + m.insert("type".into(), Value::from(kind)); + m.insert("event_id".into(), Value::from(next_event_id())); + Value::Object(m).to_string() +} + +macro_rules! ev { + ($kind:expr, $($body:tt)+) => { + event($kind, json!($($body)+)) + }; +} + +/// Decode a vendor audio chunk to PCM16 mono and its rate: raw PCM at the declared rate, or a +/// WAV container (Hume EVI sends one per chunk). +fn pcm_of(data: &[u8], declared_rate: u32) -> (u32, std::borrow::Cow<'_, [u8]>) { + if data.len() > 44 && &data[0..4] == b"RIFF" && &data[8..12] == b"WAVE" { + let (mut rate, mut channels, mut bits) = (declared_rate, 1u16, 16u16); + let mut i = 12; + while i + 8 <= data.len() { + let id = &data[i..i + 4]; + let len = + u32::from_le_bytes([data[i + 4], data[i + 5], data[i + 6], data[i + 7]]) as usize; + let body = i + 8; + if id == b"fmt " && body + 16 <= data.len() { + channels = u16::from_le_bytes([data[body + 2], data[body + 3]]); + rate = u32::from_le_bytes([ + data[body + 4], + data[body + 5], + data[body + 6], + data[body + 7], + ]); + bits = u16::from_le_bytes([data[body + 14], data[body + 15]]); + } else if id == b"data" { + let end = (body + len).min(data.len()); + let pcm = &data[body..end]; + if bits != 16 || channels == 0 { + return (rate, std::borrow::Cow::Owned(Vec::new())); + } + if channels == 1 { + return (rate, std::borrow::Cow::Borrowed(pcm)); + } + // Down-mix to mono: the first channel of every frame. + let frame = 2 * channels as usize; + let mono: Vec = pcm.chunks_exact(frame).flat_map(|f| [f[0], f[1]]).collect(); + return (rate, std::borrow::Cow::Owned(mono)); + } + i = body + len + (len & 1); + } + } + (declared_rate, std::borrow::Cow::Borrowed(data)) +} + +impl Translator { + pub(super) fn new( + vendor: &'static TranslateVendor, + deployment: String, + rules: ClientRules, + config: RealtimeConfig, + session_id: &str, + ) -> Self { + Self { + vendor, + deployment, + rules, + config, + session_id: session_id.to_string(), + vendor_session_id: None, + setup: false, + ready: false, + owed_updates: Vec::new(), + held: VecDeque::new(), + response: None, + last_response_id: None, + last_item_id: None, + user: UserTurn::default(), + input_rate: GA_RATE, + output_rate: GA_RATE, + to_vendor: StreamResampler::new(), + to_client: StreamResampler::new(), + } + } + + /// The provider config the deferred setup uses. + pub(super) fn config(&self) -> RealtimeConfig { + self.config.clone() + } + + pub(super) fn vendor_session_id(&self) -> Option { + self.vendor_session_id.clone() + } + + /// The GA session object the client is told about. + fn session_object(&self) -> Value { + let c = &self.config; + let format = json!({"type": "audio/pcm", "rate": GA_RATE}); + json!({ + "object": "realtime.session", + "type": "realtime", + "id": self.session_id, + "model": self.deployment, + "output_modalities": c.modalities.clone().unwrap_or_else(|| vec!["audio".into()]), + "instructions": c.instructions, + "tools": tools_to_ga(c.tools.as_ref()), + "audio": { + "input": { + "format": format, + "turn_detection": c.turn_detection.as_ref().map(|t| match t { + TurnDetectionConfig::None => Value::Null, + other => serde_json::to_value(other).unwrap_or(Value::Null), + }), + "transcription": c.input_audio_transcription.as_ref().map(|t| json!({"model": t.model})), + }, + "output": {"format": format, "voice": c.voice}, + }, + }) + } + + pub(super) fn session_created(&self) -> String { + ev!("session.created", {"session": self.session_object()}) + } + + fn session_updated(&self) -> String { + ev!("session.updated", {"session": self.session_object()}) + } + + fn refuse(param: &str, message: impl Into, event_id: Option<&str>) -> Act { + Act::Client(gateway_error( + "event_not_allowed", + &message.into(), + Some(param), + event_id, + )) + } + + fn untranslatable(&self, what: &str, event_id: Option<&str>) -> Act { + Self::refuse( + what, + format!( + "`{what}` cannot be served on this deployment: its vendor ({}) has no equivalent.", + self.vendor.vendor + ), + event_id, + ) + } + + /// Queue vendor-bound work: held until the vendor is ready, and preceded by the deferred + /// setup when the session has none yet. + fn queue_for_vendor(&mut self, out: &mut Vec, act: Act) { + if !self.setup { + self.setup = true; + out.push(Act::Connect); + } + if self.ready { + out.push(act); + } else if self.held.len() < MAX_PENDING { + self.held.push_back(act); + } + } + + // ---------------------------------------------------------------------------------------- + // Client → vendor + // ---------------------------------------------------------------------------------------- + + pub(super) fn client(&mut self, raw: &str) -> Vec { + let mut out = Vec::new(); + let (text, kind) = match policy::client_event(raw, &self.rules) { + ClientOutcome::Invalid(why) => { + out.push(Act::Client(gateway_error( + "invalid_event", + &format!("The event could not be read: {why}"), + None, + None, + ))); + return out; + } + ClientOutcome::Refuse(r) => { + out.push(Self::refuse(&r.param, r.message, r.event_id.as_deref())); + return out; + } + ClientOutcome::Forward { text, kind, .. } => (text.into_owned(), kind), + }; + let Ok(event) = serde_json::from_str::(&text) else { + return out; + }; + let event_id = event.get("event_id").and_then(Value::as_str); + match kind.as_str() { + "session.update" => self.session_update(&event, event_id, &mut out), + "input_audio_buffer.append" => { + match event + .get("audio") + .and_then(Value::as_str) + .map(|a| BASE64_STANDARD.decode(a)) + { + Some(Ok(pcm)) if pcm.len() % 2 == 0 => { + self.queue_for_vendor(&mut out, Act::Audio(Bytes::from(pcm))) + } + _ => out.push(Act::Client(gateway_error( + "invalid_event", + "`audio` must be base64 PCM16 (audio/pcm, 24 kHz).", + Some("audio"), + event_id, + ))), + } + } + "input_audio_buffer.commit" => { + self.queue_for_vendor(&mut out, Act::Commit); + let item = self.user_item(&mut out); + out.push(Act::Client(ev!("input_audio_buffer.committed", { + "previous_item_id": Value::Null, "item_id": item + }))); + } + "input_audio_buffer.clear" => { + self.queue_for_vendor(&mut out, Act::Clear); + out.push(Act::Client(ev!("input_audio_buffer.cleared", {}))); + } + "conversation.item.create" => self.item_create(&event, event_id, &mut out), + "response.create" => self.response_create(&event, event_id, &mut out), + "response.cancel" => { + if self.response.is_some() { + self.queue_for_vendor(&mut out, Act::Cancel); + self.finish_response("cancelled", &mut out); + } + } + other => out.push(self.untranslatable(other, event_id)), + } + out + } + + fn session_update(&mut self, event: &Value, event_id: Option<&str>, out: &mut Vec) { + let empty = Map::new(); + let s = event + .get("session") + .and_then(Value::as_object) + .unwrap_or(&empty); + // Untranslatable whatever the deployment allows (CONTRACTS C7). + if s.get("prompt").is_some_and(|p| !p.is_null()) { + out.push(self.untranslatable("session.prompt", event_id)); + return; + } + let tools = s.get("tools").and_then(Value::as_array); + if tools.is_some_and(|t| { + t.iter() + .any(|t| t.get("type").and_then(Value::as_str) != Some("function")) + }) { + out.push(self.untranslatable("session.tools.mcp", event_id)); + return; + } + let input = s.get("audio").and_then(|a| a.get("input")); + let output = s.get("audio").and_then(|a| a.get("output")); + for (fmt, param) in [ + ( + input.and_then(|i| i.get("format")), + "session.audio.input.format", + ), + ( + output.and_then(|o| o.get("format")), + "session.audio.output.format", + ), + ] { + if let Some(f) = fmt { + let pcm = f.get("type").and_then(Value::as_str) == Some("audio/pcm") + && f.get("rate") + .is_none_or(|r| r.as_u64() == Some(GA_RATE as u64)); + if !pcm { + out.push(Self::refuse( + param, + "This deployment's vendor is served in audio/pcm at 24 kHz only.", + event_id, + )); + return; + } + } + } + + let voice = output + .and_then(|o| o.get("voice")) + .or_else(|| s.get("voice")) + .and_then(Value::as_str) + .map(str::to_string); + let instructions = s + .get("instructions") + .and_then(Value::as_str) + .map(str::to_string); + let tools = tools.map(|t| tools_from_ga(t)); + let turn_detection = match input.and_then(|i| i.get("turn_detection")) { + None => None, + Some(v) => match turn_detection_from_ga(v) { + Some(td) => Some(td), + None => { + out.push(Self::refuse( + "session.audio.input.turn_detection", + "`turn_detection` is not a turn-detection object.", + event_id, + )); + return; + } + }, + }; + let modalities: Option> = s + .get("output_modalities") + .and_then(|m| serde_json::from_value(m.clone()).ok()); + + if self.setup { + // Setup-time only: a CHANGE after the vendor's opening message is refused rather + // than accepted and ignored. The same value again is fine. + let changed = |a: Option, b: Value| a.is_some_and(|a| a != b); + let as_value = |v: &Option| serde_json::to_value(v).unwrap_or(Value::Null); + let checks = [ + ( + "session.audio.output.voice", + changed( + voice.as_ref().map(|v| json!(v)), + as_value(&self.config.voice), + ), + ), + ( + "session.instructions", + changed( + instructions.as_ref().map(|v| json!(v)), + as_value(&self.config.instructions), + ), + ), + ( + "session.tools", + changed( + tools.as_ref().map(|t| tools_to_ga(Some(t))), + tools_to_ga(self.config.tools.as_ref()), + ), + ), + ( + "session.audio.input.turn_detection", + changed( + turn_detection + .as_ref() + .map(|t| serde_json::to_value(t).unwrap_or(Value::Null)), + serde_json::to_value(&self.config.turn_detection).unwrap_or(Value::Null), + ), + ), + ]; + if let Some((param, _)) = checks.iter().find(|(_, changed)| *changed) { + out.push(Self::refuse( + param, + format!( + "`{param}` is fixed when the session starts on this deployment's vendor \ + ({}); it cannot be changed mid-session.", + self.vendor.vendor + ), + event_id, + )); + return; + } + out.push(Act::Client(self.session_updated())); + return; + } + + // Before setup: fold the client's choices into the one opening message. + if voice.is_some() { + self.config.voice = voice; + } + if instructions.is_some() { + self.config.instructions = instructions; + } + if let Some(t) = tools { + self.config.tools = (!t.is_empty()).then_some(t); + } + if turn_detection.is_some() { + self.config.turn_detection = turn_detection; + } + if let Some(t) = input.and_then(|i| i.get("transcription")) { + self.config.input_audio_transcription = t + .get("model") + .and_then(Value::as_str) + .map(|model| InputTranscriptionConfig { + model: model.to_string(), + }) + .or_else(|| { + (!t.is_null()).then(|| InputTranscriptionConfig { + model: String::new(), + }) + }); + } + if modalities.is_some() { + self.config.modalities = modalities; + } + if let Some(m) = s.get("max_output_tokens").and_then(Value::as_u64) { + self.config.max_response_output_tokens = Some(m.min(i32::MAX as u64) as i32); + } + self.setup = true; + out.push(Act::Connect); + self.owed_updates.push(event_id.map(str::to_string)); + } + + fn item_create(&mut self, event: &Value, event_id: Option<&str>, out: &mut Vec) { + let Some(item) = event.get("item") else { + out.push(Self::refuse("item", "`item` is required.", event_id)); + return; + }; + let item_id = item + .get("id") + .and_then(Value::as_str) + .map(str::to_string) + .unwrap_or_else(|| new_id("item")); + match item.get("type").and_then(Value::as_str) { + Some("message") => { + if item.get("role").and_then(Value::as_str) != Some("user") { + out.push(self.untranslatable("item.role", event_id)); + return; + } + let mut text = String::new(); + for part in item + .get("content") + .and_then(Value::as_array) + .into_iter() + .flatten() + { + match part.get("type").and_then(Value::as_str) { + Some("input_text") | Some("text") => { + if !text.is_empty() { + text.push('\n'); + } + text.push_str(part.get("text").and_then(Value::as_str).unwrap_or("")); + } + Some(other) => { + out.push( + self.untranslatable(&format!("item.content.{other}"), event_id), + ); + return; + } + None => {} + } + } + self.queue_for_vendor(out, Act::Text(text)); + } + Some("function_call_output") => { + let call_id = item + .get("call_id") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + let output = item + .get("output") + .and_then(Value::as_str) + .unwrap_or_default() + .to_string(); + self.queue_for_vendor(out, Act::ToolResult { call_id, output }); + } + other => { + out.push(self.untranslatable( + &format!("item.type.{}", other.unwrap_or("missing")), + event_id, + )); + return; + } + } + let mut item = item.clone(); + item["id"] = Value::from(item_id.clone()); + item["object"] = Value::from("realtime.item"); + item["status"] = Value::from("completed"); + out.push(Act::Client(ev!("conversation.item.added", { + "previous_item_id": self.last_item_id, "item": item + }))); + out.push(Act::Client(ev!("conversation.item.done", { + "previous_item_id": self.last_item_id, "item": item + }))); + self.last_item_id = Some(item_id); + } + + fn response_create(&mut self, event: &Value, event_id: Option<&str>, out: &mut Vec) { + if let Some(r) = event.get("response").and_then(Value::as_object) { + if r.get("prompt").is_some_and(|p| !p.is_null()) { + out.push(self.untranslatable("response.prompt", event_id)); + return; + } + if r.get("conversation").and_then(Value::as_str) == Some("none") { + out.push(self.untranslatable("response.conversation", event_id)); + return; + } + if r.get("input").is_some() { + out.push(self.untranslatable("response.input", event_id)); + return; + } + // Per-response overrides the vendor would ignore. + let differs = |key: &str, current: Value| { + r.get(key).is_some_and(|v| !v.is_null() && *v != current) + }; + if differs( + "instructions", + serde_json::to_value(&self.config.instructions).unwrap_or(Value::Null), + ) { + out.push(self.untranslatable("response.instructions", event_id)); + return; + } + if r.get("tools") + .is_some_and(|t| *t != tools_to_ga(self.config.tools.as_ref())) + { + out.push(self.untranslatable("response.tools", event_id)); + return; + } + let voice = r + .get("audio") + .and_then(|a| a.get("output")) + .and_then(|o| o.get("voice")) + .or_else(|| r.get("voice")); + if voice.is_some_and(|v| { + *v != serde_json::to_value(&self.config.voice).unwrap_or(Value::Null) + }) { + out.push(self.untranslatable("response.audio.output.voice", event_id)); + return; + } + } + out.push(Act::Revalidate); + self.queue_for_vendor(out, Act::CreateResponse); + } + + // ---------------------------------------------------------------------------------------- + // The vendor's side + // ---------------------------------------------------------------------------------------- + + /// The provider is connected (the deferred setup went out); `rates` are its PCM rates. + pub(super) fn connected(&mut self, rates: Option<(u32, u32)>) -> Vec { + if let Some((input, output)) = rates { + self.input_rate = input; + self.output_rate = output; + } + if self.vendor.awaits_ready { + Vec::new() + } else { + self.became_ready() + } + } + + pub(super) fn awaits_ready(&self) -> bool { + self.vendor.awaits_ready && !self.ready + } + + fn became_ready(&mut self) -> Vec { + let mut out = Vec::new(); + if self.ready { + return out; + } + self.ready = true; + for _ in std::mem::take(&mut self.owed_updates) { + out.push(Act::Client(self.session_updated())); + } + out.extend(self.held.drain(..)); + out + } + + /// 24 kHz client PCM → the vendor's rate (streaming: the filter state carries across chunks). + pub(super) fn audio_for_vendor(&mut self, pcm: &[u8]) -> Bytes { + match resample_pcm16(&mut self.to_vendor, pcm, GA_RATE, self.input_rate) { + Some(v) => Bytes::from(v), + None => Bytes::copy_from_slice(pcm), + } + } + + /// The resampler's buffered tail at a turn boundary (a commit). + pub(super) fn audio_tail_for_vendor(&mut self) -> Option { + flush_pcm16(&mut self.to_vendor) + .filter(|t| !t.is_empty()) + .map(Bytes::from) + } + + fn user_item(&mut self, out: &mut Vec) -> String { + if let Some(id) = &self.user.item_id { + return id.clone(); + } + let id = new_id("item"); + out.push(Act::Client(ev!("conversation.item.added", { + "previous_item_id": self.last_item_id, + "item": {"id": id, "object": "realtime.item", "type": "message", "role": "user", + "status": "in_progress", "content": [{"type": "input_audio", "transcript": Value::Null}]} + }))); + self.last_item_id = Some(id.clone()); + self.user.item_id = Some(id.clone()); + id + } + + /// Close the user's transcription (a vendor that never marks it final: at the response). + fn complete_user(&mut self, out: &mut Vec, final_text: Option) { + let Some(id) = self.user.item_id.take() else { + return; + }; + let transcript = final_text + .filter(|t| !t.is_empty()) + .unwrap_or_else(|| std::mem::take(&mut self.user.transcript)); + self.user.transcript.clear(); + out.push(Act::Client( + ev!("conversation.item.input_audio_transcription.completed", { + "item_id": id, "content_index": 0, "transcript": transcript + }), + )); + } + + fn start_response(&mut self, out: &mut Vec) { + if self.response.is_some() { + return; + } + // The user spoke before this answer: their transcription is done. + let pending_user = self.user.item_id.is_some(); + if pending_user && !self.user.transcript.is_empty() { + self.complete_user(out, None); + } + let id = new_id("resp"); + out.push(Act::Client(ev!("response.created", { + "response": {"id": id, "object": "realtime.response", "status": "in_progress", + "output": [], "usage": Value::Null} + }))); + self.response = Some(Response { + id, + item_id: None, + output_index: 0, + output: Vec::new(), + transcript: String::new(), + interim_since_final: String::new(), + usage: RealtimeUsage::default(), + usage_seen: false, + cancelled: false, + }); + } + + /// The assistant message item of the response in flight, announced on first output. + fn message_item( + &mut self, + out: &mut Vec, + vendor_id: Option<&str>, + ) -> Option<(String, String)> { + self.start_response(out); + let last_item = self.last_item_id.clone(); + let r = self.response.as_mut()?; + if r.cancelled { + return None; + } + if let Some(id) = &r.item_id { + return Some((r.id.clone(), id.clone())); + } + let item_id = vendor_id + .map(str::to_string) + .unwrap_or_else(|| new_id("item")); + r.item_id = Some(item_id.clone()); + let item = json!({"id": item_id, "object": "realtime.item", "type": "message", + "role": "assistant", "status": "in_progress", "content": []}); + out.push(Act::Client(ev!("response.output_item.added", { + "response_id": r.id, "output_index": r.output_index, "item": item + }))); + out.push(Act::Client(ev!("conversation.item.added", { + "previous_item_id": last_item, "item": item + }))); + out.push(Act::Client(ev!("response.content_part.added", { + "response_id": r.id, "item_id": item_id, "output_index": r.output_index, + "content_index": 0, "part": {"type": "audio", "transcript": ""} + }))); + let rid = r.id.clone(); + self.last_item_id = Some(item_id.clone()); + Some((rid, item_id)) + } + + /// Finish the response in flight: its item events, `response.done` with the gateway's + /// `usage`, and — when the vendor reported usage — its one billed record. + fn finish_response(&mut self, status: &'static str, out: &mut Vec) { + let Some(mut r) = self.response.take() else { + return; + }; + let status = if r.cancelled { "cancelled" } else { status }; + if let Some(item_id) = r.item_id.clone() { + if let Some(tail) = flush_pcm16(&mut self.to_client).filter(|t| !t.is_empty()) { + if !r.cancelled { + out.push(Act::Client(ev!("response.output_audio.delta", { + "response_id": r.id, "item_id": item_id, "output_index": r.output_index, + "content_index": 0, "delta": BASE64_STANDARD.encode(tail) + }))); + } + } + let item_status = if status == "completed" { + "completed" + } else { + "incomplete" + }; + let item = json!({"id": item_id, "object": "realtime.item", "type": "message", + "role": "assistant", "status": item_status, + "content": [{"type": "output_audio", "transcript": r.transcript}]}); + for (kind, extra) in [ + ("response.output_audio.done", json!({})), + ( + "response.output_audio_transcript.done", + json!({"transcript": r.transcript}), + ), + ( + "response.content_part.done", + json!({"part": {"type": "audio", "transcript": r.transcript}}), + ), + ] { + let mut m = Map::new(); + m.insert("response_id".into(), json!(r.id)); + m.insert("item_id".into(), json!(item_id)); + m.insert("output_index".into(), json!(r.output_index)); + m.insert("content_index".into(), json!(0)); + if let Value::Object(x) = extra { + m.extend(x); + } + out.push(Act::Client(event(kind, Value::Object(m)))); + } + out.push(Act::Client(ev!("response.output_item.done", { + "response_id": r.id, "output_index": r.output_index, "item": item + }))); + out.push(Act::Client(ev!("conversation.item.done", {"item": item}))); + r.output.push(item); + } + let usage = r.usage_seen.then(|| ga_usage(&r.usage)); + out.push(Act::Client(ev!("response.done", { + "response": {"id": r.id, "object": "realtime.response", "status": status, + "output": r.output, "usage": usage} + }))); + if r.usage_seen { + out.push(Act::Meter { + response_id: r.id.clone(), + status, + usage: r.usage, + transcript: (!r.transcript.is_empty()).then_some(r.transcript.clone()), + }); + } + self.last_response_id = Some(r.id); + } + + /// Assistant text. An interim chunk is a delta. A FINAL one is either the next delta + /// (Gemini marks its last chunk final) or a restatement of the interim text already sent + /// (Nova Sonic's FINAL block after its SPECULATIVE one) — only its unsent part is new. + fn assistant_text(&mut self, text: &str, is_final: bool, out: &mut Vec) { + let Some((rid, item_id)) = self.message_item(out, None) else { + return; + }; + let Some(r) = self.response.as_mut() else { + return; + }; + let delta = if !is_final { + r.interim_since_final.push_str(text); + text.to_string() + } else { + let seen = std::mem::take(&mut r.interim_since_final); + if seen.is_empty() { + text.to_string() + } else if let Some(rest) = text.strip_prefix(seen.as_str()) { + rest.to_string() + } else if seen.starts_with(text) { + String::new() + } else { + text.to_string() + } + }; + if delta.is_empty() { + return; + } + r.transcript.push_str(&delta); + out.push(Act::Client(ev!("response.output_audio_transcript.delta", { + "response_id": rid, "item_id": item_id, "output_index": r.output_index, + "content_index": 0, "delta": delta + }))); + } + + fn user_text(&mut self, text: &str, is_final: bool, out: &mut Vec) { + let item_id = self.user_item(out); + if is_final { + self.complete_user(out, Some(text.to_string())); + return; + } + self.user.transcript.push_str(text); + out.push(Act::Client( + ev!("conversation.item.input_audio_transcription.delta", { + "item_id": item_id, "content_index": 0, "delta": text + }), + )); + } + + fn usage(&mut self, report: UsageReport, out: &mut Vec) { + match self.response.as_mut() { + Some(r) => { + if report.cumulative && r.usage_seen { + r.usage = report.tokens; + } else { + r.usage.add(&report.tokens); + } + r.usage_seen = true; + } + // Outside any response (a report after its `response.done`, or between turns): + // billed on its own, once. + None => out.push(Act::Meter { + response_id: self + .last_response_id + .clone() + .unwrap_or_else(|| new_id("resp")), + status: "completed", + usage: report.tokens, + transcript: None, + }), + } + } + + pub(super) fn vendor(&mut self, ev: S2sEvent) -> Vec { + let mut out = Vec::new(); + match ev { + S2sEvent::SessionReady { session_id } => { + if session_id.is_some() && self.vendor_session_id.is_none() { + self.vendor_session_id = session_id; + } + out.extend(self.became_ready()); + } + S2sEvent::Speech(SpeechEvent::Started { audio_start_ms, .. }) => { + self.user.speaking = true; + let item = self.user_item(&mut out); + out.push(Act::Client(ev!("input_audio_buffer.speech_started", { + "audio_start_ms": audio_start_ms, "item_id": item + }))); + } + S2sEvent::Speech(SpeechEvent::Stopped { audio_end_ms, .. }) => { + self.user.speaking = false; + let item = self.user_item(&mut out); + out.push(Act::Client(ev!("input_audio_buffer.speech_stopped", { + "audio_end_ms": audio_end_ms, "item_id": item + }))); + out.push(Act::Client(ev!("input_audio_buffer.committed", { + "previous_item_id": Value::Null, "item_id": item + }))); + } + S2sEvent::Transcript { + role: TranscriptRole::User, + text, + is_final, + .. + } => self.user_text(&text, is_final, &mut out), + S2sEvent::Transcript { + role: TranscriptRole::Assistant, + text, + is_final, + .. + } => self.assistant_text(&text, is_final, &mut out), + S2sEvent::Audio { data, item_id, .. } => { + let (rate, pcm) = pcm_of(&data, self.output_rate); + let pcm = match resample_pcm16(&mut self.to_client, &pcm, rate, GA_RATE) { + Some(v) => Bytes::from(v), + None => Bytes::copy_from_slice(&pcm), + }; + if let Some((rid, iid)) = self.message_item(&mut out, item_id.as_deref()) + && !pcm.is_empty() + && let Some(r) = self.response.as_ref() + { + out.push(Act::Client(ev!("response.output_audio.delta", { + "response_id": rid, "item_id": iid, "output_index": r.output_index, + "content_index": 0, "delta": BASE64_STANDARD.encode(&pcm) + }))); + } + } + S2sEvent::ItemAdded { + item_id, + role: TranscriptRole::Assistant, + } => { + let _ = self.message_item(&mut out, Some(&item_id)); + } + S2sEvent::FunctionCall(call) => { + self.start_response(&mut out); + let last_item = self.last_item_id.clone(); + if let Some(r) = self.response.as_mut() + && !r.cancelled + { + if r.item_id.is_some() { + r.output_index += 1; + } + let item_id = call.item_id.clone().unwrap_or_else(|| new_id("item")); + let item = json!({"id": item_id, "object": "realtime.item", "type": "function_call", + "status": "completed", "call_id": call.call_id, "name": call.name, + "arguments": call.arguments}); + out.push(Act::Client(ev!("response.output_item.added", { + "response_id": r.id, "output_index": r.output_index, "item": item + }))); + out.push(Act::Client(ev!("conversation.item.added", { + "previous_item_id": last_item, "item": item + }))); + out.push(Act::Client(ev!("response.function_call_arguments.done", { + "response_id": r.id, "item_id": item_id, "output_index": r.output_index, + "call_id": call.call_id, "name": call.name, "arguments": call.arguments + }))); + out.push(Act::Client(ev!("response.output_item.done", { + "response_id": r.id, "output_index": r.output_index, "item": item + }))); + out.push(Act::Client(ev!("conversation.item.done", {"item": item}))); + r.output.push(item); + self.last_item_id = Some(item_id); + } + // A GA response ends at its function call: the client runs the tool and asks + // for the next response. + self.finish_response("completed", &mut out); + } + S2sEvent::ResponseDone { .. } => { + if self.response.is_some() { + self.finish_response("completed", &mut out); + } + if self.user.item_id.is_some() && !self.user.transcript.is_empty() { + self.complete_user(&mut out, None); + } + } + S2sEvent::InterruptedByServer => { + // Barge-in: the vendor stopped its answer. A GA client stops playback on + // `speech_started`, so say so if the vendor did not. + if !self.user.speaking { + let item = self.user_item(&mut out); + out.push(Act::Client(ev!("input_audio_buffer.speech_started", { + "audio_start_ms": 0, "item_id": item + }))); + } + self.finish_response("cancelled", &mut out); + } + S2sEvent::Usage(report) => self.usage(report, &mut out), + S2sEvent::Error(e) => { + out.push(Act::Client(gateway_error( + "vendor_error", + &format!("The vendor reported an error: {e}"), + None, + None, + ))); + } + // Driver-internal, or nothing a GA client is told about. + S2sEvent::TrackPendingCall { .. } + | S2sEvent::ItemAdded { .. } + | S2sEvent::ItemDone { .. } + | S2sEvent::ResumptionHandle(_) + | S2sEvent::GoAway { .. } + | S2sEvent::SendFrame(_) + | S2sEvent::Ignore => {} + } + out + } + + /// The session is ending: bill what the response in flight already reported. + pub(super) fn close(&mut self) -> Vec { + let mut out = Vec::new(); + if let Some(r) = self.response.take() + && r.usage_seen + { + out.push(Act::Meter { + response_id: r.id, + status: "incomplete", + usage: r.usage, + transcript: (!r.transcript.is_empty()).then_some(r.transcript), + }); + } + out + } +} + +// ============================================================================================= +// The session shell: the relay's lifecycle around the translator +// ============================================================================================= + +/// What the provider's callbacks send the session. +enum VendorMsg { + Event(S2sEvent), + Reconnected(bool), +} + +struct Shell<'a> { + state: &'a AppState, + p: &'a Prepared, + plan: &'a TranslatePlan, + tr: Translator, + meter: SessionMeter, + client_tx: mpsc::Sender, + timings: Timings, + provider: Option>, + vendor_tx: mpsc::UnboundedSender, + /// Vendor-bound work waiting out a reconnect. + pending: VecDeque, + /// Set when the vendor connected and must still say it is ready. + ready_deadline: Option, + segments: Option, + last_activity: Instant, + client_missed: u32, +} + +impl Shell<'_> { + async fn to_client(&self, text: String) -> Result<(), End> { + match tokio::time::timeout( + self.timings.slow_client, + self.client_tx + .send(Outbound::Frame(Message::Text(text.into()))), + ) + .await + { + Ok(Ok(())) => Ok(()), + Ok(Err(_)) => Err(End::new("client_close", 1006)), + Err(_) => Err(End::new("client_too_slow", 1011).with_error( + "client_too_slow", + "The client did not read the session's output fast enough; audio is never dropped \ + silently, so the session is closed.", + )), + } + } + + async fn revalidate(&self) -> Result<(), End> { + if session::still_allowed(self.state, &self.p.caller.check, &self.p.endpoint_id) + .await + .is_some() + { + return Ok(()); + } + info!(endpoint_id = %self.p.endpoint_id, "realtime session revoked"); + Err(End::new("revoked", 1008).with_error( + "session_revoked", + "The credential or the deployment was revoked; the session is closed.", + )) + } + + fn upstream_error(message: impl Into) -> End { + End::new("upstream_error", 1011).with_error("upstream_error", message) + } + + /// Build the provider from the translator's config, wire its event tap, and connect. + async fn connect(&mut self) -> Result<(), End> { + if self.provider.is_some() { + return Ok(()); + } + let bedrock_http = self.state.realtime.bedrock_http_client.clone(); + let mut config = self.tr.config(); + config.max_connection = self.timings.connection_cap; + let mut provider = self.plan.provider(config, bedrock_http).map_err(|e| { + Self::upstream_error(format!("The vendor session could not be built: {e}")) + })?; + let tx = self.vendor_tx.clone(); + provider + .on_event(Arc::new(move |ev| { + let _ = tx.send(VendorMsg::Event(ev)); + Box::pin(async {}) + })) + .map_err(|e| Self::upstream_error(e.to_string()))?; + let tx = self.vendor_tx.clone(); + provider + .on_reconnection(Arc::new(move |ev| { + let _ = tx.send(VendorMsg::Reconnected(ev.success)); + Box::pin(async {}) + })) + .map_err(|e| Self::upstream_error(e.to_string()))?; + let connected = tokio::time::timeout(self.timings.connect, provider.connect()).await; + let vkey = &self.p.vkey; + match connected { + Ok(Ok(())) => { + if let Some(pol) = &self.state.policies { + pol.breakers().record_success(&self.p.endpoint_id, vkey); + } + } + other => { + let why = match other { + Ok(Err(e)) => e.to_string(), + _ => "the vendor did not answer within the connect deadline".to_string(), + }; + warn!(endpoint_id = %self.p.endpoint_id, vendor = self.plan.info.vendor, "translated vendor connect failed"); + if let Some(pol) = &self.state.policies { + let verdict = crate::core::deployment_policy::classify_message(&why, None); + pol.breakers() + .record_failure(&self.p.endpoint_id, vkey, &verdict); + } + let _ = provider.disconnect().await; + return Err(Self::upstream_error(format!( + "Could not connect to the vendor: {why}" + ))); + } + } + let rates = provider.audio_rates(); + self.provider = Some(provider); + let now = Instant::now(); + if self.plan.info.per_minute + && crate::core::realtime_cost::bills_duration(self.p.endpoint.pricing.as_ref()) + { + self.segments = Some(SegmentClock::start(now, self.timings.segment)); + } + let acts = self.tr.connected(rates); + if self.tr.awaits_ready() { + self.ready_deadline = Some(now + self.timings.hold); + } + Box::pin(self.run_acts(acts)).await + } + + /// Execute one vendor-bound act, or queue it while the vendor reconnects. + async fn vendor_call(&mut self, act: Act) -> Result<(), End> { + let ready = self.provider.as_ref().is_some_and(|p| p.is_ready()); + if !ready || !self.pending.is_empty() { + if self.pending.len() >= MAX_PENDING { + return Err(Self::upstream_error( + "The vendor did not come back in time; the session is closed.", + )); + } + self.pending.push_back(act); + return Ok(()); + } + self.send_now(act).await + } + + async fn send_now(&mut self, act: Act) -> Result<(), End> { + let Some(provider) = self.provider.as_mut() else { + return Ok(()); + }; + let result = match act { + Act::Audio(pcm) => { + let pcm = self.tr.audio_for_vendor(&pcm); + if pcm.is_empty() { + Ok(()) + } else { + provider.send_audio(pcm).await + } + } + Act::Commit => { + if let Some(tail) = self.tr.audio_tail_for_vendor() { + let _ = provider.send_audio(tail).await; + } + provider.commit_audio_buffer().await + } + Act::Clear => provider.clear_audio_buffer().await, + Act::Text(t) => provider.send_text(&t).await, + Act::CreateResponse => provider.create_response().await, + Act::Cancel => provider.cancel_response().await, + Act::ToolResult { call_id, output } => { + provider.submit_function_result(&call_id, &output).await + } + _ => Ok(()), + }; + match result { + Ok(()) => Ok(()), + Err(crate::core::realtime::RealtimeError::NotConnected) => Ok(()), + Err(e) => { + debug!(error = %e, "translated vendor call failed"); + self.to_client(gateway_error( + "vendor_error", + &format!("The vendor refused the request: {e}"), + None, + None, + )) + .await + } + } + } + + async fn flush_pending(&mut self) -> Result<(), End> { + while self.provider.as_ref().is_some_and(|p| p.is_ready()) { + let Some(act) = self.pending.pop_front() else { + break; + }; + self.send_now(act).await?; + } + Ok(()) + } + + async fn run_acts(&mut self, acts: Vec) -> Result<(), End> { + for act in acts { + match act { + Act::Client(text) => self.to_client(text).await?, + Act::Connect => self.connect().await?, + Act::Revalidate => { + self.revalidate().await?; + if !session::admit_response(&self.p.caller, &self.p.endpoint_id) { + self.to_client(gateway_error( + "quota_exceeded", + "The project's spend quota is exhausted.", + None, + None, + )) + .await?; + } + } + Act::Meter { + response_id, + status, + usage, + transcript, + } => self.meter.usage_turn( + Some(&response_id), + Some(status), + &usage, + transcript.as_deref(), + ), + vendor_bound => self.vendor_call(vendor_bound).await?, + } + } + Ok(()) + } + + async fn on_client_text(&mut self, raw: &str) -> Result<(), End> { + self.last_activity = Instant::now(); + let acts = self.tr.client(raw); + self.run_acts(acts).await + } + + async fn on_vendor(&mut self, msg: VendorMsg) -> Result<(), End> { + match msg { + VendorMsg::Event(ev) => { + self.last_activity = Instant::now(); + if matches!(ev, S2sEvent::SessionReady { .. }) { + self.ready_deadline = None; + if let Some(id) = match &ev { + S2sEvent::SessionReady { session_id } => session_id.clone(), + _ => None, + } { + self.meter.set_vendor_session_id(Some(id)); + } + } + let acts = self.tr.vendor(ev); + self.run_acts(acts).await?; + self.flush_pending().await + } + VendorMsg::Reconnected(true) => { + debug!("translated vendor reconnected"); + self.flush_pending().await + } + VendorMsg::Reconnected(false) => Err(Self::upstream_error( + "The connection to the vendor was lost and could not be restored.", + )), + } + } +} + +/// Run a translated session (the counterpart of the relay's `run`). +pub(super) async fn run( + state: Arc, + p: Prepared, + socket: WebSocket, + slot: Option, +) { + let _slot = slot; + let Engine::Translate(plan) = &p.engine else { + unreachable!("the translate engine runs translated deployments only"); + }; + let timings = state.realtime.timings.clone(); + let session_id = format!("sess_bud_{}", uuid::Uuid::new_v4().simple()); + let vendor = p.endpoint.vendor.clone(); + let meter = SessionMeter::start( + session_id.clone(), + session::attribution(&p), + p.endpoint.pricing.clone(), + ); + metrics::gauge!("waav_realtime_sessions_active", "vendor" => vendor.clone()).increment(1.0); + info!(session_id = %session_id, endpoint_id = %p.endpoint_id, vendor = %vendor, "realtime (translated) session opened"); + + let (client_sink, mut client_rx) = socket.split(); + let (client_tx, writer) = session::spawn_writer(client_sink, CLIENT_QUEUE); + let (vendor_tx, mut vendor_rx) = mpsc::unbounded_channel(); + let now = Instant::now(); + let SessionLimits { + max_len, + idle, + max_at, + warn_at, + } = SessionLimits::new(&p, &timings, now); + + let mut shell = Shell { + state: &state, + p: &p, + plan, + tr: Translator::new( + plan.info, + p.endpoint_name.clone(), + p.rules.clone(), + plan.base.clone(), + &session_id, + ), + meter, + client_tx: client_tx.clone(), + timings: timings.clone(), + provider: None, + vendor_tx, + pending: VecDeque::new(), + ready_deadline: None, + segments: None, + last_activity: now, + client_missed: 0, + }; + let mut ping = tokio::time::interval_at(now + timings.ping, timings.ping); + let mut revalidate = tokio::time::interval_at(now + timings.revalidate, timings.revalidate); + let mut warned = warn_at.is_none(); + + let first = shell.tr.session_created(); + let end: End = match shell.to_client(first).await { + Err(end) => end, + Ok(()) => loop { + let idle_at = shell.last_activity + idle; + let ready_at = shell.ready_deadline; + let segment_at = shell.segments.as_ref().map(SegmentClock::next_due); + let step: Result<(), End> = tokio::select! { + _ = state.shutdown.cancelled() => Err(End::new("drain", 1012) + .with_error("server_shutdown", "The server is restarting; reconnect.")), + msg = client_rx.next() => match msg { + None | Some(Err(_)) => Err(End::new("client_close", 1006)), + Some(Ok(Message::Close(frame))) => Err(End::new("client_close", frame.map_or(1005, |f| f.code))), + Some(Ok(Message::Pong(_))) => { shell.client_missed = 0; Ok(()) } + Some(Ok(Message::Ping(_))) => Ok(()), + Some(Ok(Message::Binary(_))) => shell.to_client(gateway_error( + "invalid_event", "Binary frames are not part of the Realtime protocol; send JSON events.", None, None, + )).await, + Some(Ok(Message::Text(t))) => shell.on_client_text(t.as_str()).await, + }, + Some(msg) = vendor_rx.recv() => shell.on_vendor(msg).await, + _ = ping.tick() => { + if shell.client_missed >= timings.max_missed_pongs { + Err(End::new("client_timeout", 1011) + .with_error("client_timeout", "The client stopped answering pings.")) + } else { + shell.client_missed += 1; + let _ = shell.client_tx.try_send(Outbound::Frame(Message::Ping(Vec::new().into()))); + Ok(()) + } + } + _ = revalidate.tick() => shell.revalidate().await, + _ = tokio::time::sleep_until(idle_at) => Err(End::new("idle", 1000) + .with_error("session_expired", format!("The session was idle for {} s.", idle.as_secs()))), + _ = tokio::time::sleep_until(warn_at.unwrap_or(max_at)), if !warned => { + warned = true; + shell.to_client(gateway_error("session_expiring", + &format!("The session reaches its maximum length in {} s.", timings.warn_before.as_secs()), + None, None)).await + } + _ = tokio::time::sleep_until(max_at) => Err(End::new("max_duration", 1000) + .with_error("session_expired", format!("The session reached its maximum length of {} s.", max_len.as_secs()))), + _ = tokio::time::sleep_until(ready_at.unwrap_or(max_at)), if ready_at.is_some() => Err(Shell::upstream_error( + "The vendor did not start the session in time.")), + _ = tokio::time::sleep_until(segment_at.unwrap_or(max_at)), if segment_at.is_some() => { + if let Some(clock) = shell.segments.as_mut() { + for secs in clock.due(Instant::now()) { + shell.meter.duration_segment(secs); + } + } + Ok(()) + } + }; + if let Err(end) = step { + break end; + } + }, + }; + + for act in shell.tr.close() { + if let Act::Meter { + response_id, + status, + usage, + transcript, + } = act + { + shell.meter.usage_turn( + Some(&response_id), + Some(status), + &usage, + transcript.as_deref(), + ); + } + } + if let Some(clock) = shell.segments.take() { + // The final partial segment: a 150 s session bills 60 + 60 + 30 (TC-XL-07). + for secs in clock.close(Instant::now()) { + shell.meter.duration_segment(secs); + } + } + if let Some(mut provider) = shell.provider.take() { + let _ = tokio::time::timeout(Duration::from_secs(2), provider.disconnect()).await; + } + let Shell { meter, tr, .. } = shell; + let mut meter = meter; + meter.set_vendor_session_id(tr.vendor_session_id()); + session::finish(meter, end, &client_tx, writer, Some(&session_id), &vendor).await; + drop(p); +} + +#[cfg(test)] +mod tests; diff --git a/gateway/src/handlers/openai_realtime/facade/tests.rs b/gateway/src/handlers/openai_realtime/facade/tests.rs new file mode 100644 index 00000000..04b25bb2 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/facade/tests.rs @@ -0,0 +1,651 @@ +//! The translator as a pure machine: GA client events in, provider calls and GA server events +//! out (FRD-023 §5.7, TC-XL-01…07). The end-to-end halves, against in-process mocks of each +//! vendor's wire protocol, are in `tests/realtime_translate.rs`. + +use super::*; +use bud_auth::RealtimePolicy; + +fn rules(policy: RealtimePolicy) -> ClientRules { + ClientRules { + policy, + session_type: "realtime".into(), + max_output_tokens: None, + } +} + +fn translator(vendor: &str) -> Translator { + Translator::new( + vendor_info(vendor).unwrap(), + "live".into(), + rules(RealtimePolicy::default()), + RealtimeConfig { + provider: vendor.into(), + api_key: "k".into(), + voice: Some("Puck".into()), + ..Default::default() + }, + "sess_bud_test", + ) +} + +/// The GA events among `acts`, parsed. +fn client_events(acts: &[Act]) -> Vec { + acts.iter() + .filter_map(|a| match a { + Act::Client(t) => serde_json::from_str(t).ok(), + _ => None, + }) + .collect() +} + +fn kinds(acts: &[Act]) -> Vec { + client_events(acts) + .iter() + .map(|e| e["type"].as_str().unwrap_or_default().to_string()) + .collect() +} + +fn refusal(acts: &[Act]) -> (String, String) { + let events = client_events(acts); + assert_eq!(events.len(), 1, "one error, got {events:?}"); + assert_eq!(events[0]["type"], "error"); + ( + events[0]["error"]["code"].as_str().unwrap().to_string(), + events[0]["error"]["param"] + .as_str() + .unwrap_or_default() + .to_string(), + ) +} + +fn ready(tr: &mut Translator) { + let _ = tr.client(&json!({"type": "session.update", "session": {}}).to_string()); + let _ = tr.connected(Some((16_000, 24_000))); + let _ = tr.vendor(S2sEvent::SessionReady { session_id: None }); +} + +fn meters(acts: &[Act]) -> Vec<(RealtimeUsage, &'static str)> { + acts.iter() + .filter_map(|a| match a { + Act::Meter { usage, status, .. } => Some((*usage, *status)), + _ => None, + }) + .collect() +} + +/// TC-XL-01 🔒 (translator half) — the setup waits for the client's first `session.update` +/// and carries it; `session.updated` follows the vendor's acceptance; afterwards a CHANGE to +/// voice or tools is refused by name, the same value again is accepted. +#[test] +fn tc_xl_01_one_setup_then_setup_time_fields_are_locked() { + let mut tr = translator("gemini"); + let tool = json!({"type": "function", "name": "lookup", "parameters": {"type": "object"}}); + let acts = tr.client( + &json!({"type": "session.update", "event_id": "c1", "session": { + "type": "realtime", "instructions": "be brief", + "audio": {"output": {"voice": "Kore"}}, "tools": [tool]}}) + .to_string(), + ); + assert_eq!( + acts, + vec![Act::Connect], + "one deferred setup, nothing else yet" + ); + let cfg = tr.config(); + assert_eq!(cfg.voice.as_deref(), Some("Kore"), "the client's voice"); + assert_eq!(cfg.instructions.as_deref(), Some("be brief")); + assert_eq!( + cfg.tools + .as_ref() + .map(|t| t[0].function.name.clone()) + .as_deref(), + Some("lookup") + ); + + assert!( + tr.connected(Some((16_000, 24_000))).is_empty(), + "Gemini: wait for setupComplete" + ); + let acts = tr.vendor(S2sEvent::SessionReady { session_id: None }); + assert_eq!(kinds(&acts), vec!["session.updated"]); + + let acts = tr.client( + &json!({"type": "session.update", "session": {"audio": {"output": {"voice": "Puck"}}}}) + .to_string(), + ); + assert_eq!( + refusal(&acts), + ( + "event_not_allowed".into(), + "session.audio.output.voice".into() + ) + ); + let other = json!({"type": "function", "name": "other"}); + let acts = + tr.client(&json!({"type": "session.update", "session": {"tools": [other]}}).to_string()); + assert_eq!( + refusal(&acts), + ("event_not_allowed".into(), "session.tools".into()) + ); + let acts = tr.client( + &json!({"type": "session.update", "session": {"instructions": "be verbose"}}).to_string(), + ); + assert_eq!(refusal(&acts).1, "session.instructions"); + + // SDKs resend the whole session: unchanged values are not a change. + let acts = tr.client( + &json!({"type": "session.update", "session": { + "instructions": "be brief", "audio": {"output": {"voice": "Kore"}}, "tools": [tool]}}) + .to_string(), + ); + assert_eq!(kinds(&acts), vec!["session.updated"]); + assert!(!acts.contains(&Act::Connect), "never a second setup"); +} + +/// Audio before the vendor is ready is held, then flushed after `session.updated`. +#[test] +fn vendor_bound_work_waits_for_the_vendor() { + let mut tr = translator("gemini"); + let pcm = BASE64_STANDARD.encode([0u8; 960]); + let acts = tr.client(&json!({"type": "input_audio_buffer.append", "audio": pcm}).to_string()); + assert_eq!( + acts, + vec![Act::Connect], + "audio first: set up with the defaults, hold it" + ); + assert!(tr.connected(Some((16_000, 24_000))).is_empty()); + let acts = tr.vendor(S2sEvent::SessionReady { session_id: None }); + assert!( + matches!(acts.as_slice(), [Act::Audio(b)] if b.len() == 960), + "{acts:?}" + ); +} + +/// FRD §5.7 / CONTRACTS C7 — what the vendor cannot do is refused by name, whatever the +/// deployment's policy allows, and the session continues. +#[test] +fn untranslatable_events_are_refused_by_name() { + let mut tr = Translator::new( + vendor_info("gemini").unwrap(), + "live".into(), + rules(RealtimePolicy { + allow_mcp_tools: Some(true), + allow_prompt_references: Some(true), + ..Default::default() + }), + RealtimeConfig::default(), + "s", + ); + ready(&mut tr); + let cases = [ + ( + json!({"type": "conversation.item.truncate", "item_id": "i", "content_index": 0, "audio_end_ms": 10}), + "conversation.item.truncate", + ), + ( + json!({"type": "conversation.item.retrieve", "item_id": "i"}), + "conversation.item.retrieve", + ), + ( + json!({"type": "conversation.item.delete", "item_id": "i"}), + "conversation.item.delete", + ), + ( + json!({"type": "output_audio_buffer.clear"}), + "output_audio_buffer.clear", + ), + (json!({"type": "future.event"}), "future.event"), + ( + json!({"type": "session.update", "session": {"tools": [{"type": "mcp", "server_url": "https://x"}]}}), + "session.tools.mcp", + ), + ( + json!({"type": "session.update", "session": {"prompt": {"id": "pmpt_1"}}}), + "session.prompt", + ), + ( + json!({"type": "session.update", "session": {"audio": {"input": {"format": {"type": "audio/pcmu"}}}}}), + "session.audio.input.format", + ), + ( + json!({"type": "conversation.item.create", "item": {"type": "message", "role": "user", + "content": [{"type": "input_image", "image_url": "data:image/png;base64,AA=="}]}}), + "item.content.input_image", + ), + ( + json!({"type": "conversation.item.create", "item": {"type": "message", "role": "user", + "content": [{"type": "input_audio", "audio": "AA=="}]}}), + "item.content.input_audio", + ), + ( + json!({"type": "response.create", "response": {"conversation": "none"}}), + "response.conversation", + ), + ( + json!({"type": "response.create", "response": {"prompt": {"id": "pmpt_1"}}}), + "response.prompt", + ), + ]; + for (event, param) in cases { + let acts = tr.client(&event.to_string()); + assert_eq!( + refusal(&acts), + ("event_not_allowed".into(), param.into()), + "{event}" + ); + } +} + +/// TC-XL-03 🔒 (translator half) — a cumulative report replaces the running total of its +/// response, so the response is billed ONCE, with the last total, and `response.done.usage` +/// says exactly that; a report outside any response is billed on its own. +#[test] +fn tc_xl_03_usage_lands_on_its_response_and_is_metered_once() { + let mut tr = translator("gemini"); + ready(&mut tr); + let report = |input_audio: u64, output_audio: u64, cumulative: bool| { + S2sEvent::Usage(UsageReport { + tokens: RealtimeUsage { + input_audio, + output_audio, + ..Default::default() + }, + seconds: None, + cumulative, + }) + }; + let mut all = tr.vendor(S2sEvent::Audio { + data: Bytes::from(vec![0u8; 480]), + item_id: None, + response_id: None, + }); + all.extend(tr.vendor(report(10, 20, true))); + all.extend(tr.vendor(report(12, 30, true))); + all.extend(tr.vendor(S2sEvent::ResponseDone { + response_id: "gemini-turn".into(), + })); + let k = kinds(&all); + assert_eq!(k.first().map(String::as_str), Some("response.created")); + assert!(k.contains(&"response.output_audio.delta".to_string())); + let done = client_events(&all) + .into_iter() + .find(|e| e["type"] == "response.done") + .unwrap(); + assert_eq!(done["response"]["status"], "completed"); + assert_eq!( + done["response"]["usage"]["input_token_details"]["audio_tokens"], + 12 + ); + assert_eq!( + done["response"]["usage"]["output_token_details"]["audio_tokens"], + 30 + ); + let billed = meters(&all); + assert_eq!(billed.len(), 1, "one record per response"); + assert_eq!( + (billed[0].0.input_audio, billed[0].0.output_audio), + (12, 30) + ); + + // Nova-style deltas add within a response. + let mut all = tr.vendor(S2sEvent::Audio { + data: Bytes::from(vec![0u8; 480]), + item_id: None, + response_id: None, + }); + all.extend(tr.vendor(report(5, 5, false))); + all.extend(tr.vendor(report(5, 7, false))); + all.extend(tr.vendor(S2sEvent::ResponseDone { + response_id: String::new(), + })); + assert_eq!( + meters(&all) + .iter() + .map(|m| (m.0.input_audio, m.0.output_audio)) + .collect::>(), + vec![(10, 12)] + ); + + // Between responses: billed alone, once. + let acts = tr.vendor(report(3, 0, false)); + assert_eq!(meters(&acts).len(), 1); + assert!(client_events(&acts).is_empty()); +} + +/// The session ends mid-response: what the vendor already reported is still billed. +#[test] +fn a_response_cut_off_by_close_keeps_its_usage() { + let mut tr = translator("nova_sonic"); + ready(&mut tr); + let _ = tr.vendor(S2sEvent::Audio { + data: Bytes::from(vec![0u8; 480]), + item_id: None, + response_id: None, + }); + let _ = tr.vendor(S2sEvent::Usage(UsageReport { + tokens: RealtimeUsage { + input_audio: 7, + ..Default::default() + }, + seconds: None, + cumulative: false, + })); + let billed = meters(&tr.close()); + assert_eq!(billed.len(), 1); + assert_eq!( + billed[0], + ( + RealtimeUsage { + input_audio: 7, + ..Default::default() + }, + "incomplete" + ) + ); +} + +fn transcript_deltas(acts: &[Act]) -> Vec { + client_events(acts) + .iter() + .filter(|e| e["type"] == "response.output_audio_transcript.delta") + .map(|e| e["delta"].as_str().unwrap().to_string()) + .collect() +} + +fn asst(text: &str, is_final: bool) -> S2sEvent { + S2sEvent::Transcript { + role: TranscriptRole::Assistant, + text: text.into(), + is_final, + item_id: None, + } +} + +/// Gemini's last transcript chunk is marked final but is a delta; Nova's FINAL block restates +/// its SPECULATIVE one. Both reach the client as each word once. +#[test] +fn assistant_transcripts_are_sent_once_each() { + let mut tr = translator("gemini"); + ready(&mut tr); + let mut acts = tr.vendor(asst("Hel", false)); + acts.extend(tr.vendor(asst("lo", false))); + acts.extend(tr.vendor(asst("!", true))); + assert_eq!(transcript_deltas(&acts), vec!["Hel", "lo", "!"]); + + let mut tr = translator("nova_sonic"); + ready(&mut tr); + let mut acts = tr.vendor(asst("Hi there.", false)); + acts.extend(tr.vendor(asst("Hi there.", true))); + acts.extend(tr.vendor(asst("Bye.", false))); + acts.extend(tr.vendor(asst("Bye.", true))); + acts.extend(tr.vendor(S2sEvent::ResponseDone { + response_id: String::new(), + })); + assert_eq!(transcript_deltas(&acts), vec!["Hi there.", "Bye."]); + let done = client_events(&acts) + .into_iter() + .find(|e| e["type"] == "response.output_audio_transcript.done") + .unwrap(); + assert_eq!(done["transcript"], "Hi there.Bye."); +} + +/// A function call ends its response (the client runs the tool, then asks for the next). +#[test] +fn a_function_call_ends_its_response() { + let mut tr = translator("gemini"); + ready(&mut tr); + let acts = tr.vendor(S2sEvent::FunctionCall( + crate::core::realtime::FunctionCallRequest { + call_id: "call_1".into(), + name: "lookup".into(), + arguments: "{\"q\":1}".into(), + item_id: None, + }, + )); + let k = kinds(&acts); + assert_eq!( + k, + vec![ + "response.created", + "response.output_item.added", + "conversation.item.added", + "response.function_call_arguments.done", + "response.output_item.done", + "conversation.item.done", + "response.done" + ] + ); + let args = &client_events(&acts)[3]; + assert_eq!(args["call_id"], "call_1"); + assert_eq!(args["arguments"], "{\"q\":1}"); + // The client's answer goes to the vendor as a tool result. + let acts = tr.client( + &json!({"type": "conversation.item.create", "item": {"type": "function_call_output", + "call_id": "call_1", "output": "{\"a\":2}"}}) + .to_string(), + ); + assert!(acts.contains(&Act::ToolResult { + call_id: "call_1".into(), + output: "{\"a\":2}".into() + })); +} + +/// `response.cancel` ends the response for the client at once; its late output is dropped. +#[test] +fn a_client_cancel_ends_the_response_and_drops_its_output() { + let mut tr = translator("gemini"); + ready(&mut tr); + let _ = tr.vendor(S2sEvent::Audio { + data: Bytes::from(vec![0u8; 480]), + item_id: None, + response_id: None, + }); + let acts = tr.client(&json!({"type": "response.cancel"}).to_string()); + assert!(acts.contains(&Act::Cancel)); + let done = client_events(&acts) + .into_iter() + .find(|e| e["type"] == "response.done") + .unwrap(); + assert_eq!(done["response"]["status"], "cancelled"); +} + +/// A barge-in the vendor reports (Gemini `interrupted`) tells a GA client to stop playing. +#[test] +fn a_vendor_interruption_is_a_speech_start_and_a_cancelled_response() { + let mut tr = translator("gemini"); + ready(&mut tr); + let _ = tr.vendor(S2sEvent::Audio { + data: Bytes::from(vec![0u8; 480]), + item_id: None, + response_id: None, + }); + let acts = tr.vendor(S2sEvent::InterruptedByServer); + let k = kinds(&acts); + assert_eq!( + k.first().map(String::as_str), + Some("conversation.item.added") + ); + assert!(k.contains(&"input_audio_buffer.speech_started".to_string())); + assert_eq!(k.last().map(String::as_str), Some("response.done")); +} + +/// Hume EVI sends a WAV container per chunk: the PCM inside it, at its own rate. +#[test] +fn wav_chunks_are_unwrapped() { + let pcm: Vec = (0..200u8).collect(); + let mut wav = Vec::new(); + wav.extend_from_slice(b"RIFF"); + wav.extend_from_slice(&(36 + pcm.len() as u32).to_le_bytes()); + wav.extend_from_slice(b"WAVEfmt "); + wav.extend_from_slice(&16u32.to_le_bytes()); + wav.extend_from_slice(&1u16.to_le_bytes()); // PCM + wav.extend_from_slice(&1u16.to_le_bytes()); // mono + wav.extend_from_slice(&48_000u32.to_le_bytes()); + wav.extend_from_slice(&96_000u32.to_le_bytes()); + wav.extend_from_slice(&2u16.to_le_bytes()); + wav.extend_from_slice(&16u16.to_le_bytes()); + wav.extend_from_slice(b"data"); + wav.extend_from_slice(&(pcm.len() as u32).to_le_bytes()); + wav.extend_from_slice(&pcm); + let (rate, body) = pcm_of(&wav, 44_100); + assert_eq!(rate, 48_000); + assert_eq!(&*body, pcm.as_slice()); + let (rate, body) = pcm_of(&pcm, 16_000); + assert_eq!((rate, body.len()), (16_000, 200), "raw PCM passes through"); +} + +/// TC-XL-02 (translator half) — 24 kHz client audio reaches the vendor at ITS rate. +#[test] +fn tc_xl_02_client_audio_is_resampled_to_the_vendor_rate() { + let mut tr = translator("gemini"); + ready(&mut tr); + // 1 s of a 440 Hz tone at 24 kHz, in 20 ms chunks. + let samples: Vec = (0..24_000) + .flat_map(|i| { + let v = (f32::sin(i as f32 * 440.0 * std::f32::consts::TAU / 24_000.0) * 8000.0) as i16; + v.to_le_bytes() + }) + .collect(); + let mut out = 0usize; + for chunk in samples.chunks(960) { + out += tr.audio_for_vendor(chunk).len(); + } + out += tr.audio_tail_for_vendor().map_or(0, |t| t.len()); + // 16 kHz: 32 000 bytes a second, within the resampler's one-chunk latency and padding. + assert!((31_000..=33_000).contains(&out), "{out}"); +} + +// --------------------------------------------------------------------------------------------- +// The plan (pre-upgrade) +// --------------------------------------------------------------------------------------------- + +fn endpoint(entry: Value) -> VoiceEndpoint { + let blob = json!({ "ep": entry }).to_string(); + bud_auth::credentials::parse_voice_blob(&blob, &bud_auth::CredentialDecryptor::disabled()) + .unwrap() + .remove("ep") + .unwrap() +} + +/// FRD-023 RT7.2 🔒 — a Nova 2 Sonic deployment without its AWS key pair or its region is +/// refused before the upgrade; nothing falls back to the gateway's AWS identity. +#[test] +fn rt7_2_nova_needs_the_deployment_key_pair_and_region() { + let entry = json!({"vendor": "nova_sonic", "endpoints": ["realtime_session"], + "model": "amazon.nova-2-sonic-v1:0", "provider_params": {"region": "us-east-1"}}); + let err = TranslatePlan::build(&endpoint(entry.clone()), None).unwrap_err(); + assert!( + matches!(err, UpstreamError::Misconfigured(ref m) if m.contains("AWS access key pair")), + "{err:?}" + ); + + let mut ep = endpoint(entry); + ep.credential_parts = Some( + [("access_key_id", "AKIDX"), ("secret_access_key", "s")] + .into_iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(), + ); + let plan = TranslatePlan::build(&ep, None).unwrap(); + assert_eq!( + plan.aws.as_ref().map(|a| a.access_key_id.as_str()), + Some("AKIDX") + ); + assert_eq!(plan.base.endpoint.as_deref(), Some("us-east-1")); + assert!( + plan.base.api_key.is_empty(), + "no API key rides a SigV4 session" + ); + + ep.provider_params.clear(); + let err = TranslatePlan::build(&ep, None).unwrap_err(); + assert!( + matches!(err, UpstreamError::Misconfigured(ref m) if m.contains("region")), + "{err:?}" + ); +} + +/// The address is the deployment's (`api_base`, http(s) converted to ws(s): F-5) and must clear +/// the SSRF validator; the defaults are the deployment's. +#[test] +fn a_plan_takes_address_credential_and_defaults_from_the_deployment() { + let mut ep = endpoint( + json!({"vendor": "gemini", "endpoints": ["realtime_session"], + "model": "gemini-3.8-live", "api_base": "https://gemini-proxy.example/ws", + "config": {"realtime": {"defaults": {"voice": "Kore", "instructions": "hi", + "turn_detection": {"type": "server_vad", "silence_duration_ms": 400}}}}}), + ); + ep.credential = Some("gkey".into()); + let settings = ep.config.realtime.clone(); + let plan = TranslatePlan::build(&ep, settings.as_ref()).unwrap(); + assert_eq!(plan.base.api_key, "gkey"); + assert_eq!(plan.base.model, "gemini-3.8-live"); + assert_eq!( + plan.base.realtime_endpoint_override.as_deref(), + Some("wss://gemini-proxy.example/ws") + ); + assert!(plan.ssrf.is_some()); + assert_eq!(plan.base.voice.as_deref(), Some("Kore")); + assert!(matches!( + plan.base.turn_detection, + Some(TurnDetectionConfig::ServerVad { + silence_duration_ms: Some(400), + .. + }) + )); + // Debug never prints the key. + assert!(!format!("{plan:?}").contains("gkey")); + + ep.credential = None; + assert_eq!( + TranslatePlan::build(&ep, settings.as_ref()).unwrap_err(), + UpstreamError::MissingCredential + ); +} + +/// RT7 vendors are speech-to-speech only (CONTRACTS C7). +#[test] +fn a_transcription_entry_on_a_translate_vendor_is_refused() { + let mut ep = endpoint( + json!({"vendor": "deepgram_voice_agent", "endpoints": ["realtime_session"], + "config": {"realtime": {"session_type": "transcription"}}}), + ); + ep.credential = Some("k".into()); + let settings = ep.config.realtime.clone(); + assert!(matches!( + TranslatePlan::build(&ep, settings.as_ref()), + Err(UpstreamError::Misconfigured(_)) + )); +} + +/// An ElevenLabs agent needs its agent id (`model`): refused before the upgrade. +#[test] +fn a_plan_the_provider_cannot_build_is_refused_up_front() { + let mut ep = + endpoint(json!({"vendor": "elevenlabs_convai", "endpoints": ["realtime_session"]})); + ep.credential = Some("xi".into()); + assert!(matches!( + TranslatePlan::build(&ep, None), + Err(UpstreamError::Misconfigured(_)) + )); + ep.model = Some("agent_123".into()); + assert!(TranslatePlan::build(&ep, None).is_ok()); +} + +#[test] +fn the_translate_vendors_are_c7s() { + for v in [ + "gemini", + "nova_sonic", + "deepgram_voice_agent", + "elevenlabs_convai", + "hume_evi", + ] { + assert!(is_translate_vendor(v), "{v}"); + } + for v in ["openai", "azure_openai", "grok", "deepgram", "elevenlabs"] { + assert!(!is_translate_vendor(v), "{v}"); + } + assert!(vendor_info("deepgram_voice_agent").unwrap().per_minute); + assert!(!vendor_info("gemini").unwrap().per_minute); +} diff --git a/gateway/src/handlers/openai_realtime/metering.rs b/gateway/src/handlers/openai_realtime/metering.rs index 18059b7a..632746be 100644 --- a/gateway/src/handlers/openai_realtime/metering.rs +++ b/gateway/src/handlers/openai_realtime/metering.rs @@ -115,6 +115,74 @@ fn response_transcript(event: &Value) -> Option { (!parts.is_empty()).then(|| parts.join(" ")) } +/// OpenAI GA `response.usage` for a usage the gateway computed itself (the translate engine's +/// `response.done`, FRD-023 §5.7). Cached counts are reported as the subset they are. +pub fn ga_usage(u: &RealtimeUsage) -> Value { + let input = u.input_text + u.input_audio + u.input_image; + let output = u.output_text + u.output_audio; + serde_json::json!({ + "total_tokens": input + output, + "input_tokens": input, + "output_tokens": output, + "input_token_details": { + "text_tokens": u.input_text, + "audio_tokens": u.input_audio, + "image_tokens": u.input_image, + "cached_tokens": u.cached_text + u.cached_audio + u.cached_image, + "cached_tokens_details": { + "text_tokens": u.cached_text, + "audio_tokens": u.cached_audio, + "image_tokens": u.cached_image, + }, + }, + "output_token_details": { + "text_tokens": u.output_text, + "audio_tokens": u.output_audio, + }, + }) +} + +/// Duration segments under a minute/second price (D-9): a billed record every `len`, and the +/// partial remainder at close — a 150 s session bills 60 + 60 + 30 (TC-XL-07). Shared by the +/// relay and the translate engine. Pure over the instants it is given. +#[derive(Debug, Clone)] +pub struct SegmentClock { + len: std::time::Duration, + /// Where the unbilled stretch begins. + last: tokio::time::Instant, +} + +impl SegmentClock { + pub fn start(now: tokio::time::Instant, len: std::time::Duration) -> Self { + Self { len, last: now } + } + + /// When the next full segment falls due. + pub fn next_due(&self) -> tokio::time::Instant { + self.last + self.len + } + + /// The full segments that fell due by `now`, in seconds each. + pub fn due(&mut self, now: tokio::time::Instant) -> Vec { + let mut out = Vec::new(); + while self.len > std::time::Duration::ZERO && now >= self.last + self.len { + self.last += self.len; + out.push(self.len.as_secs_f64()); + } + out + } + + /// The partial segment at close: every full one first, then the remainder. + pub fn close(mut self, now: tokio::time::Instant) -> Vec { + let mut out = self.due(now); + let rest = now.saturating_duration_since(self.last).as_secs_f64(); + if rest > 0.0 { + out.push(rest); + } + out + } +} + /// The meter for one session. pub struct SessionMeter { session_span: Span, @@ -216,32 +284,44 @@ impl SessionMeter { .and_then(|r| r.get("usage")) .and_then(RealtimeUsage::from_openai) .unwrap_or_default(); - let cost = realtime_response_cost(self.pricing.as_ref(), &usage); - let span = self.open_turn("response"); - record_text( - &span, - rt::RESPONSE_ID, + self.usage_turn( response.and_then(|r| r.get("id")).and_then(Value::as_str), + response + .and_then(|r| r.get("status")) + .and_then(Value::as_str), + &usage, + response_transcript(event).as_deref(), ); - let status = response - .and_then(|r| r.get("status")) - .and_then(Value::as_str); + } + + /// One billed response: a relayed `response.done`, or a translated vendor's usage report + /// (FRD-023 §5.7 — Gemini `usageMetadata`, Nova Sonic `usageEvent`), metered once. + pub fn usage_turn( + &mut self, + response_id: Option<&str>, + status: Option<&str>, + usage: &RealtimeUsage, + transcript: Option<&str>, + ) { + let cost = realtime_response_cost(self.pricing.as_ref(), usage); + let span = self.open_turn("response"); + record_text(&span, rt::RESPONSE_ID, response_id); record_text(&span, rt::RESPONSE_STATUS, status); if status == Some("failed") { span.record("otel.status_code", "ERROR"); span.record(turn::ERROR_TYPE, "vendor_error"); } - record_usage(&span, &usage); + record_usage(&span, usage); record_cost(&span, &cost); if self.capture - && let Some(t) = response_transcript(event) + && let Some(t) = transcript.filter(|t| !t.is_empty()) { span.record( turn::TRANSCRIPT, - crate::observability::trace_redact::sanitize_body(&t).as_str(), + crate::observability::trace_redact::sanitize_body(t).as_str(), ); } - self.totals.add(&usage); + self.totals.add(usage); self.account(&cost); } @@ -334,3 +414,59 @@ impl SessionMeter { // Dropping `self` ends the span. } } + +#[cfg(test)] +mod tests { + use super::*; + use std::time::Duration; + + /// TC-XL-07 🔒 — a 150 s per-minute session bills 60 + 60 + 30, never a rounded-up minute + /// and never one lump at close. + #[test] + fn tc_xl_07_a_150_s_session_bills_60_60_30() { + let t0 = tokio::time::Instant::now(); + let mut clock = SegmentClock::start(t0, Duration::from_secs(60)); + assert!(clock.due(t0 + Duration::from_secs(59)).is_empty()); + assert_eq!(clock.due(t0 + Duration::from_secs(60)), vec![60.0]); + assert_eq!(clock.next_due(), t0 + Duration::from_secs(120)); + assert_eq!(clock.due(t0 + Duration::from_secs(121)), vec![60.0]); + assert_eq!(clock.close(t0 + Duration::from_secs(150)), vec![30.0]); + } + + /// A timer that fires late still bills every elapsed segment, once. + #[test] + fn late_ticks_bill_each_elapsed_segment_once() { + let t0 = tokio::time::Instant::now(); + let mut clock = SegmentClock::start(t0, Duration::from_secs(60)); + assert_eq!( + clock.due(t0 + Duration::from_secs(185)), + vec![60.0, 60.0, 60.0] + ); + assert_eq!(clock.close(t0 + Duration::from_secs(185)), vec![5.0]); + let clock = SegmentClock::start(t0, Duration::from_secs(60)); + assert_eq!( + clock.close(t0 + Duration::from_secs(150)), + vec![60.0, 60.0, 30.0], + "a session that closes before any tick is still billed in segments" + ); + } + + #[test] + fn ga_usage_reports_cached_as_the_subset_it_is() { + let u = RealtimeUsage { + input_text: 119, + cached_text: 64, + input_audio: 13, + output_text: 30, + output_audio: 91, + ..Default::default() + }; + let v = ga_usage(&u); + assert_eq!(v["input_tokens"], 132); + assert_eq!(v["output_tokens"], 121); + assert_eq!(v["total_tokens"], 253); + assert_eq!(v["input_token_details"]["cached_tokens"], 64); + // Round-trips through the relay's reader: the same record either way. + assert_eq!(RealtimeUsage::from_openai(&v), Some(u)); + } +} diff --git a/gateway/src/handlers/openai_realtime/mod.rs b/gateway/src/handlers/openai_realtime/mod.rs index 9993283c..bad58a11 100644 --- a/gateway/src/handlers/openai_realtime/mod.rs +++ b/gateway/src/handlers/openai_realtime/mod.rs @@ -11,9 +11,12 @@ //! * [`policy`] — what crosses the relay, per event (§5.5, §5.6). //! * [`metering`] — a `voice.turn` per billed record, a `voice.session` per session (§5.10). //! * [`session`] — the session engine: admission, relay, timers, revalidation, teardown (§5.4). +//! * [`facade`] — the translate engine: GA over WaaV's native providers for vendors without a +//! GA surface (Gemini Live, Nova 2 Sonic, the per-minute voice agents; §5.7, RT7). //! * [`client_secrets`] — `POST /v1/realtime/client_secrets`, the `ek_bud_` mint (§5.8). pub mod client_secrets; +pub mod facade; pub mod handshake; pub mod metering; pub mod policy; diff --git a/gateway/src/handlers/openai_realtime/session.rs b/gateway/src/handlers/openai_realtime/session.rs index 08bdb8ba..1f1a90a7 100644 --- a/gateway/src/handlers/openai_realtime/session.rs +++ b/gateway/src/handlers/openai_realtime/session.rs @@ -30,7 +30,7 @@ use crate::state::AppState; use super::handshake::{ self, Credential, HandshakeError, MAX_MESSAGE_BYTES, REALTIME_CAPABILITY, SUBPROTOCOL, }; -use super::metering::{Attribution, SessionMeter}; +use super::metering::{Attribution, SegmentClock, SessionMeter}; use super::policy::{self, ClientOutcome, ClientRules, Tap, VendorOutcome}; use super::upstream::{self, UpstreamRequest}; @@ -51,6 +51,9 @@ pub struct Timings { pub warn_before: Duration, pub segment: Duration, pub upstream_send: Duration, + /// Replaces every translated vendor's connection cap (Nova Sonic: 8 min). `None`: each + /// vendor's own. Tests shorten it; production leaves it unset. + pub connection_cap: Option, } impl Default for Timings { @@ -67,6 +70,7 @@ impl Default for Timings { warn_before: Duration::from_secs(60), segment: Duration::from_secs(60), upstream_send: Duration::from_secs(10), + connection_cap: None, } } } @@ -101,6 +105,9 @@ pub struct RealtimeRuntime { pub timings: Timings, /// `None`: client secrets are not configured (the mint route answers 501, `ek_bud_` refused). pub client_secret_keys: Option, + /// The HTTP client Nova Sonic's Bedrock streams dial with. `None` (production): the AWS + /// SDK's own. In-process tests set a connector that speaks the Bedrock event stream. + pub bedrock_http_client: Option, } impl RealtimeRuntime { @@ -110,6 +117,7 @@ impl RealtimeRuntime { Ok(Self { timings: Timings::from_env(), client_secret_keys: ephemeral::ClientSecretKeys::from_env()?, + bedrock_http_client: None, }) } } @@ -296,6 +304,14 @@ async fn authenticate_client_secret( Ok((caller, claims.ep, alias)) } +/// The two engines behind `/v1/realtime` (FRD-023 D-2, §5.7). +pub enum Engine { + /// The vendor speaks OpenAI Realtime GA: frames are relayed after policy. + Relay(UpstreamRequest), + /// The vendor has its own protocol: WaaV's native provider, behind the GA facade. + Translate(Box), +} + /// Everything a session needs, decided before the upgrade. pub struct Prepared { pub caller: Caller, @@ -305,7 +321,8 @@ pub struct Prepared { pub alias: Option, pub settings: Option, pub rules: ClientRules, - pub upstream: UpstreamRequest, + /// How the vendor is reached: relayed (it speaks GA) or translated (FRD-023 §5.7). + pub engine: Engine, pub vkey: String, /// Held for the session: the deployment's concurrency slot (D-8). pub admission: Admission, @@ -375,18 +392,32 @@ pub async fn prepare( .as_ref() .is_some_and(RealtimeSettings::is_transcription); - // Build before admitting: a deployment the relay cannot serve must not take a slot. - let upstream_req = upstream::build(&endpoint, transcription).map_err(|e| match e { + // Build before admitting: a deployment the gateway cannot serve must not take a slot. + let refused = |e: upstream::UpstreamError| match e { upstream::UpstreamError::UnsupportedVendor(_) => HandshakeError::new( StatusCode::NOT_IMPLEMENTED, "unsupported_vendor", e.to_string(), ), + upstream::UpstreamError::Misconfigured(_) => HandshakeError::new( + StatusCode::BAD_GATEWAY, + "deployment_misconfigured", + e.to_string(), + ), other => HandshakeError::new(StatusCode::BAD_GATEWAY, "upstream_error", other.to_string()), - })?; - upstream::validate(&upstream_req).await.map_err(|e| { - HandshakeError::new(StatusCode::BAD_GATEWAY, "upstream_error", e.to_string()) - })?; + }; + let engine = if super::facade::is_translate_vendor(&endpoint.vendor) { + let plan = + super::facade::TranslatePlan::build(&endpoint, settings.as_ref()).map_err(refused)?; + plan.validate().await.map_err(refused)?; + Engine::Translate(Box::new(plan)) + } else { + let req = upstream::build(&endpoint, transcription).map_err(refused)?; + upstream::validate(&req).await.map_err(|e| { + HandshakeError::new(StatusCode::BAD_GATEWAY, "upstream_error", e.to_string()) + })?; + Engine::Relay(req) + }; debug!(endpoint_id = %endpoint_id, "realtime upstream request validated"); let vkey = vendor_key(&endpoint.vendor, endpoint.api_base.as_deref()); @@ -415,7 +446,7 @@ pub async fn prepare( endpoint, alias, settings, - upstream: upstream_req, + engine, vkey, admission, }) @@ -448,7 +479,12 @@ pub async fn realtime_ws_handler( ws.protocols([SUBPROTOCOL]) .max_message_size(MAX_MESSAGE_BYTES) .max_frame_size(MAX_MESSAGE_BYTES) - .on_upgrade(move |socket| run(state, prepared, socket, slot)) + .on_upgrade(move |socket| async move { + match prepared.engine { + Engine::Relay(_) => run(state, prepared, socket, slot).await, + Engine::Translate(_) => super::facade::run(state, prepared, socket, slot).await, + } + }) } /// How a session ended. @@ -461,7 +497,7 @@ pub struct End { } impl End { - fn new(reason: &'static str, close_code: u16) -> Self { + pub(super) fn new(reason: &'static str, close_code: u16) -> Self { Self { reason, close_code, @@ -469,13 +505,13 @@ impl End { } } - fn with_error(mut self, code: &'static str, message: impl Into) -> Self { + pub(super) fn with_error(mut self, code: &'static str, message: impl Into) -> Self { self.error = Some((code, message.into())); self } } -enum Outbound { +pub(super) enum Outbound { Frame(Message), Close { error: Option, @@ -484,11 +520,11 @@ enum Outbound { }, } -fn next_event_id() -> String { +pub(super) fn next_event_id() -> String { format!("evt_bud_{}", uuid::Uuid::new_v4().simple()) } -fn gateway_error( +pub(super) fn gateway_error( code: &str, message: &str, param: Option<&str>, @@ -510,7 +546,7 @@ fn gateway_error( /// The client-bound writer: a bounded queue drained by its own task, so a slow client applies /// backpressure without the relay dropping a frame (FR-EVT-4). -fn spawn_writer( +pub(super) fn spawn_writer( mut sink: futures_util::stream::SplitSink, capacity: usize, ) -> (mpsc::Sender, tokio::task::JoinHandle<()>) { @@ -547,7 +583,7 @@ fn spawn_writer( } /// The client-bound queue: 2 s of 24 kHz audio at 20 ms frames, plus headroom for events. -const CLIENT_QUEUE: usize = 256; +pub(super) const CLIENT_QUEUE: usize = 256; /// Client frames held while the vendor applies the defaults. const MAX_HELD: usize = 4096; @@ -805,12 +841,9 @@ impl Relay<'_> { } } -async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option) { - let _slot = slot; - let timings = state.realtime.timings.clone(); - let session_id = format!("sess_bud_{}", uuid::Uuid::new_v4().simple()); - let vendor = p.endpoint.vendor.clone(); - let attribution = Attribution { +/// Who a session bills (CONTRACTS C2), shared by both engines. +pub(super) fn attribution(p: &Prepared) -> Attribution { + Attribution { project_id: p .alias .as_ref() @@ -822,15 +855,58 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< user_id: p.caller.principal.user_id.clone(), api_key_project_id: p.caller.principal.project_id.clone(), endpoint_name: p.endpoint_name.clone(), - vendor: vendor.clone(), + vendor: p.endpoint.vendor.clone(), model: p.endpoint.model.clone(), session_type: p .settings .as_ref() .and_then(|s| s.session_type.clone()) .unwrap_or_else(|| "realtime".into()), - }; - let meter = SessionMeter::start(session_id.clone(), attribution, p.endpoint.pricing.clone()); + } +} + +/// A session's length limits (D-13): the deployment's, capped by the gateway's. +pub(super) struct SessionLimits { + pub max_len: Duration, + pub idle: Duration, + pub max_at: Instant, + pub warn_at: Option, +} + +impl SessionLimits { + pub(super) fn new(p: &Prepared, timings: &Timings, now: Instant) -> Self { + let limits = p + .settings + .as_ref() + .and_then(|s| s.limits.clone()) + .unwrap_or_default(); + let max_len = limits + .max_session_seconds + .map(Duration::from_secs) + .map_or(timings.max_session, |d| d.min(timings.max_session)); + let idle = limits + .idle_timeout_seconds + .map(Duration::from_secs) + .unwrap_or(timings.default_idle); + Self { + max_len, + idle, + max_at: now + max_len, + warn_at: max_len.checked_sub(timings.warn_before).map(|d| now + d), + } + } +} + +async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option) { + let _slot = slot; + let timings = state.realtime.timings.clone(); + let session_id = format!("sess_bud_{}", uuid::Uuid::new_v4().simple()); + let vendor = p.endpoint.vendor.clone(); + let meter = SessionMeter::start( + session_id.clone(), + attribution(&p), + p.endpoint.pricing.clone(), + ); metrics::gauge!("waav_realtime_sessions_active", "vendor" => vendor.clone()).increment(1.0); info!(session_id = %session_id, endpoint_id = %p.endpoint_id, vendor = %vendor, "realtime session opened"); @@ -838,7 +914,10 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< let (client_sink, mut client_rx) = socket.split(); let (client_tx, writer) = spawn_writer(client_sink, CLIENT_QUEUE); - let upstream_socket = match upstream::connect(&p.upstream, timings.connect).await { + let Engine::Relay(upstream_req) = &p.engine else { + unreachable!("the relay runs relayed deployments only"); + }; + let upstream_socket = match upstream::connect(upstream_req, timings.connect).await { Ok(s) => { if let Some(pol) = &state.policies { pol.breakers().record_success(&p.endpoint_id, &p.vkey); @@ -862,21 +941,12 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< let (up_tx, mut up_rx) = upstream_socket.split(); let now = Instant::now(); - let limits = p - .settings - .as_ref() - .and_then(|s| s.limits.clone()) - .unwrap_or_default(); - let max_len = limits - .max_session_seconds - .map(Duration::from_secs) - .map_or(timings.max_session, |d| d.min(timings.max_session)); - let idle = limits - .idle_timeout_seconds - .map(Duration::from_secs) - .unwrap_or(timings.default_idle); - let max_at = now + max_len; - let warn_at = max_len.checked_sub(timings.warn_before).map(|d| now + d); + let SessionLimits { + max_len, + idle, + max_at, + warn_at, + } = SessionLimits::new(&p, &timings, now); let bills_duration = crate::core::realtime_cost::bills_duration(p.endpoint.pricing.as_ref()); let mut relay = Relay { @@ -898,8 +968,7 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< }; let mut ping = tokio::time::interval_at(now + timings.ping, timings.ping); let mut revalidate = tokio::time::interval_at(now + timings.revalidate, timings.revalidate); - let mut segment = tokio::time::interval_at(now + timings.segment, timings.segment); - let mut last_segment = now; + let mut segments = bills_duration.then(|| SegmentClock::start(now, timings.segment)); let mut warned = warn_at.is_none(); let end: End = loop { @@ -962,9 +1031,13 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< .with_error("session_expired", format!("The session reached its maximum length of {} s.", max_len.as_secs()))), _ = tokio::time::sleep_until(hold_at.unwrap_or(max_at)), if hold_at.is_some() => Err(End::new("upstream_error", 1011) .with_error("upstream_error", "The vendor did not start or configure the session in time.")), - _ = segment.tick(), if bills_duration => { - relay.meter.duration_segment(timings.segment.as_secs_f64()); - last_segment = Instant::now(); + _ = tokio::time::sleep_until(segments.as_ref().map_or(max_at, SegmentClock::next_due)), + if segments.is_some() => { + if let Some(clock) = segments.as_mut() { + for secs in clock.due(Instant::now()) { + relay.meter.duration_segment(secs); + } + } Ok(()) } }; @@ -973,11 +1046,11 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< } }; - if bills_duration { + if let Some(clock) = segments { // The final partial segment: a 150 s session bills 60 + 60 + 30 (TC-XL-07). - relay - .meter - .duration_segment(last_segment.elapsed().as_secs_f64()); + for secs in clock.close(Instant::now()) { + relay.meter.duration_segment(secs); + } } let Relay { meter, mut up_tx, .. @@ -997,7 +1070,7 @@ async fn run(state: Arc, p: Prepared, socket: WebSocket, slot: Option< } /// The common teardown: the error event and close, the session record, the metrics. -async fn finish( +pub(super) async fn finish( meter: SessionMeter, end: End, client_tx: &mpsc::Sender, diff --git a/gateway/src/handlers/openai_realtime/upstream.rs b/gateway/src/handlers/openai_realtime/upstream.rs index ff5f376d..f9af4882 100644 --- a/gateway/src/handlers/openai_realtime/upstream.rs +++ b/gateway/src/handlers/openai_realtime/upstream.rs @@ -40,6 +40,9 @@ pub enum UpstreamError { Connect(String), /// No answer within the connect deadline. Timeout, + /// The deployment's entry cannot be served as published (a missing region, a key of the + /// wrong shape, a session type the vendor does not have). + Misconfigured(String), } impl std::fmt::Display for UpstreamError { @@ -57,6 +60,7 @@ impl std::fmt::Display for UpstreamError { Self::InvalidApiBase(why) => write!(f, "the deployment's api_base was refused: {why}"), Self::Connect(why) => write!(f, "could not connect to the vendor: {why}"), Self::Timeout => write!(f, "the vendor did not answer within the connect deadline"), + Self::Misconfigured(why) => write!(f, "the deployment is misconfigured: {why}"), } } } @@ -83,7 +87,7 @@ impl std::fmt::Debug for UpstreamRequest { } /// `https://…` → `wss://…`, `http://…` → `ws://…`; `ws(s)` kept. Trailing slashes trimmed. -fn to_ws_base(base: &str) -> Result { +pub(super) fn to_ws_base(base: &str) -> Result { let base = base.trim().trim_end_matches('/'); let converted = if let Some(rest) = base.strip_prefix("https://") { format!("wss://{rest}") diff --git a/gateway/tests/openai_realtime_relay.rs b/gateway/tests/openai_realtime_relay.rs index f14bdc75..1d81f99e 100644 --- a/gateway/tests/openai_realtime_relay.rs +++ b/gateway/tests/openai_realtime_relay.rs @@ -493,6 +493,7 @@ fn fast_timings() -> Timings { warn_before: Duration::from_secs(1), segment: Duration::from_secs(1), upstream_send: Duration::from_secs(2), + connection_cap: None, } } @@ -633,6 +634,7 @@ async fn gateway(setup: Setup) -> Gateway { client_secret_keys: setup .client_secret_keys .map(|k| waav_gateway::auth::ephemeral::ClientSecretKeys::parse(k).unwrap()), + bedrock_http_client: None, }); } let app = waav_gateway::routes::openai_realtime::create_openai_realtime_router() @@ -1856,6 +1858,63 @@ async fn tc_xl_06_xai_relay_quirks() { } } +/// CONTRACTS C7 — an xAI deployment priced PER MINUTE (xAI bills its Voice Agent API by the +/// minute) is billed by duration segments; its responses carry tokens and no cost. +#[tokio::test] +async fn xai_priced_per_minute_bills_duration_segments() { + let cap = Capture::install(); + let vendor = MockVendor::start(Behaviour { + xai: true, + ..Default::default() + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "grok-min", + "a1a1a1a1-0000-4000-8000-0000000000a7", + rt_entry( + &vendor, + json!({"vendor": "grok", "model": "grok-voice-2", + "pricing": {"unit": "minute", "cost_per_unit": 0.08, "per_units": 1}}), + ), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "grok-min").await; + until_type(&mut c, "session.created").await; + send(&mut c, json!({"type": "response.create"})).await; + until_type(&mut c, "response.done").await; + keep_alive(&mut c, Duration::from_millis(1500)).await; + c.close(None).await.unwrap(); + let _ = until_close(&mut c).await; + let session = cap.wait_for("voice.session", 1).await.remove(0); + let turns: Vec = cap + .spans() + .into_iter() + .filter(|s| s.name == "voice.turn") + .collect(); + let segments: Vec<&SpanData> = turns + .iter() + .filter(|t| text(t, "bud.voice.rt.component").as_deref() == Some("duration_segment")) + .collect(); + assert!(segments.len() >= 2, "a full segment and the remainder"); + for s in &segments { + let secs = number(s, "bud.voice.billed_seconds").unwrap(); + assert!((number(s, "bud.voice.cost").unwrap() - secs / 60.0 * 0.08).abs() < 1e-12); + } + let response = turns + .iter() + .find(|t| text(t, "bud.voice.rt.component").as_deref() == Some("response")) + .expect("the response is still recorded"); + assert_eq!( + number(response, "bud.voice.cost"), + None, + "a minute price bills no tokens" + ); + assert!(number(&session, "bud.voice.billed_seconds").unwrap() > 1.5); +} + // ============================================================================================= // Credentials never logged — TC-SEC-07 // ============================================================================================= diff --git a/gateway/tests/realtime_translate.rs b/gateway/tests/realtime_translate.rs new file mode 100644 index 00000000..079cd9dc --- /dev/null +++ b/gateway/tests/realtime_translate.rs @@ -0,0 +1,1440 @@ +//! FRD-023 RT7 end to end, in process: `/v1/realtime` speaking OpenAI Realtime GA to real +//! WebSocket clients while WaaV's native providers speak each vendor's own protocol to an +//! in-process MOCK of that protocol (TC-XL-01…05, TC-XL-07). +//! +//! * Gemini Live — a WebSocket mock of `BidiGenerateContent` (`setup` / `setupComplete`, +//! `realtimeInput`, `clientContent`, `serverContent` + `usageMetadata`, `goAway`, session +//! resumption). +//! * Nova 2 Sonic — an aws-smithy `HttpConnector` that speaks the Bedrock +//! `InvokeModelWithBidirectionalStream` HTTP event stream: it decodes the SDK's SIGNED input +//! event-stream frames and answers with output frames, so the SigV4 request, the event-stream +//! framing and the Nova JSON events are all exercised (the precedent: the Transcribe mock in +//! `mock_endpoint_e2e.rs`). +//! * The per-minute agents (Deepgram Voice Agent, ElevenLabs Agents, Hume EVI) — WebSocket mocks +//! that open the vendor's session and record what they receive. +//! +//! No real vendor is called and no vendor key exists. + +use std::collections::HashMap; +use std::net::SocketAddr; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll}; +use std::time::Duration; + +use base64::Engine as _; +use base64::prelude::BASE64_STANDARD; +use bytes::{Bytes, BytesMut}; +use futures_util::{SinkExt, StreamExt}; +use opentelemetry::trace::TracerProvider as _; +use opentelemetry_sdk::error::OTelSdkResult; +use opentelemetry_sdk::trace::{SdkTracerProvider, SpanData, SpanExporter}; +use serde_json::{Value as Json, json}; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tracing_subscriber::layer::SubscriberExt; + +use waav_gateway::config::{DAGTimeoutsConfig, PluginConfig, ServerConfig}; +use waav_gateway::handlers::openai_realtime::{RealtimeRuntime, Timings}; +use waav_gateway::state::AppState; + +// ============================================================================================= +// Fixtures +// ============================================================================================= + +const KEY: &str = "bud_realtime_translate_test_key"; +const PROJECT: &str = "5b0c7e1d-0000-4000-8000-00000000ab01"; +const USER: &str = "5b0c7e1d-0000-4000-8000-00000000bc01"; +const API_KEY_ID: &str = "5b0c7e1d-0000-4000-8000-00000000cd01"; +const MODEL_ID: &str = "5b0c7e1d-0000-4000-8000-00000000de01"; + +/// bud-auth's fixture ciphertext; the plaintext is `VENDOR_KEY`. +const TEST_CREDENTIAL: &str = include_str!("../../bud-auth/tests/fixtures/test_cred_encrypted.hex"); +const VENDOR_KEY: &str = "dg_vendor_key_abc123"; + +fn test_pem() -> String { + std::fs::read_to_string(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../bud-auth/tests/fixtures/test_cred_private.pem" + )) + .expect("bud-auth's fixture key (git-ignored *.pem) must be present locally") +} + +/// Encrypt a credential exactly as budapp does (RSA-OAEP-SHA-256 blocks, hex). +fn encrypt_like_budapp(plain: &str) -> String { + use rsa::pkcs8::DecodePrivateKey; + let pem = test_pem(); + let key = rsa::RsaPrivateKey::from_pkcs8_pem(&pem) + .or_else(|_| { + use rsa::pkcs1::DecodeRsaPrivateKey; + rsa::RsaPrivateKey::from_pkcs1_pem(&pem) + }) + .expect("the fixture key loads"); + let public = key.to_public_key(); + let mut out = Vec::new(); + for chunk in plain + .as_bytes() + .chunks(rsa::traits::PublicKeyParts::size(&key) - 66) + { + out.extend( + public + .encrypt( + &mut rsa::rand_core::OsRng, + rsa::Oaep::new::(), + chunk, + ) + .unwrap(), + ); + } + hex::encode(out) +} + +fn allow_loopback() { + static ONCE: std::sync::Once = std::sync::Once::new(); + ONCE.call_once(|| unsafe { std::env::set_var("WAAV_ALLOW_LOOPBACK_ENDPOINTS", "1") }); +} + +// ============================================================================================= +// Span capture (per test, thread-local) +// ============================================================================================= + +#[derive(Debug, Clone, Default)] +struct Exported(Arc>>); + +impl SpanExporter for Exported { + fn export( + &self, + batch: Vec, + ) -> impl std::future::Future + Send { + self.0.lock().unwrap().extend(batch); + std::future::ready(Ok(())) + } +} + +struct Capture { + exported: Exported, + _guard: tracing::subscriber::DefaultGuard, + _provider: SdkTracerProvider, +} + +impl Capture { + fn install() -> Self { + let exported = Exported::default(); + let provider = SdkTracerProvider::builder() + .with_simple_exporter(exported.clone()) + .build(); + let subscriber = tracing_subscriber::registry() + .with(tracing_subscriber::filter::LevelFilter::DEBUG) + .with( + tracing_opentelemetry::layer().with_tracer(provider.tracer("realtime-translate")), + ); + let guard = tracing::subscriber::set_default(subscriber); + Self { + exported, + _guard: guard, + _provider: provider, + } + } + + fn spans(&self) -> Vec { + self.exported.0.lock().unwrap().clone() + } + + async fn wait_for(&self, name: &str, n: usize) -> Vec { + for _ in 0..300 { + let found: Vec = self + .spans() + .into_iter() + .filter(|s| s.name == name) + .collect(); + if found.len() >= n { + return found; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("expected {n} `{name}` spans, got {}", self.spans().len()); + } +} + +fn attr(span: &SpanData, key: &str) -> Option { + span.attributes + .iter() + .find(|kv| kv.key.as_str() == key) + .map(|kv| kv.value.clone()) +} + +fn text(span: &SpanData, key: &str) -> Option { + attr(span, key).map(|v| v.as_str().into_owned()) +} + +fn number(span: &SpanData, key: &str) -> Option { + match attr(span, key)? { + opentelemetry::Value::I64(i) => Some(i as f64), + opentelemetry::Value::F64(f) => Some(f), + opentelemetry::Value::String(s) => s.as_str().parse().ok(), + _ => None, + } +} + +// ============================================================================================= +// The gateway +// ============================================================================================= + +fn config() -> ServerConfig { + ServerConfig { + host: "127.0.0.1".to_string(), + port: 0, + tls: None, + livekit_url: "ws://localhost:7880".to_string(), + livekit_public_url: "http://localhost:7880".to_string(), + livekit_api_key: None, + livekit_api_secret: None, + deepgram_api_key: Some("dg-process-canary".to_string()), + elevenlabs_api_key: None, + google_credentials: None, + azure_speech_subscription_key: None, + azure_speech_region: None, + cartesia_api_key: None, + openai_api_key: None, + azure_openai_api_key: None, + azure_openai_endpoint: None, + grok_api_key: None, + inworld_api_key: None, + gemini_api_key: Some("gemini-process-canary".to_string()), + ultravox_api_key: None, + speechmatics_api_key: None, + yandex_api_key: None, + yandex_folder_id: None, + assemblyai_api_key: None, + hume_api_key: None, + groq_api_key: None, + ibm_watson_api_key: None, + ibm_watson_instance_id: None, + ibm_watson_region: None, + aws_access_key_id: Some("AKIDPROCESSCANARY".to_string()), + aws_secret_access_key: Some("process-canary".to_string()), + aws_region: None, + gnani_token: None, + gnani_access_key: None, + gnani_certificate_path: None, + recording_s3_bucket: None, + recording_s3_region: None, + recording_s3_endpoint: None, + recording_s3_access_key: None, + recording_s3_secret_key: None, + recording_s3_prefix: None, + cache_path: None, + cache_ttl_seconds: Some(3600), + auth_service_url: None, + auth_signing_key_path: None, + auth_api_secrets: Vec::new(), + auth_timeout_seconds: 5, + auth_required: false, + sip: None, + cors_allowed_origins: None, + rate_limit_requests_per_second: 60, + rate_limit_burst_size: 10, + max_websocket_connections: None, + max_connections_per_ip: 1000, + ws_processing_timeout_secs: 10, + realtime_processing_timeout_secs: 30, + sip_max_participants: 3, + realtime_endpoint_overrides: Default::default(), + plugins: PluginConfig::default(), + dag_timeouts: DAGTimeoutsConfig::default(), + aliases: Default::default(), + } +} + +fn fast_timings() -> Timings { + Timings { + ping: Duration::from_millis(300), + max_missed_pongs: 3, + revalidate: Duration::from_millis(300), + connect: Duration::from_millis(3000), + hold: Duration::from_millis(1500), + slow_client: Duration::from_millis(600), + max_session: Duration::from_secs(3600), + default_idle: Duration::from_secs(300), + warn_before: Duration::from_secs(1), + segment: Duration::from_secs(1), + upstream_send: Duration::from_secs(2), + connection_cap: None, + } +} + +/// Gemini-shaped token rates (CONTRACTS C7), per 1 000 000 tokens. +fn token_pricing() -> Json { + json!({"unit": "token", "per_units": 1000000, "currency": "USD", + "rates": {"input_text": 0.5, "input_audio": 3.0, "cached_input_text": 0.05, + "cached_input_audio": 0.3, "output_text": 2.0, "output_audio": 12.0}}) +} + +struct Gateway { + addr: SocketAddr, +} + +struct Setup { + endpoints: Vec<(String, String, Json)>, + timings: Timings, + bedrock: Option, +} + +impl Default for Setup { + fn default() -> Self { + Self { + endpoints: Vec::new(), + timings: fast_timings(), + bedrock: None, + } + } +} + +async fn gateway(setup: Setup) -> Gateway { + allow_loopback(); + let store = Arc::new(bud_auth::MemoryStore::new()); + let mut m = serde_json::Map::new(); + for (alias, id, _) in &setup.endpoints { + m.insert( + alias.clone(), + json!({"endpoint_id": id, "model_id": MODEL_ID, "project_id": PROJECT, "kind": "model"}), + ); + } + m.insert( + "__metadata__".into(), + json!({"api_key_id": API_KEY_ID, "user_id": USER, "api_key_project_id": PROJECT}), + ); + store.set( + &format!("api_key:{}", bud_auth::hash_api_key(KEY)), + &Json::Object(m).to_string(), + ); + for (_, id, entry) in &setup.endpoints { + store.set( + &format!("voice_table:{id}"), + &json!({ id.as_str(): entry }).to_string(), + ); + } + let plane = Arc::new(bud_auth::BudPlane::with_decryptor( + store.clone() as Arc, + None, + bud_auth::CredentialDecryptor::from_pem(&test_pem()).unwrap(), + )); + plane.boot().await.unwrap(); + + let mut state = AppState::new(config()).await; + { + let s = Arc::get_mut(&mut state).expect("unshared"); + s.bud_mode = Some(waav_gateway::auth::bud_mode::BudMode::for_plane(plane.clone()).unwrap()); + s.policies = Some(waav_gateway::core::deployment_policy::DeploymentPolicies::local()); + s.realtime = Arc::new(RealtimeRuntime { + timings: setup.timings, + client_secret_keys: None, + bedrock_http_client: setup.bedrock, + }); + } + let app = waav_gateway::routes::openai_realtime::create_openai_realtime_router() + .layer(axum::middleware::from_fn_with_state( + state.clone(), + waav_gateway::middleware::connection_limit_middleware, + )) + .with_state(state.clone()); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + axum::serve( + listener, + app.into_make_service_with_connect_info::(), + ) + .await + .unwrap(); + }); + Gateway { addr } +} + +fn ep(alias: &str, id: &str, entry: Json) -> (String, String, Json) { + (alias.to_string(), id.to_string(), entry) +} + +// ============================================================================================= +// The client +// ============================================================================================= + +type Client = + tokio_tungstenite::WebSocketStream>; + +async fn connect(gw: &Gateway, model: &str) -> Client { + let mut req = format!("ws://{}/v1/realtime?model={model}", gw.addr) + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("authorization", format!("Bearer {KEY}").parse().unwrap()); + match tokio::time::timeout( + Duration::from_secs(15), + tokio_tungstenite::connect_async(req), + ) + .await + .expect("the handshake did not complete within 15 s") + { + Ok((ws, _)) => ws, + Err(tokio_tungstenite::tungstenite::Error::Http(resp)) => { + let body = resp + .body() + .as_ref() + .map(|b| String::from_utf8_lossy(b).into_owned()); + panic!("connect refused {}: {body:?}", resp.status()) + } + Err(e) => panic!("connect failed: {e}"), + } +} + +async fn send(ws: &mut Client, v: Json) { + ws.send(Message::Text(v.to_string().into())).await.unwrap(); +} + +/// The next text event; a close or a silent 5 s fails the test with what was seen. +async fn next_json(ws: &mut Client) -> Json { + loop { + match tokio::time::timeout(Duration::from_secs(5), ws.next()).await { + Ok(Some(Ok(Message::Text(t)))) => return serde_json::from_str(t.as_str()).unwrap(), + Ok(Some(Ok(Message::Ping(_) | Message::Pong(_)))) => continue, + other => panic!("expected a text event, got {other:?}"), + } + } +} + +/// Read until `kind`; every event read on the way is returned too. +async fn until_type(ws: &mut Client, kind: &str) -> (Json, Vec) { + let mut seen = Vec::new(); + for _ in 0..500 { + let v = next_json(ws).await; + if std::env::var("RT7_DEBUG").is_ok() { + eprintln!("<- {v}"); + } + if v["type"] == kind { + return (v, seen); + } + seen.push(v); + } + panic!( + "no {kind}; saw {:?}", + seen.iter().map(|e| e["type"].clone()).collect::>() + ); +} + +/// Wait while still reading (the reads answer the server's pings); returns what arrived. +async fn keep_alive(ws: &mut Client, dur: Duration) -> Vec { + let deadline = tokio::time::Instant::now() + dur; + let mut seen = Vec::new(); + while tokio::time::Instant::now() < deadline { + match tokio::time::timeout_at(deadline, ws.next()).await { + Ok(Some(Ok(Message::Text(t)))) => seen.push(serde_json::from_str(t.as_str()).unwrap()), + Ok(Some(Ok(Message::Close(f)))) => panic!("the session closed: {f:?}; saw {seen:?}"), + Ok(None) => panic!("the session ended; saw {seen:?}"), + _ => {} + } + } + seen +} + +async fn close_and_drain(ws: &mut Client) { + let _ = ws.close(None).await; + for _ in 0..500 { + match tokio::time::timeout(Duration::from_secs(5), ws.next()).await { + Ok(Some(Ok(Message::Close(_)))) | Ok(None) | Ok(Some(Err(_))) | Err(_) => return, + _ => {} + } + } +} + +fn user_text(text: &str) -> Json { + json!({"type": "conversation.item.create", "item": {"type": "message", "role": "user", + "content": [{"type": "input_text", "text": text}]}}) +} + +// ============================================================================================= +// Gemini Live mock (BidiGenerateContent over WebSocket) +// ============================================================================================= + +/// 20 ms of 24 kHz PCM16 with a recognisable byte pattern. +fn gemini_out_pcm() -> Vec { + (0..960u32).map(|i| (i % 251) as u8).collect() +} + +#[derive(Clone, Default)] +struct GeminiBehaviour { + /// After the first turn on the first connection: `goAway`, then keep the socket open (the + /// gateway must reconnect on its own, with the resumption handle). + go_away_after_first_turn: bool, +} + +#[derive(Default)] +struct GeminiLog { + /// Per connection: the URL it was opened with. + urls: Vec, + /// Per connection: every `setup` it received. + setups: Vec>, + /// Every `realtimeInput.mediaChunks[]`: (connection, mime, decoded bytes). + audio: Vec<(usize, String, usize)>, + /// Every `clientContent` text: (connection, text). + texts: Vec<(usize, String)>, +} + +#[derive(Clone)] +struct GeminiMock { + addr: SocketAddr, + log: Arc>, +} + +impl GeminiMock { + async fn start(b: GeminiBehaviour) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let log = Arc::new(Mutex::new(GeminiLog::default())); + let shared = log.clone(); + tokio::spawn(async move { + let mut n = 0usize; + loop { + let Ok((stream, _)) = listener.accept().await else { + return; + }; + let conn = n; + n += 1; + let log = shared.clone(); + let b = b.clone(); + tokio::spawn(async move { + let captured = log.clone(); + let callback = move |req: &tokio_tungstenite::tungstenite::handshake::server::Request, + resp: tokio_tungstenite::tungstenite::handshake::server::Response| { + let mut l = captured.lock().unwrap(); + l.urls.push(req.uri().to_string()); + l.setups.push(Vec::new()); + Ok(resp) + }; + let Ok(mut ws) = tokio_tungstenite::accept_hdr_async(stream, callback).await + else { + return; + }; + let mut turns = 0usize; + while let Some(Ok(msg)) = ws.next().await { + let Message::Text(t) = msg else { continue }; + let v: Json = serde_json::from_str(t.as_str()).unwrap_or(Json::Null); + let mut out: Vec = Vec::new(); + if let Some(setup) = v.get("setup") { + log.lock().unwrap().setups[conn].push(setup.clone()); + out.push(json!({"setupComplete": {}})); + out.push(json!({"sessionResumptionUpdate": { + "newHandle": format!("handle-{conn}"), "resumable": true}})); + } else if let Some(ri) = v.get("realtimeInput") { + for c in ri["mediaChunks"].as_array().into_iter().flatten() { + let bytes = BASE64_STANDARD + .decode(c["data"].as_str().unwrap_or("")) + .map(|b| b.len()) + .unwrap_or(0); + log.lock().unwrap().audio.push(( + conn, + c["mimeType"].as_str().unwrap_or("").to_string(), + bytes, + )); + } + } else if let Some(cc) = v.get("clientContent") { + let said = cc["turns"][0]["parts"][0]["text"] + .as_str() + .unwrap_or("") + .to_string(); + log.lock().unwrap().texts.push((conn, said)); + turns += 1; + out.push(json!({"serverContent": { + "modelTurn": {"parts": [{"inlineData": { + "mimeType": "audio/pcm;rate=24000", + "data": BASE64_STANDARD.encode(gemini_out_pcm())}}]}, + "outputTranscription": {"text": "Hello"}}})); + out.push(json!({"serverContent": { + "outputTranscription": {"text": " there"}, "turnComplete": true}, + "usageMetadata": { + "promptTokenCount": 150, "cachedContentTokenCount": 40, + "responseTokenCount": 60, "totalTokenCount": 210, + "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 100}, + {"modality": "AUDIO", "tokenCount": 50}], + "cacheTokensDetails": [{"modality": "TEXT", "tokenCount": 40}], + "responseTokensDetails": [{"modality": "AUDIO", "tokenCount": 50}, + {"modality": "TEXT", "tokenCount": 10}]}})); + if b.go_away_after_first_turn && conn == 0 && turns == 1 { + out.push(json!({"goAway": {"timeLeft": "5s"}})); + } + } + for o in out { + if ws.send(Message::Text(o.to_string().into())).await.is_err() { + return; + } + } + } + }); + } + }); + Self { addr, log } + } + + fn base(&self) -> String { + format!("http://{}/ws/gemini", self.addr) + } + + fn setups(&self) -> Vec> { + self.log.lock().unwrap().setups.clone() + } +} + +fn gemini_entry(mock: &GeminiMock, extra: Json) -> Json { + let mut e = json!({ + "vendor": "gemini", + "api_base": mock.base(), + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["realtime_session"], + "model": "gemini-3.8-live", + "pricing": token_pricing(), + "config": {"realtime": {"defaults": {"voice": "Puck", "instructions": "deployment prompt"}, + "limits": {"max_session_seconds": 900}}} + }); + for (k, v) in extra.as_object().into_iter().flatten() { + e[k] = v.clone(); + } + e +} + +/// TC-XL-01 🔒 — the Gemini setup is built from the client's first `session.update` (over the +/// deployment defaults) and sent ONCE; a later change of voice or tools is refused with +/// `event_not_allowed` naming the field; the session continues and no second setup is sent. +/// The deployment's key authenticates the vendor leg, never the process's. +#[tokio::test] +async fn tc_xl_01_gemini_setup_from_the_ga_session_update() { + let mock = GeminiMock::start(GeminiBehaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "live", + "c7c7c7c7-0000-4000-8000-000000000101", + gemini_entry(&mock, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "live").await; + let (created, _) = until_type(&mut c, "session.created").await; + assert_eq!(created["session"]["model"], "live"); + assert_eq!( + created["session"]["audio"]["output"]["voice"], "Puck", + "the deployment's default" + ); + + let tool = json!({"type": "function", "name": "get_weather", "description": "d", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}}); + send( + &mut c, + json!({"type": "session.update", "event_id": "c1", "session": { + "type": "realtime", "model": "gpt-realtime", "instructions": "client prompt", + "audio": {"output": {"voice": "Kore"}}, "tools": [tool]}}), + ) + .await; + let (updated, _) = until_type(&mut c, "session.updated").await; + assert_eq!(updated["session"]["audio"]["output"]["voice"], "Kore"); + assert_eq!( + updated["session"]["model"], "live", + "the client's `model` is never used" + ); + + let setups = mock.setups(); + assert_eq!(setups.len(), 1, "one connection"); + assert_eq!(setups[0].len(), 1, "one setup frame"); + let setup = &setups[0][0]; + assert_eq!(setup["model"], "models/gemini-3.8-live"); + assert_eq!( + setup["generationConfig"]["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"], + "Kore" + ); + assert_eq!( + setup["systemInstruction"]["parts"][0]["text"], + "client prompt" + ); + assert_eq!( + setup["tools"][0]["functionDeclarations"][0]["name"], + "get_weather" + ); + let url = mock.log.lock().unwrap().urls[0].clone(); + assert!( + url.contains(&format!("key={VENDOR_KEY}")), + "the deployment's key: {url}" + ); + assert!(!url.contains("gemini-process-canary")); + + // A later voice change: refused by name, not silently ignored. + send( + &mut c, + json!({"type": "session.update", "event_id": "c2", "session": { + "audio": {"output": {"voice": "Charon"}}}}), + ) + .await; + let (e, _) = until_type(&mut c, "error").await; + assert_eq!(e["error"]["code"], "event_not_allowed"); + assert_eq!(e["error"]["param"], "session.audio.output.voice"); + assert_eq!(e["error"]["event_id"], "c2"); + // A later tools change: the same. + send( + &mut c, + json!({"type": "session.update", "session": {"tools": []}}), + ) + .await; + let (e, _) = until_type(&mut c, "error").await; + assert_eq!(e["error"]["param"], "session.tools"); + // `conversation.item.truncate` has no Gemini equivalent (FRD §5.7). + send(&mut c, json!({"type": "conversation.item.truncate", "item_id": "i", "content_index": 0, "audio_end_ms": 5})).await; + let (e, _) = until_type(&mut c, "error").await; + assert_eq!(e["error"]["param"], "conversation.item.truncate"); + + // The session is still alive and still set up once. + send(&mut c, user_text("hi")).await; + until_type(&mut c, "response.done").await; + assert_eq!(mock.setups()[0].len(), 1, "no second setup"); + close_and_drain(&mut c).await; +} + +/// TC-XL-02 — 24 kHz client audio reaches Gemini resampled to 16 kHz (declared as such); the +/// vendor's 24 kHz output reaches the client unchanged as `response.output_audio.delta`. +#[tokio::test] +async fn tc_xl_02_gemini_audio_rates() { + let mock = GeminiMock::start(GeminiBehaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "live", + "c7c7c7c7-0000-4000-8000-000000000102", + gemini_entry(&mock, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "live").await; + until_type(&mut c, "session.created").await; + // 1 s of a tone at 24 kHz, in 20 ms appends (the first one sets the session up). + let pcm: Vec = (0..24_000) + .flat_map(|i| { + ((f32::sin(i as f32 * 440.0 * std::f32::consts::TAU / 24_000.0) * 8000.0) as i16) + .to_le_bytes() + }) + .collect(); + for chunk in pcm.chunks(960) { + send( + &mut c, + json!({"type": "input_audio_buffer.append", + "audio": BASE64_STANDARD.encode(chunk)}), + ) + .await; + } + send(&mut c, json!({"type": "input_audio_buffer.commit"})).await; + until_type(&mut c, "input_audio_buffer.committed").await; + let mut received = 0usize; + for _ in 0..100 { + let audio = mock.log.lock().unwrap().audio.clone(); + assert!( + audio + .iter() + .all(|(_, mime, _)| mime == "audio/pcm;rate=16000"), + "every chunk declared 16 kHz" + ); + received = audio.iter().map(|(_, _, n)| n).sum(); + if received >= 31_000 { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + // 1 s at 16 kHz PCM16 = 32 000 bytes (± the resampler's one chunk). + assert!( + (31_000..=33_000).contains(&received), + "{received} bytes at 16 kHz" + ); + + send(&mut c, user_text("say something")).await; + let (_, seen) = until_type(&mut c, "response.done").await; + let audio: Vec = seen + .iter() + .filter(|e| e["type"] == "response.output_audio.delta") + .flat_map(|e| { + BASE64_STANDARD + .decode(e["delta"].as_str().unwrap()) + .unwrap() + }) + .collect(); + assert_eq!( + audio, + gemini_out_pcm(), + "24 kHz output passes through unchanged" + ); + let order: Vec = seen + .iter() + .map(|e| e["type"].as_str().unwrap().to_string()) + .filter(|t| t.starts_with("response.")) + .collect(); + assert_eq!(order.first().map(String::as_str), Some("response.created")); + close_and_drain(&mut c).await; +} + +/// TC-XL-03 🔒 — Gemini's `usageMetadata` becomes the `usage` of the `response.done` it closes +/// and exactly one priced `voice.turn` per response (cached text a subset of input text). +#[tokio::test] +async fn tc_xl_03_gemini_usage_is_the_response_usage_and_a_priced_turn() { + let cap = Capture::install(); + let mock = GeminiMock::start(GeminiBehaviour::default()).await; + let gw = gateway(Setup { + endpoints: vec![ep( + "live", + "c7c7c7c7-0000-4000-8000-000000000103", + gemini_entry(&mock, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "live").await; + until_type(&mut c, "session.created").await; + for i in 0..2 { + send(&mut c, user_text(&format!("turn {i}"))).await; + let (done, seen) = until_type(&mut c, "response.done").await; + let usage = &done["response"]["usage"]; + assert_eq!(usage["input_tokens"], 150); + assert_eq!(usage["output_tokens"], 60); + assert_eq!(usage["input_token_details"]["text_tokens"], 100); + assert_eq!(usage["input_token_details"]["audio_tokens"], 50); + assert_eq!( + usage["input_token_details"]["cached_tokens_details"]["text_tokens"], + 40 + ); + assert_eq!(usage["output_token_details"]["audio_tokens"], 50); + let transcript: String = seen + .iter() + .filter(|e| e["type"] == "response.output_audio_transcript.delta") + .map(|e| e["delta"].as_str().unwrap().to_string()) + .collect(); + assert_eq!(transcript, "Hello there"); + } + close_and_drain(&mut c).await; + + let turns = cap.wait_for("voice.turn", 2).await; + assert_eq!(turns.len(), 2, "one billed record per response, never two"); + let session = cap.wait_for("voice.session", 1).await.remove(0); + // (100 − 40)·0.5 + 40·0.05 + 50·3 + 10·2 + 50·12 = 30 + 2 + 150 + 20 + 600 = 802 per 10⁶. + for t in &turns { + assert_eq!( + text(t, "bud.voice.rt.component").as_deref(), + Some("response") + ); + assert_eq!(text(t, "bud.voice.rt.vendor").as_deref(), Some("gemini")); + assert_eq!( + text(t, "bud.voice.rt.model").as_deref(), + Some("gemini-3.8-live") + ); + assert_eq!( + text(t, "bud.voice.rt.response_status").as_deref(), + Some("completed") + ); + assert_eq!(number(t, "bud.voice.rt.input_text_tokens"), Some(100.0)); + assert_eq!(number(t, "bud.voice.rt.cached_text_tokens"), Some(40.0)); + assert_eq!(number(t, "bud.voice.rt.output_audio_tokens"), Some(50.0)); + let cost = number(t, "bud.voice.cost").expect("priced"); + assert!((cost - 802.0e-6).abs() < 1e-12, "cost {cost}"); + assert_eq!(text(t, "bud.voice.pricing_unit").as_deref(), Some("token")); + assert_eq!( + text(t, "bud.endpoint_id").as_deref(), + Some("c7c7c7c7-0000-4000-8000-000000000103") + ); + assert_eq!(text(t, "bud.api_key_id").as_deref(), Some(API_KEY_ID)); + } + assert_eq!(number(&session, "bud.voice.session.turns"), Some(2.0)); + assert!((number(&session, "bud.voice.cost").unwrap() - 2.0 * 802.0e-6).abs() < 1e-12); +} + +/// TC-XL-04 — Gemini's `goAway`: WaaV reconnects on its own with the resumption handle the +/// vendor issued, and the client sees no close and no error — the next turn simply works. +#[tokio::test] +async fn tc_xl_04_gemini_go_away_is_a_transparent_resumption() { + let mock = GeminiMock::start(GeminiBehaviour { + go_away_after_first_turn: true, + }) + .await; + let gw = gateway(Setup { + endpoints: vec![ep( + "live", + "c7c7c7c7-0000-4000-8000-000000000104", + gemini_entry(&mock, json!({})), + )], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "live").await; + until_type(&mut c, "session.created").await; + send(&mut c, user_text("first")).await; + until_type(&mut c, "response.done").await; + // The goAway follows the turn; the gateway replaces the connection meanwhile. + let seen = keep_alive(&mut c, Duration::from_millis(800)).await; + assert!( + seen.iter().all(|e| e["type"] != "error"), + "the client is told nothing: {seen:?}" + ); + let setups = mock.setups(); + assert_eq!(setups.len(), 2, "a second connection"); + assert_eq!( + setups[1][0]["sessionResumption"]["handle"], "handle-0", + "resumed with the handle the vendor issued" + ); + assert!(setups[0][0]["sessionResumption"].get("handle").is_none()); + + send(&mut c, user_text("second")).await; + until_type(&mut c, "response.done").await; + let texts = mock.log.lock().unwrap().texts.clone(); + assert_eq!( + texts, + vec![(0, "first".to_string()), (1, "second".to_string())] + ); + close_and_drain(&mut c).await; +} + +// ============================================================================================= +// Nova 2 Sonic mock: the Bedrock bidirectional event stream, behind an aws-smithy connector +// ============================================================================================= + +mod bedrock { + use super::*; + use aws_smithy_eventstream::frame::{ + DecodedFrame, MessageFrameDecoder, read_message_from, write_message_to, + }; + use aws_smithy_runtime_api::client::http::{ + HttpClient, HttpConnector, HttpConnectorFuture, HttpConnectorSettings, SharedHttpClient, + SharedHttpConnector, + }; + use aws_smithy_runtime_api::client::orchestrator::HttpRequest; + use aws_smithy_runtime_api::client::runtime_components::RuntimeComponents; + use aws_smithy_runtime_api::http::{Response, StatusCode}; + use aws_smithy_types::body::SdkBody; + use aws_smithy_types::event_stream::{Header, HeaderValue, Message as EsMessage}; + use http_body_util::BodyExt; + + /// One stream's output half, fed by the mock. + struct ChanBody(mpsc::UnboundedReceiver); + + impl http_body::Body for ChanBody { + type Data = Bytes; + type Error = std::convert::Infallible; + fn poll_frame( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>>> { + self.0 + .poll_recv(cx) + .map(|o| o.map(|b| Ok(http_body::Frame::data(b)))) + } + } + + #[derive(Default)] + pub struct NovaLog { + /// Per stream: (uri, Authorization header). + pub calls: Vec<(String, String)>, + /// Per stream: the Nova events received, in order. + pub events: Vec>, + } + + /// The Nova events a stream answers a user text turn with: a speculative and a final + /// assistant transcript, one audio block, then a usage report whose `delta` is this turn's + /// and whose `total` is the running sum ACROSS streams — a meter that read totals would bill + /// twice. + fn answer(stream: usize, turn: usize, total_turns: u64) -> Vec { + let content = format!("ct-{stream}-{turn}"); + vec![ + json!({"event": {"completionStart": {"completionId": format!("cmp-{stream}-{turn}")}}}), + json!({"event": {"contentStart": {"type": "TEXT", "role": "ASSISTANT", + "contentId": format!("tx-{stream}-{turn}"), + "additionalModelFields": "{\"generationStage\":\"SPECULATIVE\"}"}}}), + json!({"event": {"textOutput": {"content": "Sure."}}}), + json!({"event": {"contentStart": {"type": "TEXT", "role": "ASSISTANT", + "contentId": format!("txf-{stream}-{turn}"), + "additionalModelFields": "{\"generationStage\":\"FINAL\"}"}}}), + json!({"event": {"textOutput": {"content": "Sure."}}}), + json!({"event": {"contentStart": {"type": "AUDIO", "role": "ASSISTANT", "contentId": content}}}), + json!({"event": {"audioOutput": {"contentId": content, + "content": BASE64_STANDARD.encode(vec![7u8; 480])}}}), + json!({"event": {"contentEnd": {"contentId": content, "stopReason": "END_TURN"}}}), + json!({"event": {"usageEvent": {"completionId": "c", "details": { + "delta": {"input": {"speechTokens": 40, "textTokens": 3}, + "output": {"speechTokens": 50, "textTokens": 7}}, + "total": {"input": {"speechTokens": 40 * total_turns, "textTokens": 3 * total_turns}, + "output": {"speechTokens": 50 * total_turns, "textTokens": 7 * total_turns}}}}}}), + json!({"event": {"completionEnd": {"completionId": format!("cmp-{stream}-{turn}")}}}), + ] + } + + fn output_frame(event: &Json) -> Bytes { + let payload = json!({"bytes": BASE64_STANDARD.encode(event.to_string())}).to_string(); + let msg = EsMessage::new(payload.into_bytes()) + .add_header(Header::new( + ":message-type", + HeaderValue::String("event".into()), + )) + .add_header(Header::new( + ":event-type", + HeaderValue::String("chunk".into()), + )) + .add_header(Header::new( + ":content-type", + HeaderValue::String("application/json".into()), + )); + let mut out = Vec::new(); + write_message_to(&msg, &mut out).unwrap(); + Bytes::from(out) + } + + /// The Nova event inside one input frame: SigV4 wraps each event in an outer message whose + /// payload is the inner `chunk` message, whose payload is `{"bytes": base64(event json)}`. + fn input_event(outer: &EsMessage) -> Option { + let signed = outer + .headers() + .iter() + .any(|h| h.name().as_str() == ":chunk-signature"); + let inner = if signed { + if outer.payload().is_empty() { + return None; // the closing empty signed frame + } + read_message_from(outer.payload().as_ref()).ok()? + } else { + outer.clone() + }; + let body: Json = serde_json::from_slice(inner.payload()).ok()?; + let bytes = BASE64_STANDARD.decode(body["bytes"].as_str()?).ok()?; + serde_json::from_slice(&bytes).ok() + } + + #[derive(Clone, Default)] + pub struct NovaMock { + pub log: Arc>, + turns: Arc>, + } + + impl std::fmt::Debug for NovaMock { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("NovaMock") + } + } + + impl NovaMock { + pub fn client(&self) -> SharedHttpClient { + SharedHttpClient::new(self.clone()) + } + } + + impl HttpConnector for NovaMock { + fn call(&self, req: HttpRequest) -> HttpConnectorFuture { + let stream = { + let mut log = self.log.lock().unwrap(); + log.calls.push(( + req.uri().to_string(), + req.headers() + .get("authorization") + .unwrap_or_default() + .to_string(), + )); + log.events.push(Vec::new()); + log.calls.len() - 1 + }; + let (tx, rx) = mpsc::unbounded_channel::(); + let log = self.log.clone(); + let turns = self.turns.clone(); + let mut body = req.into_body(); + tokio::spawn(async move { + let mut buf = BytesMut::new(); + let mut decoder = MessageFrameDecoder::new(); + let mut turn = 0usize; + let mut user_text_open = false; + while let Some(Ok(frame)) = body.frame().await { + let Ok(data) = frame.into_data() else { + continue; + }; + buf.extend_from_slice(&data); + while let Ok(DecodedFrame::Complete(outer)) = decoder.decode_frame(&mut buf) { + let Some(ev) = input_event(&outer) else { + continue; + }; + log.lock().unwrap().events[stream].push(ev.clone()); + let e = &ev["event"]; + if e["contentStart"]["role"] == "USER" + && e["contentStart"]["type"] == "TEXT" + { + user_text_open = true; + } + if e.get("textInput").is_some() && user_text_open { + user_text_open = false; + turn += 1; + let total = { + let mut t = turns.lock().unwrap(); + *t += 1; + *t + }; + for out in answer(stream, turn, total) { + if tx.send(output_frame(&out)).is_err() { + return; + } + } + } + } + } + }); + HttpConnectorFuture::new(async move { + let mut resp = Response::new( + StatusCode::try_from(200u16).unwrap(), + SdkBody::from_body_1_x(ChanBody(rx)), + ); + resp.headers_mut() + .insert("content-type", "application/vnd.amazon.eventstream"); + Ok(resp) + }) + } + } + + impl HttpClient for NovaMock { + fn http_connector( + &self, + _s: &HttpConnectorSettings, + _c: &RuntimeComponents, + ) -> SharedHttpConnector { + SharedHttpConnector::new(self.clone()) + } + } +} + +fn nova_entry() -> Json { + json!({ + "vendor": "nova_sonic", + "credential": encrypt_like_budapp( + r#"{"access_key_id":"AKIDNOVADEPLOYMENT","secret_access_key":"nova-deployment-secret"}"#, + ), + "endpoints": ["realtime_session"], + "model": "amazon.nova-2-sonic-v1:0", + "provider_params": {"region": "eu-north-1"}, + "pricing": {"unit": "token", "per_units": 1000000, "currency": "USD", + "rates": {"input_text": 0.06, "input_audio": 3.4, "output_text": 0.24, "output_audio": 13.6}}, + "config": {"realtime": {"defaults": {"voice": "tiffany", "instructions": "be helpful"}}} + }) +} + +/// TC-XL-05 🔒 — Nova 2 Sonic: the Bedrock stream is SigV4-signed with the DEPLOYMENT's key +/// pair in its region (never the gateway's AWS identity); the connection cap (shortened here +/// from 8 minutes) is met by a reconnect INSIDE the translator — a new stream with the session's +/// history, invisible to the client — and every `usageEvent` is metered exactly once, from its +/// `delta`, even though the second stream's `total` counts both. +#[tokio::test] +async fn tc_xl_05_nova_sonic_cap_reconnect_and_usage_metered_once() { + let cap = Capture::install(); + let nova = bedrock::NovaMock::default(); + let gw = gateway(Setup { + endpoints: vec![ep( + "sonic", + "c7c7c7c7-0000-4000-8000-000000000105", + nova_entry(), + )], + timings: Timings { + connection_cap: Some(Duration::from_millis(1200)), + ..fast_timings() + }, + bedrock: Some(nova.client()), + }) + .await; + let mut c = connect(&gw, "sonic").await; + until_type(&mut c, "session.created").await; + send(&mut c, user_text("first")).await; + let (done, seen) = until_type(&mut c, "response.done").await; + assert_eq!( + done["response"]["usage"]["input_token_details"]["audio_tokens"], + 40 + ); + assert_eq!( + done["response"]["usage"]["output_token_details"]["text_tokens"], + 7 + ); + let transcript: String = seen + .iter() + .filter(|e| e["type"] == "response.output_audio_transcript.delta") + .map(|e| e["delta"].as_str().unwrap().to_string()) + .collect(); + assert_eq!( + transcript, "Sure.", + "the FINAL block restates the SPECULATIVE one" + ); + assert!( + seen.iter() + .any(|e| e["type"] == "response.output_audio.delta") + ); + + // Past the cap: a second stream, and the client notices nothing. + let seen = keep_alive(&mut c, Duration::from_millis(1600)).await; + assert!(seen.iter().all(|e| e["type"] != "error"), "{seen:?}"); + send(&mut c, user_text("second")).await; + let (done, _) = until_type(&mut c, "response.done").await; + assert_eq!( + done["response"]["usage"]["input_token_details"]["audio_tokens"], 40, + "the delta, not the total" + ); + close_and_drain(&mut c).await; + + let (calls, events) = { + let log = nova.log.lock().unwrap(); + (log.calls.clone(), log.events.clone()) + }; + assert!( + calls.len() >= 2, + "the cap replaced the stream: {} streams", + calls.len() + ); + for (uri, auth) in &calls { + assert!( + uri.starts_with("https://bedrock-runtime.eu-north-1.amazonaws.com/"), + "{uri}" + ); + assert!(uri.contains("amazon.nova-2-sonic-v1"), "{uri}"); + assert!(auth.contains("Credential=AKIDNOVADEPLOYMENT/"), "{auth}"); + assert!( + !auth.contains("AKIDPROCESSCANARY"), + "never the gateway's identity" + ); + } + // Every stream opens a Nova session with the deployment's voice and prompt. + for events in events.iter().filter(|e| !e.is_empty()) { + assert!( + events[0]["event"].get("sessionStart").is_some(), + "{:?}", + events[0] + ); + let prompt_start = events + .iter() + .find(|e| e["event"].get("promptStart").is_some()) + .unwrap(); + assert_eq!( + prompt_start["event"]["promptStart"]["audioOutputConfiguration"]["voiceId"], + "tiffany" + ); + } + // The second stream carries the conversation so far. + let second = &events[1]; + assert!( + second + .iter() + .any(|e| e["event"]["textInput"]["content"] == "Sure."), + "history replayed on the new stream: {second:?}" + ); + + // The session record is written at close, after every turn. + let session = cap.wait_for("voice.session", 1).await.remove(0); + let turns: Vec = cap + .spans() + .into_iter() + .filter(|s| s.name == "voice.turn") + .collect(); + assert_eq!( + turns.len(), + 2, + "two usage reports, two billed records — not three, not four" + ); + assert_eq!(number(&session, "bud.voice.session.turns"), Some(2.0)); + assert_eq!( + number(&session, "bud.voice.rt.input_audio_tokens"), + Some(80.0) + ); + for t in &turns { + assert_eq!( + text(t, "bud.voice.rt.vendor").as_deref(), + Some("nova_sonic") + ); + assert_eq!(number(t, "bud.voice.rt.input_audio_tokens"), Some(40.0)); + assert_eq!(number(t, "bud.voice.rt.output_audio_tokens"), Some(50.0)); + // (3·0.06 + 40·3.4 + 7·0.24 + 50·13.6) / 10⁶ + let want = (3.0 * 0.06 + 40.0 * 3.4 + 7.0 * 0.24 + 50.0 * 13.6) / 1e6; + assert!((number(t, "bud.voice.cost").unwrap() - want).abs() < 1e-12); + } +} + +/// FRD-023 RT7.2 🔒 — a Nova deployment without its AWS key pair is refused before the +/// upgrade, although the gateway process holds AWS keys of its own. +#[tokio::test] +async fn nova_without_a_key_pair_is_refused_before_the_upgrade() { + let mut entry = nova_entry(); + entry["credential"] = json!(TEST_CREDENTIAL.trim()); // a plain key, not a pair + let nova = bedrock::NovaMock::default(); + let gw = gateway(Setup { + endpoints: vec![ep("sonic", "c7c7c7c7-0000-4000-8000-000000000106", entry)], + bedrock: Some(nova.client()), + ..Default::default() + }) + .await; + let mut req = format!("ws://{}/v1/realtime?model=sonic", gw.addr) + .into_client_request() + .unwrap(); + req.headers_mut() + .insert("authorization", format!("Bearer {KEY}").parse().unwrap()); + let Err(tokio_tungstenite::tungstenite::Error::Http(resp)) = + tokio_tungstenite::connect_async(req).await + else { + panic!("upgraded without the deployment's key pair"); + }; + assert_eq!(resp.status().as_u16(), 502); + let body: Json = serde_json::from_slice(resp.body().as_ref().unwrap()).unwrap(); + assert_eq!(body["error"]["code"], "deployment_misconfigured"); + assert!(nova.log.lock().unwrap().calls.is_empty(), "no Bedrock call"); +} + +// ============================================================================================= +// The per-minute agents: Deepgram Voice Agent, ElevenLabs Agents, Hume EVI +// ============================================================================================= + +#[derive(Clone)] +struct AgentMock { + addr: SocketAddr, + /// (path?query, lower-cased headers) per upgrade. + upgrades: Arc)>>>, +} + +impl AgentMock { + async fn start(hello: Json) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let upgrades = Arc::new(Mutex::new(Vec::new())); + let shared = upgrades.clone(); + tokio::spawn(async move { + loop { + let Ok((stream, _)) = listener.accept().await else { + return; + }; + let hello = hello.clone(); + let captured = shared.clone(); + tokio::spawn(async move { + let callback = move |req: &tokio_tungstenite::tungstenite::handshake::server::Request, + resp: tokio_tungstenite::tungstenite::handshake::server::Response| { + let headers = req + .headers() + .iter() + .map(|(k, v)| (k.as_str().to_ascii_lowercase(), v.to_str().unwrap_or("").to_string())) + .collect(); + let pq = req.uri().path_and_query().map(|p| p.to_string()).unwrap_or_default(); + captured.lock().unwrap().push((pq, headers)); + Ok(resp) + }; + let Ok(mut ws) = tokio_tungstenite::accept_hdr_async(stream, callback).await + else { + return; + }; + let _ = ws.send(Message::Text(hello.to_string().into())).await; + while let Some(Ok(_)) = ws.next().await {} + }); + } + }); + Self { addr, upgrades } + } + + fn base(&self) -> String { + format!("http://{}/agent", self.addr) + } +} + +/// TC-XL-07 🔒 — a per-minute vendor bills DURATION SEGMENTS from the moment its connection +/// opens: full segments while the session lives and the partial remainder at close (segments +/// shortened to 1 s here, so ~2.5 s bills 1 + 1 + ~0.5; the 60 + 60 + 30 arithmetic is the +/// metering unit test), each priced per minute and none as tokens — for all three vendors. +#[tokio::test] +async fn tc_xl_07_per_minute_vendors_bill_duration_segments() { + let cap = Capture::install(); + let cases = [ + ( + "deepgram_voice_agent", + json!({"type": "Welcome", "request_id": "dg-req-1"}), + json!({"model": "gpt-4o-mini"}), + ), + ( + "elevenlabs_convai", + json!({"type": "conversation_initiation_metadata", + "conversation_initiation_metadata_event": {"conversation_id": "conv_1"}}), + json!({"model": "agent_abc"}), + ), + ( + "hume_evi", + json!({"type": "chat_metadata", "chat_id": "chat_1", "chat_group_id": "g_1", "request_id": "r_1"}), + json!({}), + ), + ]; + for (i, (vendor, hello, extra)) in cases.into_iter().enumerate() { + let mock = AgentMock::start(hello).await; + let mut entry = json!({ + "vendor": vendor, + "api_base": mock.base(), + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["realtime_session"], + "pricing": {"unit": "minute", "cost_per_unit": 0.08, "per_units": 1, "currency": "USD"} + }); + for (k, v) in extra.as_object().into_iter().flatten() { + entry[k] = v.clone(); + } + let id = format!("c7c7c7c7-0000-4000-8000-00000000017{i}"); + let gw = gateway(Setup { + endpoints: vec![ep("agent", &id, entry)], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "agent").await; + until_type(&mut c, "session.created").await; + send( + &mut c, + json!({"type": "session.update", "session": {"instructions": "hello"}}), + ) + .await; + until_type(&mut c, "session.updated").await; + keep_alive(&mut c, Duration::from_millis(2500)).await; + close_and_drain(&mut c).await; + + let sessions = cap.wait_for("voice.session", i + 1).await; + let session = sessions + .iter() + .find(|s| text(s, "bud.endpoint_id").as_deref() == Some(id.as_str())) + .unwrap(); + let segments: Vec = cap + .spans() + .into_iter() + .filter(|s| { + s.name == "voice.turn" && text(s, "bud.endpoint_id").as_deref() == Some(id.as_str()) + }) + .collect(); + assert!(!segments.is_empty(), "{vendor}: no segment"); + for s in &segments { + assert_eq!( + text(s, "bud.voice.rt.component").as_deref(), + Some("duration_segment"), + "{vendor}: a per-minute vendor bills no token turns" + ); + assert_eq!(text(s, "bud.voice.pricing_unit").as_deref(), Some("minute")); + let secs = number(s, "bud.voice.billed_seconds").unwrap(); + assert!( + secs <= 1.0 + 1e-9, + "{vendor}: a segment is at most the segment length" + ); + let cost = number(s, "bud.voice.cost").unwrap(); + assert!((cost - secs / 60.0 * 0.08).abs() < 1e-12, "{vendor}"); + assert_eq!(text(s, "bud.voice.rt.vendor").as_deref(), Some(vendor)); + } + let billed: Vec = segments + .iter() + .map(|s| number(s, "bud.voice.billed_seconds").unwrap()) + .collect(); + let full = billed.iter().filter(|s| (**s - 1.0).abs() < 1e-9).count(); + assert!(full >= 2, "{vendor}: full segments while live: {billed:?}"); + let total: f64 = billed.iter().sum(); + assert!(total > 2.3 && total < 3.6, "{vendor}: {billed:?}"); + assert!((total - number(session, "bud.voice.billed_seconds").unwrap()).abs() < 1e-9); + // The deployment's key reached the vendor. + let (pq, headers) = mock.upgrades.lock().unwrap()[0].clone(); + let carried = pq.contains(VENDOR_KEY) || headers.values().any(|v| v.contains(VENDOR_KEY)); + assert!(carried, "{vendor}: the deployment's key is used"); + assert!( + !headers.values().any(|v| v.contains("dg-process-canary")), + "{vendor}: never the process key" + ); + } +} From 6cd3850d4e241e25baace0890761eb4dae18b1c3 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 07:27:39 +0530 Subject: [PATCH 12/17] fix(gateway): TCP_NODELAY on accepted client sockets (FRD-023 TC-PERF-01) WaaV never set TCP_NODELAY on the sockets it accepts. A realtime client streams 20 ms frames and reads small events back; with Nagle on, each relayed write waited for the client's next frame to acknowledge the previous one, adding a frame interval to every frame: TC-PERF-01 measured p50 19.5 ms / p99 20.4 ms added against a 5 ms budget. The vendor sockets already set it. server::nodelay_listener (axum::serve) and server::tls_server (axum-server's NoDelayAcceptor under rustls) now serve every route; main.rs uses both. Tests: server::accepted_sockets_set_tcp_nodelay; the TC-PERF-01 harness (tests/openai_realtime_perf.rs, #[ignore]) now serves through the production listener: p50 0.37 ms / p99 0.70 ms added (debug build). Co-Authored-By: Claude Opus 5.5 --- gateway/src/lib.rs | 1 + gateway/src/main.rs | 7 +- gateway/src/server.rs | 53 +++ gateway/tests/openai_realtime_perf.rs | 583 ++++++++++++++++++++++++++ 4 files changed, 641 insertions(+), 3 deletions(-) create mode 100644 gateway/src/server.rs create mode 100644 gateway/tests/openai_realtime_perf.rs diff --git a/gateway/src/lib.rs b/gateway/src/lib.rs index b0d47766..c426b2cf 100644 --- a/gateway/src/lib.rs +++ b/gateway/src/lib.rs @@ -25,6 +25,7 @@ pub mod middleware; pub mod observability; pub mod plugin; pub mod routes; +pub mod server; pub mod state; pub mod utils; diff --git a/gateway/src/main.rs b/gateway/src/main.rs index 85b522c9..78176996 100644 --- a/gateway/src/main.rs +++ b/gateway/src/main.rs @@ -539,7 +539,8 @@ async fn main() -> anyhow::Result<()> { }); // Create and run the server - axum_server::bind_rustls(socket_addr, rustls_config) + // TCP_NODELAY on accept: see `server` (FRD-023 TC-PERF-01). + waav_gateway::server::tls_server(socket_addr, rustls_config) .handle(handle) .serve(app.into_make_service_with_connect_info::()) .await @@ -549,9 +550,9 @@ async fn main() -> anyhow::Result<()> { let listener = TcpListener::bind(&socket_addr).await?; - // Use axum::serve with graceful shutdown + // Use axum::serve with graceful shutdown; TCP_NODELAY on accept (FRD-023 TC-PERF-01). axum::serve( - listener, + waav_gateway::server::nodelay_listener(listener), app.into_make_service_with_connect_info::(), ) .with_graceful_shutdown(async move { diff --git a/gateway/src/server.rs b/gateway/src/server.rs new file mode 100644 index 00000000..33dd7de3 --- /dev/null +++ b/gateway/src/server.rs @@ -0,0 +1,53 @@ +//! How the gateway accepts client connections. +//! +//! Every accepted socket sets `TCP_NODELAY`. A realtime client streams audio in 20 ms frames and +//! reads small events back; with Nagle on, each small write the gateway relays is held until the +//! client's next frame acknowledges the previous one, which adds a whole frame interval (~20 ms) +//! to every relayed frame (FRD-023 TC-PERF-01: p50 19.5 ms added, against a 5 ms budget). The +//! gateway's vendor sockets already set it; the client side did not. + +use std::net::SocketAddr; + +use axum::serve::{ListenerExt as _, TapIo}; +use axum_server::accept::NoDelayAcceptor; +use axum_server::tls_rustls::{RustlsAcceptor, RustlsConfig}; +use tokio::net::{TcpListener, TcpStream}; + +fn set_nodelay(tcp: &mut TcpStream) { + if let Err(e) = tcp.set_nodelay(true) { + tracing::warn!(error = %e, "could not set TCP_NODELAY on an accepted socket"); + } +} + +/// The plain-TCP listener the gateway serves on (`axum::serve`), with `TCP_NODELAY` on accept. +pub fn nodelay_listener(listener: TcpListener) -> TapIo { + listener.tap_io(set_nodelay as fn(&mut TcpStream)) +} + +/// The TLS server (`axum_server`), with `TCP_NODELAY` on accept. +pub fn tls_server( + addr: SocketAddr, + config: RustlsConfig, +) -> axum_server::Server> { + axum_server::bind(addr).acceptor(RustlsAcceptor::new(config).acceptor(NoDelayAcceptor::new())) +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::serve::Listener as _; + + /// TC-PERF-01's cause, pinned: an accepted client socket has Nagle off. + #[tokio::test] + async fn accepted_sockets_set_tcp_nodelay() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let mut listener = nodelay_listener(listener); + let _client = TcpStream::connect(addr).await.unwrap(); + let (accepted, _) = listener.accept().await; + assert!( + accepted.nodelay().unwrap(), + "TCP_NODELAY must be set on accept" + ); + } +} diff --git a/gateway/tests/openai_realtime_perf.rs b/gateway/tests/openai_realtime_perf.rs new file mode 100644 index 00000000..18147c92 --- /dev/null +++ b/gateway/tests/openai_realtime_perf.rs @@ -0,0 +1,583 @@ +//! FRD-023 TC-PERF-01 — the relay's added latency per audio frame, in process and exact. +//! +//! One client streams 24 kHz PCM16 audio as `input_audio_buffer.append` events at 20 ms frames +//! (960 bytes → 1280 base64 characters, real-time pacing). An echo vendor answers every append +//! at once with a `response.output_audio.delta` carrying the same audio, so both directions carry +//! a frame every 20 ms, as in a live conversation. The client times each frame from send to echo: +//! +//! * **direct** — client ↔ vendor; +//! * **relay** — client ↔ WaaV `/v1/realtime` (Bud mode, in-memory control plane) ↔ vendor. +//! +//! Direct and relay phases alternate (`PERF_ROUNDS`) so drift on a shared host hits both. The added +//! latency is the relay's percentile minus the direct path's at the same percentile; the FRD budget +//! is p99 < 5 ms (R-2). Every frame must come back (FR-EVT-4: audio is never dropped). +//! +//! Serving is `main.rs`'s (`axum::serve` on `waav_gateway::server::nodelay_listener`, +//! `connection_limit_middleware` in front) and the production `Timings` (ping 20 s, revalidate 30 s), with the process-global +//! Prometheus recorder installed as `AppState::new` does in production. +//! +//! Run explicitly (about two minutes): +//! +//! ```text +//! cargo test --release --no-default-features --features dag-routing,turn-ensemble,noise-filter,openapi \ +//! --test openai_realtime_perf -- --ignored --nocapture +//! ``` +//! +//! Knobs: `PERF_FRAMES` (measured frames per path, default 3000 = 60 s of audio), `PERF_WARMUP` +//! (discarded frames per phase, default 50), `PERF_ROUNDS` (default 3), `PERF_BACKGROUND` (extra +//! sessions streaming alongside the measured one on the same path, default 0), `PERF_BUDGET_MS` +//! (default 5). Where the shell environment does not reach the test binary (a builder container), +//! cargo's `--config 'env.PERF_FRAMES="600"'` sets them. +//! +//! First run (2026-09-28, before `server::nodelay_listener`): added p50 19.5 / p99 20.4 ms — Nagle on +//! the accepted client socket held each relayed frame for one frame interval. + +use std::net::SocketAddr; +use std::sync::{Arc, Mutex}; +use std::time::{Duration, Instant}; + +use base64::Engine as _; +use futures_util::{SinkExt, StreamExt}; +use serde_json::{Value as Json, json}; +use tokio::net::TcpListener; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; + +use waav_gateway::config::{DAGTimeoutsConfig, PluginConfig, ServerConfig}; +use waav_gateway::handlers::openai_realtime::{RealtimeRuntime, Timings}; + +// ============================================================================================= +// Fixtures (as `openai_realtime_relay.rs`) +// ============================================================================================= + +const KEY: &str = "bud_realtime_perf_test_key"; +const PROJECT: &str = "5b0c7e1d-0000-4000-8000-00000000ef01"; +const USER: &str = "5b0c7e1d-0000-4000-8000-00000000ef02"; +const API_KEY_ID: &str = "5b0c7e1d-0000-4000-8000-00000000ef03"; +const MODEL_ID: &str = "5b0c7e1d-0000-4000-8000-00000000ef04"; +const ENDPOINT_ID: &str = "f0f0f0f0-0000-4000-8000-00000000ef05"; +const ALIAS: &str = "rt-perf"; +const VENDOR_MODEL: &str = "gpt-realtime-2.1"; + +/// bud-auth's fixture ciphertext; the plaintext is its `PLAIN`. +const TEST_CREDENTIAL: &str = include_str!("../../bud-auth/tests/fixtures/test_cred_encrypted.hex"); + +/// 24 kHz mono PCM16 at 20 ms: 480 samples. +const SAMPLES_PER_FRAME: usize = 24_000 / 50; +const FRAME_INTERVAL: Duration = Duration::from_millis(20); + +fn test_pem() -> String { + std::fs::read_to_string(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../bud-auth/tests/fixtures/test_cred_private.pem" + )) + .expect("bud-auth's fixture key (git-ignored *.pem) must be present locally") +} + +/// The loopback escape hatch, so the vendor on 127.0.0.1 passes SSRF validation. +fn allow_loopback() { + static ONCE: std::sync::Once = std::sync::Once::new(); + ONCE.call_once(|| unsafe { std::env::set_var("WAAV_ALLOW_LOOPBACK_ENDPOINTS", "1") }); +} + +fn env_usize(name: &str, default: usize) -> usize { + std::env::var(name) + .ok() + .and_then(|v| v.trim().parse().ok()) + .unwrap_or(default) +} + +/// One 20 ms frame of a 300 Hz tone, base64 (the payload size a real client sends). +fn audio_frame_b64() -> String { + let mut pcm = Vec::with_capacity(SAMPLES_PER_FRAME * 2); + for i in 0..SAMPLES_PER_FRAME { + let s = (6000.0 * (2.0 * std::f64::consts::PI * 300.0 * i as f64 / 24_000.0).sin()) as i16; + pcm.extend_from_slice(&s.to_le_bytes()); + } + base64::engine::general_purpose::STANDARD.encode(pcm) +} + +// ============================================================================================= +// The echo vendor (OpenAI Realtime GA) +// ============================================================================================= + +/// `session.created` on connect, `session.update` → `session.updated`, and every +/// `input_audio_buffer.append` echoed at once as a `response.output_audio.delta` with the same +/// audio and `event_id` `echo_`. +async fn start_echo_vendor() -> SocketAddr { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((stream, _)) = listener.accept().await else { + return; + }; + // A vendor's edge answers small frames without Nagle delay. + let _ = stream.set_nodelay(true); + tokio::spawn(async move { + let Ok(mut ws) = tokio_tungstenite::accept_async(stream).await else { + return; + }; + let created = json!({"type": "session.created", "event_id": "evt_v0", + "session": {"id": "sess_vendor_perf", "object": "realtime.session", "type": "realtime", "model": VENDOR_MODEL}}); + if ws + .send(Message::Text(created.to_string().into())) + .await + .is_err() + { + return; + } + while let Some(Ok(msg)) = ws.next().await { + let Message::Text(t) = msg else { continue }; + let v: Json = serde_json::from_str(t.as_str()).unwrap_or(Json::Null); + let out = match v["type"].as_str().unwrap_or_default() { + "session.update" => json!({"type": "session.updated", "event_id": "evt_vu", + "session": {"id": "sess_vendor_perf", "model": VENDOR_MODEL, "type": v["session"]["type"]}}), + "input_audio_buffer.append" => { + let seq = v["event_id"] + .as_str() + .and_then(|e| e.strip_prefix("perf_")) + .unwrap_or("x"); + json!({"type": "response.output_audio.delta", "event_id": format!("echo_{seq}"), + "response_id": "resp_perf", "item_id": "item_perf", "output_index": 0, + "content_index": 0, "delta": v["audio"]}) + } + _ => continue, + }; + if ws + .send(Message::Text(out.to_string().into())) + .await + .is_err() + { + return; + } + } + }); + } + }); + addr +} + +// ============================================================================================= +// The gateway (Bud mode over an in-memory control plane, served as main.rs serves it) +// ============================================================================================= + +fn config() -> ServerConfig { + ServerConfig { + host: "127.0.0.1".to_string(), + port: 0, + tls: None, + livekit_url: "ws://localhost:7880".to_string(), + livekit_public_url: "http://localhost:7880".to_string(), + livekit_api_key: None, + livekit_api_secret: None, + deepgram_api_key: None, + elevenlabs_api_key: None, + google_credentials: None, + azure_speech_subscription_key: None, + azure_speech_region: None, + cartesia_api_key: None, + openai_api_key: None, + azure_openai_api_key: None, + azure_openai_endpoint: None, + grok_api_key: None, + inworld_api_key: None, + gemini_api_key: None, + ultravox_api_key: None, + speechmatics_api_key: None, + yandex_api_key: None, + yandex_folder_id: None, + assemblyai_api_key: None, + hume_api_key: None, + groq_api_key: None, + ibm_watson_api_key: None, + ibm_watson_instance_id: None, + ibm_watson_region: None, + aws_access_key_id: None, + aws_secret_access_key: None, + aws_region: None, + gnani_token: None, + gnani_access_key: None, + gnani_certificate_path: None, + recording_s3_bucket: None, + recording_s3_region: None, + recording_s3_endpoint: None, + recording_s3_access_key: None, + recording_s3_secret_key: None, + recording_s3_prefix: None, + cache_path: None, + cache_ttl_seconds: Some(3600), + auth_service_url: None, + auth_signing_key_path: None, + auth_api_secrets: Vec::new(), + auth_timeout_seconds: 5, + auth_required: false, + sip: None, + cors_allowed_origins: None, + rate_limit_requests_per_second: 60, + rate_limit_burst_size: 10, + max_websocket_connections: None, + max_connections_per_ip: 1000, + ws_processing_timeout_secs: 10, + realtime_processing_timeout_secs: 30, + sip_max_participants: 3, + realtime_endpoint_overrides: Default::default(), + plugins: PluginConfig::default(), + dag_timeouts: DAGTimeoutsConfig::default(), + aliases: Default::default(), + } +} + +/// A realtime deployment on the vendor, configured as the live harness configures one. +fn rt_entry(vendor: SocketAddr) -> Json { + json!({ + "vendor": "openai", + "api_base": format!("http://{vendor}/v1"), + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["realtime_session"], + "model": VENDOR_MODEL, + "pricing": {"unit": "token", "per_units": 1000000, "currency": "USD", + "rates": {"input_text": 4.0, "input_audio": 32.0, "input_image": 5.0, + "cached_input_text": 0.4, "cached_input_audio": 0.4, "cached_input_image": 0.5, + "output_text": 24.0, "output_audio": 64.0, "transcription_per_minute": 0.003}}, + "config": {"realtime": { + "defaults": {"voice": "marin", "instructions": "You are the TC-PERF-01 agent.", + "turn_detection": {"type": "server_vad"}}, + "limits": {"max_session_seconds": 3600, "idle_timeout_seconds": 300} + }} + }) +} + +async fn start_gateway(vendor: SocketAddr) -> SocketAddr { + allow_loopback(); + let store = Arc::new(bud_auth::MemoryStore::new()); + store.set( + &format!("api_key:{}", bud_auth::hash_api_key(KEY)), + &json!({ + ALIAS: {"endpoint_id": ENDPOINT_ID, "model_id": MODEL_ID, "project_id": PROJECT, "kind": "model"}, + "__metadata__": {"api_key_id": API_KEY_ID, "user_id": USER, "api_key_project_id": PROJECT} + }) + .to_string(), + ); + store.set( + &format!("voice_table:{ENDPOINT_ID}"), + &json!({ ENDPOINT_ID: rt_entry(vendor) }).to_string(), + ); + let plane = Arc::new(bud_auth::BudPlane::with_decryptor( + store.clone() as Arc, + None, + bud_auth::CredentialDecryptor::from_pem(&test_pem()).unwrap(), + )); + plane.boot().await.unwrap(); + + let mut state = waav_gateway::state::AppState::new(config()).await; + { + let s = Arc::get_mut(&mut state).expect("unshared"); + s.bud_mode = Some(waav_gateway::auth::bud_mode::BudMode::for_plane(plane.clone()).unwrap()); + s.policies = Some(waav_gateway::core::deployment_policy::DeploymentPolicies::local()); + s.realtime = Arc::new(RealtimeRuntime { + timings: Timings::default(), + client_secret_keys: None, + }); + } + let app = waav_gateway::routes::openai_realtime::create_openai_realtime_router() + .layer(axum::middleware::from_fn_with_state( + state.clone(), + waav_gateway::middleware::connection_limit_middleware, + )) + .with_state(state.clone()); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let service = app.into_make_service_with_connect_info::(); + tokio::spawn(async move { + // Served exactly as main.rs serves it (the production listener, TCP_NODELAY on accept). + axum::serve(waav_gateway::server::nodelay_listener(listener), service) + .await + .unwrap(); + drop(plane); + }); + addr +} + +// ============================================================================================= +// The client +// ============================================================================================= + +#[derive(Clone, Copy, Debug, PartialEq)] +enum Path { + Direct, + Relay, +} + +struct Phase { + rtts: Vec, + sent: usize, + received: usize, +} + +/// One session: connect, configure, stream `warmup + frames` appends at 20 ms, time each echo. +async fn stream_session( + path: Path, + vendor: SocketAddr, + gateway: SocketAddr, + warmup: usize, + frames: usize, + audio: Arc, +) -> Phase { + let url = match path { + Path::Direct => format!("ws://{vendor}/v1/realtime?model={VENDOR_MODEL}"), + Path::Relay => format!("ws://{gateway}/v1/realtime?model={ALIAS}"), + }; + let mut req = url.into_client_request().unwrap(); + req.headers_mut() + .insert("authorization", format!("Bearer {KEY}").parse().unwrap()); + // Browsers and asyncio clients set TCP_NODELAY on their sockets. + let (ws, _) = tokio::time::timeout( + Duration::from_secs(15), + tokio_tungstenite::connect_async_with_config(req, None, true), + ) + .await + .expect("handshake within 15 s") + .unwrap_or_else(|e| panic!("{path:?} connect failed: {e}")); + let (mut sink, mut stream) = ws.split(); + + // session.created, then the client's own session.update and its answer. + let mut seen_created = false; + while !seen_created { + match tokio::time::timeout(Duration::from_secs(10), stream.next()).await { + Ok(Some(Ok(Message::Text(t)))) => { + let v: Json = serde_json::from_str(t.as_str()).unwrap(); + seen_created = v["type"] == "session.created"; + } + Ok(Some(Ok(_))) => {} + other => panic!("{path:?}: no session.created: {other:?}"), + } + } + let update = json!({"type": "session.update", "event_id": "perf_cfg", + "session": {"type": "realtime", "audio": {"input": {"turn_detection": null}}}}); + sink.send(Message::Text(update.to_string().into())) + .await + .unwrap(); + loop { + match tokio::time::timeout(Duration::from_secs(10), stream.next()).await { + Ok(Some(Ok(Message::Text(t)))) => { + let v: Json = serde_json::from_str(t.as_str()).unwrap(); + if v["type"] == "session.updated" { + break; + } + assert_ne!(v["type"], "error", "{path:?}: {v}"); + } + Ok(Some(Ok(_))) => {} + other => panic!("{path:?}: no session.updated: {other:?}"), + } + } + + let total = warmup + frames; + let sent_at: Arc>>> = Arc::new(Mutex::new(vec![None; total])); + let sender = { + let sent_at = sent_at.clone(); + let audio = audio.clone(); + tokio::spawn(async move { + let mut tick = tokio::time::interval(FRAME_INTERVAL); + tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); + for seq in 0..total { + tick.tick().await; + let frame = format!( + r#"{{"type":"input_audio_buffer.append","event_id":"perf_{seq}","audio":"{audio}"}}"# + ); + sent_at.lock().unwrap()[seq] = Some(Instant::now()); + if sink.send(Message::Text(frame.into())).await.is_err() { + return (sink, seq); + } + } + (sink, total) + }) + }; + + let mut rtts = Vec::with_capacity(frames); + let mut received = 0usize; + let deadline = + tokio::time::Instant::now() + FRAME_INTERVAL * total as u32 + Duration::from_secs(10); + while received < total { + match tokio::time::timeout_at(deadline, stream.next()).await { + Ok(Some(Ok(Message::Text(t)))) => { + let now = Instant::now(); + let v: Json = serde_json::from_str(t.as_str()).unwrap(); + if v["type"] != "response.output_audio.delta" { + assert_ne!(v["type"], "error", "{path:?}: {v}"); + continue; + } + let Some(seq) = v["event_id"] + .as_str() + .and_then(|e| e.strip_prefix("echo_")) + .and_then(|s| s.parse::().ok()) + else { + continue; + }; + received += 1; + let sent = sent_at.lock().unwrap()[seq].expect("echo before send"); + if seq >= warmup { + rtts.push(now - sent); + } + } + Ok(Some(Ok(_))) => {} + Ok(Some(Err(e))) => panic!("{path:?}: read error after {received} frames: {e}"), + Ok(None) => panic!("{path:?}: closed after {received} frames"), + Err(_) => break, + } + } + let (mut sink, sent) = sender.await.unwrap(); + let _ = sink.send(Message::Close(None)).await; + let _ = sink.close().await; + Phase { + rtts, + sent, + received, + } +} + +fn percentile(sorted: &[Duration], p: f64) -> Duration { + if sorted.is_empty() { + return Duration::ZERO; + } + let rank = ((p / 100.0) * sorted.len() as f64).ceil() as usize; + sorted[rank.clamp(1, sorted.len()) - 1] +} + +fn ms(d: Duration) -> f64 { + d.as_secs_f64() * 1e3 +} + +const PCTS: [f64; 5] = [50.0, 90.0, 99.0, 99.9, 100.0]; + +fn row(label: &str, sorted: &[Duration]) -> String { + let mean = sorted.iter().map(|d| d.as_secs_f64()).sum::() / sorted.len().max(1) as f64; + let cells: Vec = PCTS + .iter() + .map(|p| format!("{:>8.3}", ms(percentile(sorted, *p)))) + .collect(); + format!("{label:<8}{}{:>9.3}", cells.join(""), mean * 1e3) +} + +// ============================================================================================= +// TC-PERF-01 +// ============================================================================================= + +/// TC-PERF-01 — added latency per 20 ms frame of 24 kHz audio, client → vendor → client, relay +/// vs. a direct connection to the same vendor: p99 added < 5 ms, no frame lost. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +#[ignore = "performance: run explicitly with --ignored --nocapture (about two minutes)"] +async fn tc_perf_01_relay_added_latency_per_frame() { + let frames = env_usize("PERF_FRAMES", 3000); + let warmup = env_usize("PERF_WARMUP", 50); + let rounds = env_usize("PERF_ROUNDS", 3).max(1); + let background = env_usize("PERF_BACKGROUND", 0); + let budget_ms = env_usize("PERF_BUDGET_MS", 5) as f64; + let per_round = frames.div_ceil(rounds); + + let vendor = start_echo_vendor().await; + let gateway = start_gateway(vendor).await; + let audio = Arc::new(audio_frame_b64()); + assert_eq!( + audio.len(), + 1280, + "20 ms of 24 kHz PCM16 is 1280 base64 characters" + ); + + let mut all: [(Path, Vec, usize, usize); 2] = [ + (Path::Direct, Vec::new(), 0, 0), + (Path::Relay, Vec::new(), 0, 0), + ]; + for round in 0..rounds { + for slot in all.iter_mut() { + let path = slot.0; + // Background sessions stream on the same path for the whole phase (plus slack). + let bg: Vec<_> = (0..background) + .map(|_| { + tokio::spawn(stream_session( + path, + vendor, + gateway, + 0, + per_round + warmup + 100, + audio.clone(), + )) + }) + .collect(); + let phase = + stream_session(path, vendor, gateway, warmup, per_round, audio.clone()).await; + for h in bg { + let b = h.await.unwrap(); + assert_eq!( + b.received, b.sent, + "{path:?} background session lost frames" + ); + } + eprintln!( + "round {}/{rounds} {path:?}: {} frames, p50 {:.3} ms, p99 {:.3} ms", + round + 1, + phase.received, + ms(percentile(&sorted(&phase.rtts), 50.0)), + ms(percentile(&sorted(&phase.rtts), 99.0)), + ); + slot.1.extend(phase.rtts); + slot.2 += phase.sent; + slot.3 += phase.received; + } + } + + let direct = sorted(&all[0].1); + let relay = sorted(&all[1].1); + let header = PCTS + .iter() + .map(|p| { + if *p == 100.0 { + format!("{:>8}", "max") + } else { + format!("{:>8}", format!("p{p}")) + } + }) + .collect::(); + let added: Vec = PCTS + .iter() + .map(|p| { + format!( + "{:>8.3}", + ms(percentile(&relay, *p)) - ms(percentile(&direct, *p)) + ) + }) + .collect(); + let added_p50 = ms(percentile(&relay, 50.0)) - ms(percentile(&direct, 50.0)); + let added_p99 = ms(percentile(&relay, 99.0)) - ms(percentile(&direct, 99.0)); + eprintln!( + "\nTC-PERF-01 round trip per 20 ms frame (24 kHz PCM16, 1280 B base64), {} measured frames per path \ + ({rounds} rounds, {warmup} warm-up frames per phase, {background} background sessions), {} build, \ + gateway accept as main.rs (TCP_NODELAY on accept)\n\ + ms {header} mean\n{}\n{}\nadded {}\n\ + => added latency p50 {added_p50:.3} ms, p99 {added_p99:.3} ms (budget p99 < {budget_ms} ms)\n", + relay.len(), + if cfg!(debug_assertions) { + "debug" + } else { + "release" + }, + row("direct", &direct), + row("relay", &relay), + added.join(""), + ); + + for (path, rtts, sent, received) in &all { + assert_eq!(received, sent, "{path:?}: frames lost"); + assert_eq!(rtts.len(), per_round * rounds, "{path:?}: measured frames"); + } + assert!( + added_p99 < budget_ms, + "TC-PERF-01: relay p99 adds {added_p99:.3} ms over the direct path (budget {budget_ms} ms)" + ); +} + +fn sorted(v: &[Duration]) -> Vec { + let mut s = v.to_vec(); + s.sort_unstable(); + s +} From c66776904f3c716ba7e15db1f7b8044ee8c4d187 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 07:30:19 +0530 Subject: [PATCH 13/17] fix(gateway): a per-minute vendor's time is metered even without a duration price A Deepgram Voice Agent deployment published with no price held a 72 s session on pde-ditto and left no usage record: the facade only segmented a per-minute vendor when the deployment carried a minute/second price. The vendor bills by time whatever Bud charges, so per-minute vendors always meter duration segments; a segment with no minute/second price is recorded with its billed seconds, no cost, and unpriced_components = "duration" (never $0). Tests: realtime_cost an_unpriced_duration_segment_names_what_is_unpriced; realtime_translate tc_xl_07_an_unpriced_per_minute_vendor_still_meters_its_time (red with the old condition restored). Co-Authored-By: Claude Opus 5.5 --- gateway/src/core/realtime_cost.rs | 27 +++++++- .../src/handlers/openai_realtime/facade.rs | 6 +- gateway/tests/realtime_translate.rs | 66 +++++++++++++++++++ 3 files changed, 93 insertions(+), 6 deletions(-) diff --git a/gateway/src/core/realtime_cost.rs b/gateway/src/core/realtime_cost.rs index 33076bf5..f83766ec 100644 --- a/gateway/src/core/realtime_cost.rs +++ b/gateway/src/core/realtime_cost.rs @@ -280,16 +280,26 @@ pub fn realtime_transcription_cost( /// A duration segment under a minute or second price (D-9: per 60 s, so a socket that drops at /// minute 40 was still billed for 39 minutes). `None` under any other unit. pub fn realtime_duration_cost(pricing: Option<&VoicePricing>, seconds: f64) -> RealtimeCost { + // Billed time with no minute/second price is named unpriced, never zero-priced: a per-minute + // vendor billed it whatever Bud charges (RT7). + let unpriced = || RealtimeCost { + cost: None, + unit: None, + unpriced: vec!["duration"], + }; let Some(pricing) = pricing else { - return RealtimeCost::none(); + return unpriced(); }; - if pricing.per_units == 0 || !seconds.is_finite() || seconds < 0.0 { + if !seconds.is_finite() || seconds < 0.0 { return RealtimeCost::none(); } + if pricing.per_units == 0 { + return unpriced(); + } let (units, unit) = match pricing.unit.as_str() { "minute" => (seconds / 60.0, "minute"), "second" => (seconds, "second"), - _ => return RealtimeCost::none(), + _ => return unpriced(), }; let cost = units * pricing.cost_per_unit / pricing.per_units as f64; RealtimeCost { @@ -309,6 +319,17 @@ mod tests { use super::*; use std::collections::BTreeMap; + /// A duration segment without a minute/second price is unpriced (`duration`), never $0 and + /// never silently costless: the vendor billed that time (RT7 per-minute vendors). + #[test] + fn an_unpriced_duration_segment_names_what_is_unpriced() { + for pricing in [None, Some(token_price(&[("input_audio", 1.0)]))] { + let c = realtime_duration_cost(pricing.as_ref(), 30.0); + assert_eq!(c.cost, None); + assert_eq!(c.unpriced, vec!["duration"]); + } + } + fn token_price(rates: &[(&str, f64)]) -> VoicePricing { VoicePricing { unit: "token".into(), diff --git a/gateway/src/handlers/openai_realtime/facade.rs b/gateway/src/handlers/openai_realtime/facade.rs index f82726c8..f014aff0 100644 --- a/gateway/src/handlers/openai_realtime/facade.rs +++ b/gateway/src/handlers/openai_realtime/facade.rs @@ -1518,9 +1518,9 @@ impl Shell<'_> { let rates = provider.audio_rates(); self.provider = Some(provider); let now = Instant::now(); - if self.plan.info.per_minute - && crate::core::realtime_cost::bills_duration(self.p.endpoint.pricing.as_ref()) - { + // A per-minute vendor bills by time whatever Bud charges, so its time is always metered — + // priced when the deployment has a minute/second price, marked unpriced otherwise. + if self.plan.info.per_minute { self.segments = Some(SegmentClock::start(now, self.timings.segment)); } let acts = self.tr.connected(rates); diff --git a/gateway/tests/realtime_translate.rs b/gateway/tests/realtime_translate.rs index 079cd9dc..2dda482b 100644 --- a/gateway/tests/realtime_translate.rs +++ b/gateway/tests/realtime_translate.rs @@ -1336,6 +1336,72 @@ impl AgentMock { } } +/// TC-XL-07 🔒 — a per-minute vendor with NO duration price still meters its time: the vendor +/// bills by the minute whatever Bud charges, so the segments are recorded with their seconds and +/// marked unpriced (`duration`) — never a session that used 72 s of vendor time and left no record. +#[tokio::test] +async fn tc_xl_07_an_unpriced_per_minute_vendor_still_meters_its_time() { + let cap = Capture::install(); + let mock = AgentMock::start(json!({"type": "Welcome", "request_id": "dg-req-2"})).await; + let entry = json!({ + "vendor": "deepgram_voice_agent", + "api_base": mock.base(), + "credential": TEST_CREDENTIAL.trim(), + "endpoints": ["realtime_session"], + "model": "gpt-4o-mini" + }); + let id = "c7c7c7c7-0000-4000-8000-000000000180"; + let gw = gateway(Setup { + endpoints: vec![ep("agent", id, entry)], + ..Default::default() + }) + .await; + let mut c = connect(&gw, "agent").await; + until_type(&mut c, "session.created").await; + send( + &mut c, + json!({"type": "session.update", "session": {"instructions": "hello"}}), + ) + .await; + until_type(&mut c, "session.updated").await; + keep_alive(&mut c, Duration::from_millis(2500)).await; + close_and_drain(&mut c).await; + + let sessions = cap.wait_for("voice.session", 1).await; + let session = sessions + .iter() + .find(|s| text(s, "bud.endpoint_id").as_deref() == Some(id)) + .unwrap(); + let segments: Vec = cap + .spans() + .into_iter() + .filter(|s| s.name == "voice.turn" && text(s, "bud.endpoint_id").as_deref() == Some(id)) + .collect(); + assert!( + segments.len() >= 2, + "segments are recorded without a price: {}", + segments.len() + ); + for s in &segments { + assert_eq!( + text(s, "bud.voice.rt.component").as_deref(), + Some("duration_segment") + ); + assert!(number(s, "bud.voice.billed_seconds").unwrap() > 0.0); + assert!(number(s, "bud.voice.cost").is_none(), "unpriced, never $0"); + assert_eq!( + text(s, "bud.voice.unpriced_components").as_deref(), + Some("duration") + ); + } + let total: f64 = segments + .iter() + .map(|s| number(s, "bud.voice.billed_seconds").unwrap()) + .sum(); + assert!(total > 2.3 && total < 3.6, "{total}"); + assert!((total - number(session, "bud.voice.billed_seconds").unwrap()).abs() < 1e-9); +} + /// TC-XL-07 🔒 — a per-minute vendor bills DURATION SEGMENTS from the moment its connection /// opens: full segments while the session lives and the partial remainder at close (segments /// shortened to 1 s here, so ~2.5 s bills 1 + 1 + ~0.5; the 60 + 60 + 30 arithmetic is the From 1aa7b2194743030e377a7e552e6330514a89d86a Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 11:15:49 +0530 Subject: [PATCH 14/17] fix(gateway): a translated session with no turn_detection reports server VAD Every translate vendor detects turns itself unless told otherwise, but the facade echoed an unset turn_detection as null, which a GA client reads as push-to-talk: the playground showed a manual Send beside a Deepgram Voice Agent session that answers on its own. Unset now reports {type: server_vad}; an explicit null still reports off. Co-Authored-By: Claude Opus 5.5 --- .../src/handlers/openai_realtime/facade.rs | 11 +++++--- .../handlers/openai_realtime/facade/tests.rs | 26 +++++++++++++++++++ 2 files changed, 33 insertions(+), 4 deletions(-) diff --git a/gateway/src/handlers/openai_realtime/facade.rs b/gateway/src/handlers/openai_realtime/facade.rs index f014aff0..0eeba396 100644 --- a/gateway/src/handlers/openai_realtime/facade.rs +++ b/gateway/src/handlers/openai_realtime/facade.rs @@ -557,10 +557,13 @@ impl Translator { "audio": { "input": { "format": format, - "turn_detection": c.turn_detection.as_ref().map(|t| match t { - TurnDetectionConfig::None => Value::Null, - other => serde_json::to_value(other).unwrap_or(Value::Null), - }), + // Unset = the vendor's own turn detection, which every translate vendor + // runs by default; `null` only when the client turned it off. + "turn_detection": match c.turn_detection.as_ref() { + None => json!({"type": "server_vad"}), + Some(TurnDetectionConfig::None) => Value::Null, + Some(other) => serde_json::to_value(other).unwrap_or(Value::Null), + }, "transcription": c.input_audio_transcription.as_ref().map(|t| json!({"model": t.model})), }, "output": {"format": format, "voice": c.voice}, diff --git a/gateway/src/handlers/openai_realtime/facade/tests.rs b/gateway/src/handlers/openai_realtime/facade/tests.rs index 04b25bb2..5004d4e7 100644 --- a/gateway/src/handlers/openai_realtime/facade/tests.rs +++ b/gateway/src/handlers/openai_realtime/facade/tests.rs @@ -649,3 +649,29 @@ fn the_translate_vendors_are_c7s() { assert!(vendor_info("deepgram_voice_agent").unwrap().per_minute); assert!(!vendor_info("gemini").unwrap().per_minute); } + +/// Every translate vendor detects turns itself unless told otherwise, so a session that set no +/// `turn_detection` reports server VAD — `null` would tell a GA client (the playground) the session +/// is push-to-talk and put a manual "Send" in front of a user who only has to speak. An explicit +/// `turn_detection: null` is still reported as off. +#[test] +fn an_unset_turn_detection_is_reported_as_the_vendors_own_vad() { + for vendor in ["gemini", "nova_sonic", "deepgram_voice_agent", "elevenlabs_convai", "hume_evi"] { + let tr = translator(vendor); + let created: Value = serde_json::from_str(&tr.session_created()).unwrap(); + assert_eq!( + created["session"]["audio"]["input"]["turn_detection"], + json!({"type": "server_vad"}), + "{vendor}" + ); + } + let mut tr = translator("deepgram_voice_agent"); + let acts = tr.client( + &json!({"type": "session.update", + "session": {"audio": {"input": {"turn_detection": null}}}}) + .to_string(), + ); + let _ = acts; + let created: Value = serde_json::from_str(&tr.session_created()).unwrap(); + assert!(created["session"]["audio"]["input"]["turn_detection"].is_null()); +} From eae903aec1f5a993d1f0b619a2af48d0ca600b51 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 12:10:27 +0530 Subject: [PATCH 15/17] fix(gateway): space the sentences a translated vendor sends as separate finals MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Deepgram's Voice Agent (ConversationText) and Hume's EVI (assistant_message) send an answer as one finished sentence per message, and Nova Sonic as separate blocks, with no whitespace between them. The facade forwarded each verbatim as a response.output_audio_transcript.delta, so a GA client that appends deltas showed "…culture.Iconic landmarks…", and the transcript in response.output_audio_transcript.done read the same (seen live on a Deepgram Voice Agent deployment). A new segment after earlier text in the same response now gets one space, unless either side is already whitespace or the script does not space sentences (Chinese, Japanese). A streamed chunk continuing a segment stays verbatim, so Gemini's word chunks are unchanged. Also gives the ignored perf test the bedrock_http_client field the RT7 merge added, so --all-targets compiles again. Co-Authored-By: Claude Opus 5.5 --- .../src/handlers/openai_realtime/facade.rs | 25 ++++++++- .../handlers/openai_realtime/facade/tests.rs | 52 +++++++++++++++++-- gateway/tests/openai_realtime_perf.rs | 1 + 3 files changed, 73 insertions(+), 5 deletions(-) diff --git a/gateway/src/handlers/openai_realtime/facade.rs b/gateway/src/handlers/openai_realtime/facade.rs index 0eeba396..bf0c03db 100644 --- a/gateway/src/handlers/openai_realtime/facade.rs +++ b/gateway/src/handlers/openai_realtime/facade.rs @@ -1189,12 +1189,17 @@ impl Translator { return; }; let delta = if !is_final { + let delta = if r.interim_since_final.is_empty() { + spaced(&r.transcript, text) + } else { + text.to_string() + }; r.interim_since_final.push_str(text); - text.to_string() + delta } else { let seen = std::mem::take(&mut r.interim_since_final); if seen.is_empty() { - text.to_string() + spaced(&r.transcript, text) } else if let Some(rest) = text.strip_prefix(seen.as_str()) { rest.to_string() } else if seen.starts_with(text) { @@ -1820,5 +1825,21 @@ pub(super) async fn run( drop(p); } +/// The text of a new transcript segment, spaced from what the response already said. Deepgram's +/// Voice Agent and Hume's EVI send an answer as separate finished sentences, and Nova Sonic as +/// separate blocks, none with whitespace between them; a GA client just appends deltas. Scripts +/// that do not space sentences (Chinese, Japanese) are left as they are. +fn spaced(transcript: &str, text: &str) -> String { + let unspaced_script = |c: char| matches!(c, '\u{3000}'..='\u{30FF}' | '\u{3400}'..='\u{9FFF}' | '\u{FF00}'..='\u{FFEF}'); + match (transcript.chars().last(), text.chars().next()) { + (Some(last), Some(first)) + if !last.is_whitespace() && !first.is_whitespace() && !unspaced_script(last) => + { + format!(" {text}") + } + _ => text.to_string(), + } +} + #[cfg(test)] mod tests; diff --git a/gateway/src/handlers/openai_realtime/facade/tests.rs b/gateway/src/handlers/openai_realtime/facade/tests.rs index 5004d4e7..afc1d7ec 100644 --- a/gateway/src/handlers/openai_realtime/facade/tests.rs +++ b/gateway/src/handlers/openai_realtime/facade/tests.rs @@ -381,12 +381,52 @@ fn assistant_transcripts_are_sent_once_each() { acts.extend(tr.vendor(S2sEvent::ResponseDone { response_id: String::new(), })); - assert_eq!(transcript_deltas(&acts), vec!["Hi there.", "Bye."]); + assert_eq!(transcript_deltas(&acts), vec!["Hi there.", " Bye."]); let done = client_events(&acts) .into_iter() .find(|e| e["type"] == "response.output_audio_transcript.done") .unwrap(); - assert_eq!(done["transcript"], "Hi there.Bye."); + assert_eq!(done["transcript"], "Hi there. Bye."); +} + +/// Deepgram's Voice Agent and Hume's EVI send an answer as one finished sentence per message, with +/// no whitespace between them; the client appends deltas, so the facade supplies the separator. +/// Text the vendor already spaced, and a streamed chunk inside one segment, stay verbatim. +#[test] +fn sentences_sent_as_separate_finals_are_joined_with_a_space() { + let mut tr = translator("deepgram_voice_agent"); + ready(&mut tr); + let mut acts = tr.vendor(asst("Paris is the capital.", true)); + acts.extend(tr.vendor(asst("It is on the Seine.", true))); + acts.extend(tr.vendor(asst(" Its river has bridges.", true))); + acts.extend(tr.vendor(S2sEvent::ResponseDone { + response_id: String::new(), + })); + assert_eq!( + transcript_deltas(&acts), + vec![ + "Paris is the capital.", + " It is on the Seine.", + " Its river has bridges." + ] + ); + let done = client_events(&acts) + .into_iter() + .find(|e| e["type"] == "response.output_audio_transcript.done") + .unwrap(); + assert_eq!( + done["transcript"], + "Paris is the capital. It is on the Seine. Its river has bridges." + ); + + let mut tr = translator("hume_evi"); + ready(&mut tr); + let mut acts = tr.vendor(asst("東京です。", true)); + acts.extend(tr.vendor(asst("人口は多いです。", true))); + assert_eq!( + transcript_deltas(&acts), + vec!["東京です。", "人口は多いです。"] + ); } /// A function call ends its response (the client runs the tool, then asks for the next). @@ -656,7 +696,13 @@ fn the_translate_vendors_are_c7s() { /// `turn_detection: null` is still reported as off. #[test] fn an_unset_turn_detection_is_reported_as_the_vendors_own_vad() { - for vendor in ["gemini", "nova_sonic", "deepgram_voice_agent", "elevenlabs_convai", "hume_evi"] { + for vendor in [ + "gemini", + "nova_sonic", + "deepgram_voice_agent", + "elevenlabs_convai", + "hume_evi", + ] { let tr = translator(vendor); let created: Value = serde_json::from_str(&tr.session_created()).unwrap(); assert_eq!( diff --git a/gateway/tests/openai_realtime_perf.rs b/gateway/tests/openai_realtime_perf.rs index 18147c92..b434d69c 100644 --- a/gateway/tests/openai_realtime_perf.rs +++ b/gateway/tests/openai_realtime_perf.rs @@ -278,6 +278,7 @@ async fn start_gateway(vendor: SocketAddr) -> SocketAddr { s.realtime = Arc::new(RealtimeRuntime { timings: Timings::default(), client_secret_keys: None, + bedrock_http_client: None, }); } let app = waav_gateway::routes::openai_realtime::create_openai_realtime_router() From 06ac404369818c188e83cdb8111a64d722483d20 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 18:18:43 +0530 Subject: [PATCH 16/17] fix(gateway): send a deployment's "turn detection off" to OpenAI as null Bud stores turn detection off as {"type":"none"}. The relay forwarded it as-is in the defaults session.update, and OpenAI GA refuses the type ("Invalid value: 'none'. Supported values are: 'server_vad' and 'semantic_vad'", seen live on gpt-realtime-2.1-mini) -- which drops the whole update, so a deployment set to "None" silently lost its voice, instructions and every other default and ran on server VAD. GA turns detection off with null; send that. Co-Authored-By: Claude Opus 5.5 --- .../src/handlers/openai_realtime/policy.rs | 31 ++++++++++++++++++- 1 file changed, 30 insertions(+), 1 deletion(-) diff --git a/gateway/src/handlers/openai_realtime/policy.rs b/gateway/src/handlers/openai_realtime/policy.rs index 4f5bf423..52702aa1 100644 --- a/gateway/src/handlers/openai_realtime/policy.rs +++ b/gateway/src/handlers/openai_realtime/policy.rs @@ -434,7 +434,13 @@ pub fn defaults_update( input.insert("transcription".into(), Value::Object(transcription_cfg)); } if let Some(td) = &defaults.turn_detection { - input.insert("turn_detection".into(), td.clone()); + // Bud stores "off" as `{"type":"none"}`; GA turns detection off with `null` and refuses the + // type, which would drop the whole update. + let off = td.get("type").and_then(Value::as_str) == Some("none"); + input.insert( + "turn_detection".into(), + if off { Value::Null } else { td.clone() }, + ); } if let Some(nr) = &defaults.noise_reduction { input.insert("noise_reduction".into(), serde_json::json!({ "type": nr })); @@ -796,6 +802,29 @@ mod tests { assert!(s.get("model").is_none()); } + /// Bud stores "turn detection off" as `{"type":"none"}`; OpenAI GA turns it off with `null` and + /// refuses the type (`invalid_value: Supported values are: 'server_vad' and 'semantic_vad'`, seen live + /// on gpt-realtime-2.1-mini) -- which also drops every other default in the same update. + #[test] + fn turn_detection_none_is_sent_as_null() { + let settings: RealtimeSettings = serde_json::from_value(serde_json::json!({ + "session_type": "realtime", + "defaults": {"voice": "cedar", "turn_detection": {"type": "none"}} + })) + .unwrap(); + let update: Value = serde_json::from_str( + &defaults_update(Some(&settings), Some("gpt-realtime-2.1"), "e").unwrap(), + ) + .unwrap(); + let input = &update["session"]["audio"]["input"]; + assert!( + input.as_object().unwrap().contains_key("turn_detection"), + "off must be sent, not omitted: omitting keeps the vendor's server VAD" + ); + assert!(input["turn_detection"].is_null(), "{input}"); + assert_eq!(update["session"]["audio"]["output"]["voice"], "cedar"); + } + #[test] fn a_realtime_deployment_without_defaults_sends_no_update() { assert!(defaults_update(None, Some("m"), "e").is_none()); From 451967e1a22fe3d3b64051d98e59fa8b76a91038 Mon Sep 17 00:00:00 2001 From: dittops Date: Mon, 28 Sep 2026 21:00:53 +0530 Subject: [PATCH 17/17] fix(gateway): build without dag-routing, test keys from a clean checkout, openapi regenerated CI on #12: - bud_legs::bind_dag used crate::dag unconditionally, so every build without the dag-routing feature (the default matrix entry, server boot, VAD accuracy, openapi drift) failed to compile. It and its DAG-template tests are gated on dag-routing, as its caller already was. - The /ws leg tests read bud-auth's git-ignored fixture key, which a clean checkout does not have (26 failures). test_support now generates a key pair once per test binary and encrypts the test credential to it the way budapp does (RSA-OAEP-SHA-256, hex). - docs/openapi.yaml regenerated for RT6's /ws config (provider and base_url optional under the Bud control plane). Default features: 6865 lib tests pass; production features: all but the four ONNX tests that need the runtime library (green in CI); openapi drift passes. Co-Authored-By: Claude Opus 5.5 --- gateway/docs/openapi.yaml | 20 +++++---- gateway/src/handlers/ws/bud_legs.rs | 11 +++-- gateway/src/handlers/ws/config_handler.rs | 4 +- gateway/src/test_support.rs | 50 +++++++++++++++++------ 4 files changed, 58 insertions(+), 27 deletions(-) diff --git a/gateway/docs/openapi.yaml b/gateway/docs/openapi.yaml index 57105d94..300f58d1 100644 --- a/gateway/docs/openapi.yaml +++ b/gateway/docs/openapi.yaml @@ -965,7 +965,6 @@ components: history and barge-in. When absent, the gateway keeps its raw STT/TTS behavior (fully backward-compatible). required: - - base_url - model properties: allow_interruption: @@ -988,7 +987,10 @@ components: minimum: 0 base_url: type: string - description: OpenAI-compatible base URL for the LLM (e.g. `https://api.openai.com/v1`). + description: |- + OpenAI-compatible base URL for the LLM (e.g. `https://api.openai.com/v1`). Omitted — and + refused if present — under the Bud control plane, where the LLM leg is a Bud chat + deployment reached through the Bud gateway (FRD-023 RT6). example: https://api.openai.com/v1 degradation_message: type: @@ -2661,7 +2663,6 @@ components: type: object description: STT configuration for WebSocket messages (with optional API key) required: - - provider - language - sample_rate - channels @@ -2716,7 +2717,9 @@ components: example: nova-2 provider: type: string - description: Provider name (e.g., "deepgram") + description: |- + Provider name (e.g., "deepgram"). Under the Bud control plane it is taken from the + deployment `model` names, and may be omitted (FRD-023 RT6). example: deepgram punctuation: type: boolean @@ -2930,9 +2933,6 @@ components: TTSWebSocketConfig: type: object description: TTS configuration for WebSocket messages (with optional API key) - required: - - provider - - model properties: api_key: type: @@ -3049,7 +3049,7 @@ components: (Deepgram first) honor them END-TO-END (W1 keystone, closing BRUTAL_REVIEW.md S1/S5). model: type: string - description: Model to use for TTS + description: 'Model to use for TTS. Under the Bud control plane: the text-to-speech DEPLOYMENT.' example: aura-asteria-en pronunciations: type: array @@ -3058,7 +3058,9 @@ components: description: Pronunciation replacements to apply before TTS provider: type: string - description: Provider name (e.g., "deepgram", "hume", "elevenlabs") + description: |- + Provider name (e.g., "deepgram", "hume", "elevenlabs"). Under the Bud control plane it is + taken from the deployment `model` names, and may be omitted (FRD-023 RT6). example: deepgram request_timeout: type: diff --git a/gateway/src/handlers/ws/bud_legs.rs b/gateway/src/handlers/ws/bud_legs.rs index d23cb74d..6e0a857e 100644 --- a/gateway/src/handlers/ws/bud_legs.rs +++ b/gateway/src/handlers/ws/bud_legs.rs @@ -702,6 +702,7 @@ pub async fn prepare( } /// Headers a template may not send to the Bud gateway: the caller's own credential is the one. +#[cfg(feature = "dag-routing")] const LLM_CREDENTIAL_HEADERS: &[&str] = &[ "authorization", "api-key", @@ -722,6 +723,7 @@ const LLM_CREDENTIAL_HEADERS: &[&str] = &[ /// relay, and the native engines address deployments from RT7. /// /// Returns the admissions to hold for the session. +#[cfg(feature = "dag-routing")] pub async fn bind_dag( state: &Arc, definition: &mut crate::dag::definition::DAGDefinition, @@ -859,7 +861,7 @@ mod tests { //! deployment, which is what these assert; the live streams are the pde-ditto E2E. use super::*; - use crate::test_support::{TEST_CREDENTIAL, TEST_CREDENTIAL_PLAIN, bud_state_with_credentials}; + use crate::test_support::{TEST_CREDENTIAL_PLAIN, bud_state_with_credentials, test_credential}; use serde_json::{Value as Json, json}; use std::collections::HashMap; use std::sync::Mutex; @@ -881,7 +883,7 @@ mod tests { merge( json!({ "vendor": "deepgram", - "credential": TEST_CREDENTIAL.trim(), + "credential": test_credential(), "endpoints": ["audio_transcription"], "model": "nova-3", "language": "en-US", @@ -897,7 +899,7 @@ mod tests { merge( json!({ "vendor": "elevenlabs", - "credential": TEST_CREDENTIAL.trim(), + "credential": test_credential(), "endpoints": ["text_to_speech"], "model": "eleven_flash_v2_5", "voice": "JBFqnCBsd6RMkjVDRZzb", @@ -1542,6 +1544,7 @@ mod tests { // TC-WS-10 — DAG templates // --------------------------------------------------------------------------------------- + #[cfg(feature = "dag-routing")] fn template(nodes: Json) -> crate::dag::definition::DAGDefinition { serde_json::from_value(json!({ "id": "t", "name": "t", "nodes": nodes, "edges": [], @@ -1553,6 +1556,7 @@ mod tests { /// TC-WS-10 🔒 — a template's TTS node resolves a deployment (vendor, model, voice, /// credential, admission, revalidation); its LLM node goes to the Bud gateway with the /// caller's credential and none of the template's own. + #[cfg(feature = "dag-routing")] #[tokio::test] #[serial_test::serial] async fn tc_ws_10_template_nodes_address_deployments() { @@ -1643,6 +1647,7 @@ mod tests { /// TC-WS-10 — a template node naming a deployment the caller cannot reach, or a realtime /// provider node, is refused. + #[cfg(feature = "dag-routing")] #[tokio::test] #[serial_test::serial] async fn tc_ws_10_unreachable_and_realtime_nodes_are_refused() { diff --git a/gateway/src/handlers/ws/config_handler.rs b/gateway/src/handlers/ws/config_handler.rs index 3bfdc06c..b3c1a089 100644 --- a/gateway/src/handlers/ws/config_handler.rs +++ b/gateway/src/handlers/ws/config_handler.rs @@ -4174,14 +4174,14 @@ mod frd023_bud_mode_tests { async fn leg_plane(stt_extra: serde_json::Value) -> Arc { let mut stt = serde_json::json!({ - "vendor": "deepgram", "credential": crate::test_support::TEST_CREDENTIAL.trim(), + "vendor": "deepgram", "credential": crate::test_support::test_credential(), "endpoints": ["audio_transcription"], "model": "nova-3" }); for (k, v) in stt_extra.as_object().unwrap() { stt[k] = v.clone(); } let tts = serde_json::json!({ - "vendor": "elevenlabs", "credential": crate::test_support::TEST_CREDENTIAL.trim(), + "vendor": "elevenlabs", "credential": crate::test_support::test_credential(), "endpoints": ["text_to_speech"], "model": "eleven_flash_v2_5", "voice": "v1" }); let blob = serde_json::json!({ diff --git a/gateway/src/test_support.rs b/gateway/src/test_support.rs index ed10daa1..00ff74a2 100644 --- a/gateway/src/test_support.rs +++ b/gateway/src/test_support.rs @@ -93,9 +93,9 @@ pub(crate) async fn bud_state(keys: &[(&str, &str)]) -> Arc { state } -/// [`bud_state`] whose plane opens `voice_table` credentials with bud-auth's fixture key (the -/// plaintext of `test_cred_encrypted.hex` is `dg_vendor_key_abc123`) and enforces deployment -/// policies (FRD-023 RT6 tests). Returns the store, to mutate the control plane mid-test. +/// [`bud_state`] whose plane opens `voice_table` credentials with the test key pair +/// ([`test_credential`] decrypts to [`TEST_CREDENTIAL_PLAIN`]) and enforces deployment policies +/// (FRD-023 RT6 tests). Returns the store, to mutate the control plane mid-test. pub(crate) async fn bud_state_with_credentials( keys: &[(&str, &str)], ) -> (Arc, Arc) { @@ -103,15 +103,11 @@ pub(crate) async fn bud_state_with_credentials( for (k, v) in keys { store.set(k, v); } - let pem = std::fs::read_to_string(concat!( - env!("CARGO_MANIFEST_DIR"), - "/../bud-auth/tests/fixtures/test_cred_private.pem" - )) - .expect("bud-auth's fixture key (git-ignored *.pem) must be present locally"); let plane = Arc::new(bud_auth::BudPlane::with_decryptor( store.clone() as Arc, None, - bud_auth::CredentialDecryptor::from_pem(&pem).expect("fixture key parses"), + bud_auth::CredentialDecryptor::from_pem(&credential_fixture().0) + .expect("fixture key parses"), )); plane.boot().await.expect("plane boots"); let mut state = AppState::new(minimal_config()).await; @@ -123,8 +119,36 @@ pub(crate) async fn bud_state_with_credentials( (state, store) } -/// bud-auth's fixture ciphertext, for `voice_table` entries in tests. -pub(crate) const TEST_CREDENTIAL: &str = - include_str!("../../bud-auth/tests/fixtures/test_cred_encrypted.hex"); -/// Its plaintext. +/// The plaintext of [`test_credential`]. pub(crate) const TEST_CREDENTIAL_PLAIN: &str = "dg_vendor_key_abc123"; + +/// A test key pair (PKCS#8 PEM) and [`TEST_CREDENTIAL_PLAIN`] encrypted to it the way budapp encrypts +/// credentials (RSA-OAEP-SHA-256, hex), made once per test binary. bud-auth's `*.pem` fixtures are +/// git-ignored, so a clean checkout -- CI -- has no key to read and no way to open a committed +/// ciphertext. +fn credential_fixture() -> &'static (String, String) { + static FIXTURE: std::sync::OnceLock<(String, String)> = std::sync::OnceLock::new(); + FIXTURE.get_or_init(|| { + use rsa::pkcs8::{EncodePrivateKey, LineEnding}; + let mut rng = rsa::rand_core::OsRng; + let key = rsa::RsaPrivateKey::new(&mut rng, 2048).expect("test key generates"); + let pem = key + .to_pkcs8_pem(LineEnding::LF) + .expect("test key encodes") + .to_string(); + let ciphertext = key + .to_public_key() + .encrypt( + &mut rng, + rsa::Oaep::new::(), + TEST_CREDENTIAL_PLAIN.as_bytes(), + ) + .expect("test credential encrypts"); + (pem, hex::encode(ciphertext)) + }) +} + +/// The test credential's ciphertext, for `voice_table` entries in tests. +pub(crate) fn test_credential() -> &'static str { + &credential_fixture().1 +}