diff --git a/bud-auth/src/credentials.rs b/bud-auth/src/credentials.rs index 67afe32f..62d590bc 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, @@ -357,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"], _ => &[], @@ -741,6 +812,7 @@ mod tests { cost_per_unit: 0.0001, currency: Some("USD".into()), per_units: 1, + rates: BTreeMap::new(), }) ); } @@ -1176,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/bud-auth/src/endpoint_config.rs b/bud-auth/src/endpoint_config.rs index d3035c0c..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. @@ -185,6 +205,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 +329,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. @@ -265,11 +406,14 @@ const KNOWN_STT: &[&str] = &[ "alternatives", "sentiment", "noise_suppression", + "streaming", ]; 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 +430,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 +463,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)), } } } @@ -413,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/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..527ab08f 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,210 @@ 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..3276ac17 --- /dev/null +++ b/bud-auth/tests/realtime_contract.rs @@ -0,0 +1,197 @@ +//! 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()); +} diff --git a/gateway/Cargo.lock b/gateway/Cargo.lock index 7d943b92..a5a10ac2 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", @@ -7730,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", @@ -7766,6 +7831,7 @@ dependencies = [ "reqwest 0.12.28", "resil", "rhai", + "rsa", "rtrb", "rubato 0.16.2", "rustls 0.23.45", @@ -8091,7 +8157,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..966242af 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 @@ -329,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/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/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/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/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..f4b74b13 100644 --- a/gateway/src/auth/mod.rs +++ b/gateway/src/auth/mod.rs @@ -2,12 +2,13 @@ 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 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/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/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/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 399f8fbc..98dbffa6 100644 --- a/gateway/src/core/realtime/gemini/protocol.rs +++ b/gateway/src/core/realtime/gemini/protocol.rs @@ -42,17 +42,24 @@ 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"; -/// 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 @@ -239,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`). @@ -301,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 @@ -415,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); - } - - // 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]; + // 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); } - - // 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 { @@ -561,6 +655,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] @@ -813,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() { @@ -870,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 e626e3bd..7c76ce82 100644 --- a/gateway/src/core/realtime/nova_sonic/protocol.rs +++ b/gateway/src/core/realtime/nova_sonic/protocol.rs @@ -45,18 +45,20 @@ 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; -/// 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 | …). @@ -75,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` @@ -113,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 { @@ -196,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 { @@ -239,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()), }) } @@ -246,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 @@ -283,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) @@ -402,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, @@ -505,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). @@ -520,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], } } @@ -585,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 / @@ -672,6 +826,96 @@ 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() { + assert_eq!(DEFAULT_MODEL, "amazon.nova-2-sonic-v1:0"); + } + #[test] fn from_config_defaults_model_and_voice() { let cfg = RealtimeConfig { @@ -1088,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()); @@ -1097,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/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/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 cc2fe124..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 @@ -675,6 +775,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 +959,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 } @@ -839,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 } @@ -898,6 +1054,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/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/src/core/realtime_cost.rs b/gateway/src/core/realtime_cost.rs new file mode 100644 index 00000000..f83766ec --- /dev/null +++ b/gateway/src/core/realtime_cost.rs @@ -0,0 +1,500 @@ +//! 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 { + // 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 unpriced(); + }; + 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 unpriced(), + }; + 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; + + /// 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(), + 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/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 8eff6797..d77a9379 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, } } } @@ -235,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 { @@ -258,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(), } } @@ -359,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, ) @@ -435,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 61cc7ab3..c631636a 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,46 @@ 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)" + )) +} + +/// 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 @@ -511,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 { @@ -523,6 +593,7 @@ impl TTSProviderNode { model: None, config: serde_json::Value::Null, max_audio_bytes: DEFAULT_MAX_TTS_AUDIO_BYTES, + bud: None, } } @@ -547,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 @@ -631,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 { @@ -2433,3 +2528,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/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_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/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/facade.rs b/gateway/src/handlers/openai_realtime/facade.rs new file mode 100644 index 00000000..bf0c03db --- /dev/null +++ b/gateway/src/handlers/openai_realtime/facade.rs @@ -0,0 +1,1845 @@ +//! 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, + // 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}, + }, + }) + } + + 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 { + let delta = if r.interim_since_final.is_empty() { + spaced(&r.transcript, text) + } else { + text.to_string() + }; + r.interim_since_final.push_str(text); + delta + } else { + let seen = std::mem::take(&mut r.interim_since_final); + if seen.is_empty() { + spaced(&r.transcript, text) + } 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(); + // 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); + 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); +} + +/// 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 new file mode 100644 index 00000000..afc1d7ec --- /dev/null +++ b/gateway/src/handlers/openai_realtime/facade/tests.rs @@ -0,0 +1,723 @@ +//! 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."); +} + +/// 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). +#[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); +} + +/// 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()); +} 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..632746be --- /dev/null +++ b/gateway/src/handlers/openai_realtime/metering.rs @@ -0,0 +1,472 @@ +//! 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(" ")) +} + +/// 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, + 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(); + 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(), + ); + } + + /// 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_cost(&span, &cost); + if self.capture + && let Some(t) = transcript.filter(|t| !t.is_empty()) + { + 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. + } +} + +#[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 new file mode 100644 index 00000000..bad58a11 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/mod.rs @@ -0,0 +1,32 @@ +//! `/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). +//! * [`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; +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..52702aa1 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/policy.rs @@ -0,0 +1,864 @@ +//! 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 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"); + 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 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); + } + 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, + }, + /// 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), +} + +/// 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) + } + "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, + }, + "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 { + // 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 })); + } + 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(), + ) +} + +/// 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, + 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()); + } + + /// 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()); + } + + #[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..1f1a90a7 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/session.rs @@ -0,0 +1,1114 @@ +//! 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, SegmentClock, 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, + /// 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 { + 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), + connection_cap: None, + } + } +} + +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, + /// 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 { + /// 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()?, + bedrock_http_client: None, + }) + } +} + +/// 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)) +} + +/// 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, + pub endpoint_id: String, + pub endpoint_name: String, + pub endpoint: VoiceEndpoint, + pub alias: Option, + pub settings: Option, + pub rules: ClientRules, + /// 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, +} + +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 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()), + }; + 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()); + 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, + engine, + 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| 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. +#[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 { + pub(super) fn new(reason: &'static str, close_code: u16) -> Self { + Self { + reason, + close_code, + error: None, + } + } + + pub(super) fn with_error(mut self, code: &'static str, message: impl Into) -> Self { + self.error = Some((code, message.into())); + self + } +} + +pub(super) enum Outbound { + Frame(Message), + Close { + error: Option, + code: u16, + reason: String, + }, +} + +pub(super) fn next_event_id() -> String { + format!("evt_bud_{}", uuid::Uuid::new_v4().simple()) +} + +pub(super) 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). +pub(super) 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. +pub(super) 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 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, + 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 send_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.send_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 + } + } + } + + /// 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) { + 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.to_client(text).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?; + 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, + } + } + } + } +} + +/// Who a session bills (CONTRACTS C2), shared by both engines. +pub(super) fn attribution(p: &Prepared) -> 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: 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()), + } +} + +/// 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"); + + 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 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); + } + 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 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 { + state: &state, + p: &p, + meter, + client_tx: client_tx.clone(), + up_tx, + 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), + 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 segments = bills_duration.then(|| SegmentClock::start(now, timings.segment)); + 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.")), + _ = 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(()) + } + }; + if let Err(end) = step { + break end; + } + }; + + if let Some(clock) = segments { + // The final partial segment: a 150 s session bills 60 + 60 + 30 (TC-XL-07). + for secs in clock.close(Instant::now()) { + relay.meter.duration_segment(secs); + } + } + 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. +pub(super) 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..f9af4882 --- /dev/null +++ b/gateway/src/handlers/openai_realtime/upstream.rs @@ -0,0 +1,414 @@ +//! 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): 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>; + +/// 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, + /// 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 { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + 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!( + 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"), + Self::Misconfigured(why) => write!(f, "the deployment is misconfigured: {why}"), + } + } +} + +/// 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. +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}") + } 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() { + // 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(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-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() { + 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/handlers/realtime/handler.rs b/gateway/src/handlers/realtime/handler.rs index d855dbeb..7d4383b3 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 @@ -970,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() @@ -1040,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}; + use crate::core::realtime::InputTranscriptionConfig; - 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 - } - }); - - 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 { @@ -1116,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::*; @@ -1389,3 +1442,257 @@ 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")); + } + + /// 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/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..6e0a857e --- /dev/null +++ b/gateway/src/handlers/ws/bud_legs.rs @@ -0,0 +1,1699 @@ +//! `/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. +#[cfg(feature = "dag-routing")] +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. +#[cfg(feature = "dag-routing")] +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_PLAIN, bud_state_with_credentials, test_credential}; + 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(), + "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(), + "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 + // --------------------------------------------------------------------------------------- + + #[cfg(feature = "dag-routing")] + 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. + #[cfg(feature = "dag-routing")] + #[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. + #[cfg(feature = "dag-routing")] + #[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 b0fb9906..b3c1a089 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 @@ -235,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); @@ -252,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; @@ -305,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) => { @@ -514,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 { @@ -564,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; @@ -618,6 +727,212 @@ 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 +} + +/// 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 @@ -996,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); @@ -1031,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; }, ); })); @@ -1275,6 +1600,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, @@ -1294,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()); } @@ -1314,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: {}", @@ -1323,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 @@ -1420,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 @@ -2707,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 { @@ -2723,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, @@ -3452,8 +3823,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 +3841,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 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_api_key_falls_back_under_bud_mode() { + 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 +3901,7 @@ mod tests { None, "elevenlabs", "tts", - false, + true, &config_with_deepgram_key(), &tx, ) @@ -3671,3 +4072,297 @@ 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" + ); + } + + // ----------------------------------------------------------------------------------------- + // 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(), + "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(), + "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/lib.rs b/gateway/src/lib.rs index ee5ba53c..c426b2cf 100644 --- a/gateway/src/lib.rs +++ b/gateway/src/lib.rs @@ -25,9 +25,13 @@ pub mod middleware; pub mod observability; pub mod plugin; pub mod routes; +pub mod server; 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/main.rs b/gateway/src/main.rs index b353191f..78176996 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)); @@ -530,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 @@ -540,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/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/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() { 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/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/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/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/src/state/mod.rs b/gateway/src/state/mod.rs index 367bae0d..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, })) } @@ -534,6 +542,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 +553,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 +932,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); + } +} diff --git a/gateway/src/test_support.rs b/gateway/src/test_support.rs new file mode 100644 index 00000000..00ff74a2 --- /dev/null +++ b/gateway/src/test_support.rs @@ -0,0 +1,154 @@ +//! 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 +} + +/// [`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) { + let store = Arc::new(bud_auth::MemoryStore::new()); + for (k, v) in keys { + store.set(k, v); + } + let plane = Arc::new(bud_auth::BudPlane::with_decryptor( + store.clone() as Arc, + None, + 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; + { + 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) +} + +/// 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 +} 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:?}" + ); +} diff --git a/gateway/tests/openai_realtime_perf.rs b/gateway/tests/openai_realtime_perf.rs new file mode 100644 index 00000000..b434d69c --- /dev/null +++ b/gateway/tests/openai_realtime_perf.rs @@ -0,0 +1,584 @@ +//! 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, + bedrock_http_client: 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 +} diff --git a/gateway/tests/openai_realtime_relay.rs b/gateway/tests/openai_realtime_relay.rs new file mode 100644 index 00000000..1d81f99e --- /dev/null +++ b/gateway/tests/openai_realtime_relay.rs @@ -0,0 +1,2154 @@ +//! 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, + /// 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 { + 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(), + xai: false, + } + } +} + +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 = 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 + .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 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", + "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), + connection_cap: None, + } +} + +/// 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()), + bedrock_http_client: 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(); + 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")); + } +} + +// ============================================================================================= +// 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()); + } +} + +/// 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 +// ============================================================================================= + +/// 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/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 { 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; diff --git a/gateway/tests/realtime_translate.rs b/gateway/tests/realtime_translate.rs new file mode 100644 index 00000000..2dda482b --- /dev/null +++ b/gateway/tests/realtime_translate.rs @@ -0,0 +1,1506 @@ +//! 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 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 +/// 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" + ); + } +} 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",