feat(gateway): add Codex Live and OpenAI Realtime

Implement preflighted Live/Realtime WebSocket transports, protocol-aware authentication, usage auditing, UI filtering, and legacy Codex permission migration.
This commit is contained in:
ZheFox
2026-08-21 04:27:34 +08:00
parent fe38dcd294
commit 2c89202001
105 changed files with 7553 additions and 947 deletions
@@ -52,19 +52,27 @@ fn apply_admin_usage_status_filter(query: &mut UsageAuditListQuery, status: Opti
let Some(status) = status
.map(str::trim)
.filter(|candidate| !candidate.is_empty())
.map(str::to_ascii_lowercase)
else {
return;
};
match status {
"stream" => query.is_stream = Some(true),
"standard" => query.is_stream = Some(false),
match status.as_str() {
"stream" => {
query.is_stream = Some(true);
query.is_websocket = Some(false);
}
"standard" => {
query.is_stream = Some(false);
query.is_websocket = Some(false);
}
"websocket" | "ws" => query.is_websocket = Some(true),
"error" | "failed" => query.error_only = true,
"active" => {
query.statuses = Some(vec!["pending".to_string(), "streaming".to_string()]);
}
"pending" | "streaming" | "completed" | "cancelled" => {
query.statuses = Some(vec![status.to_string()]);
query.statuses = Some(vec![status]);
}
"has_fallback" | "has_retry" => {}
_ => {}
@@ -655,6 +663,7 @@ fn build_admin_usage_keyword_search_query(
statuses: base_query.statuses.clone(),
exclude_status_codes: base_query.exclude_status_codes.clone(),
is_stream: base_query.is_stream,
is_websocket: base_query.is_websocket,
error_only: base_query.error_only,
keywords,
matched_user_ids_by_keyword: search_context.matched_user_ids_by_keyword,
@@ -1027,7 +1036,10 @@ mod tests {
};
use serde_json::json;
use super::admin_usage_terminal_candidate_state_override;
use super::{
admin_usage_terminal_candidate_state_override, build_admin_usage_keyword_search_query,
build_admin_usage_records_query, AdminUsageSearchContext,
};
fn sample_candidate(
candidate_index: i32,
@@ -1127,4 +1139,49 @@ mod tests {
assert!(payload.is_none());
}
#[test]
fn admin_usage_transport_statuses_are_disjoint_in_list_and_keyword_queries() {
for status in ["websocket", "ws", "WS"] {
let raw_query = format!("status={status}");
let list_query =
build_admin_usage_records_query(100, 200, Some(&raw_query), None, None);
assert_eq!(list_query.is_websocket, Some(true));
assert_eq!(list_query.is_stream, None);
let keyword_query = build_admin_usage_keyword_search_query(
&list_query,
vec!["live".to_string()],
None,
AdminUsageSearchContext::default(),
false,
false,
None,
None,
);
assert_eq!(keyword_query.is_websocket, Some(true));
}
for (status, expected_stream) in [("stream", true), ("standard", false)] {
let raw_query = format!("status={status}");
let list_query =
build_admin_usage_records_query(100, 200, Some(&raw_query), None, None);
assert_eq!(list_query.is_stream, Some(expected_stream));
assert_eq!(list_query.is_websocket, Some(false));
let keyword_query = build_admin_usage_keyword_search_query(
&list_query,
vec!["live".to_string()],
None,
AdminUsageSearchContext::default(),
false,
false,
None,
None,
);
assert_eq!(keyword_query.is_stream, Some(expected_stream));
assert_eq!(keyword_query.is_websocket, Some(false));
}
}
}
@@ -10,6 +10,7 @@ use self::local::{
maybe_build_local_admin_proxy_response, maybe_build_local_internal_proxy_response,
};
pub(crate) use self::websocket::live::{live_websocket, maybe_handle_live_http};
pub(crate) use self::websocket::realtime::realtime_websocket;
pub(crate) use self::websocket::responses::responses_websocket;
use super::internal::resolve_local_proxy_execution_path;
pub(crate) use super::public::matches_model_mapping_for_models;
@@ -0,0 +1,395 @@
//! Session-level audit records for Codex Live transports.
//!
//! Frameless Bidi does not expose an authoritative token/cost usage object.
//! These records therefore capture exactly one bounded lifecycle summary per
//! connection and are explicitly void for billing. They never infer tokens,
//! audio duration, or cost from frame sizes.
use std::time::Duration;
use aether_ai_serving::AiStreamAttempt;
use aether_contracts::ExecutionPlan;
use aether_data_contracts::repository::usage::{
LIVE_SESSION_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY,
WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY,
};
use aether_usage_runtime::build_usage_event_data_seed;
use serde_json::{json, Map, Value};
use tracing::warn;
use crate::usage::{UsageEvent, UsageEventData, UsageEventType};
use crate::AppState;
const LIVE_AUDIT_WRITE_WAIT: Duration = Duration::from_secs(5);
const LIVE_AUDIT_SCHEMA_VERSION: &str = "1";
const LIVE_AUDIT_LOG_TARGET: &str = "aether_gateway::handlers::proxy::codex_live";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum LiveAuditTransport {
WebRtc,
DirectWebSocket,
Sideband,
}
impl LiveAuditTransport {
const fn transport(self) -> &'static str {
match self {
Self::WebRtc => "webrtc",
Self::DirectWebSocket => "websocket",
Self::Sideband => "sideband",
}
}
const fn mode(self) -> &'static str {
match self {
Self::WebRtc => "call_create",
Self::DirectWebSocket => "direct",
Self::Sideband => "sideband",
}
}
const fn websocket_transport(self) -> Option<&'static str> {
match self {
Self::WebRtc => None,
Self::DirectWebSocket => Some("codex_live_direct"),
Self::Sideband => Some("codex_live_sideband"),
}
}
}
/// Marks the existing synchronous SDP call-create audit row as an unmetered
/// WebRTC control exchange. The media leg bypasses Aether after this request.
pub(super) fn mark_live_call_create_report_context(report_context: &mut Option<Value>) {
attach_live_base_metadata(report_context, LiveAuditTransport::WebRtc);
}
fn attach_live_base_metadata(report_context: &mut Option<Value>, transport: LiveAuditTransport) {
let object = report_context_object(report_context);
object.insert(USAGE_AVAILABLE_METADATA_KEY.to_string(), Value::Bool(false));
object.insert(
USAGE_PRICING_AVAILABLE_METADATA_KEY.to_string(),
Value::Bool(false),
);
object.insert(
WEBSOCKET_MODE_METADATA_KEY.to_string(),
Value::Bool(transport.websocket_transport().is_some()),
);
if let Some(websocket_transport) = transport.websocket_transport() {
object.insert(
WEBSOCKET_TRANSPORT_METADATA_KEY.to_string(),
Value::String(websocket_transport.to_string()),
);
} else {
object.remove(WEBSOCKET_TRANSPORT_METADATA_KEY);
}
object.insert(
LIVE_SESSION_METADATA_KEY.to_string(),
json!({
"schema_version": LIVE_AUDIT_SCHEMA_VERSION,
"transport": transport.transport(),
"mode": transport.mode(),
"usage_state": "unavailable",
}),
);
}
fn report_context_object(report_context: &mut Option<Value>) -> &mut Map<String, Value> {
if !matches!(report_context, Some(Value::Object(_))) {
let seed = report_context.take();
let mut object = Map::new();
if let Some(seed) = seed.filter(|value| !value.is_null()) {
object.insert("seed".to_string(), seed);
}
*report_context = Some(Value::Object(object));
}
report_context
.as_mut()
.and_then(Value::as_object_mut)
.expect("Live audit report context was normalized to an object")
}
pub(super) struct LiveSessionAudit {
plan: ExecutionPlan,
report_context: Option<Value>,
transport: LiveAuditTransport,
}
impl LiveSessionAudit {
pub(super) fn from_attempt(attempt: &AiStreamAttempt, transport: LiveAuditTransport) -> Self {
let mut report_context = attempt.report_context.clone();
attach_live_base_metadata(&mut report_context, transport);
Self {
plan: attempt.plan.clone(),
report_context,
transport,
}
}
/// Persists one terminal lifecycle row. The spawned write remains alive if
/// the bounded caller wait elapses, so closing a socket cannot silently
/// cancel the only audit write for that connection.
pub(super) async fn finish(self, state: &AppState, terminal: LiveSessionTerminal) {
let request_id = self.plan.request_id.clone();
let event = self.build_terminal_event(terminal);
let usage_runtime = std::sync::Arc::clone(&state.usage_runtime);
let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state());
let task = tokio::spawn(async move {
usage_runtime
.record_terminal_event_direct(usage_data.as_ref(), event)
.await;
});
match tokio::time::timeout(LIVE_AUDIT_WRITE_WAIT, task).await {
Ok(Ok(())) => {}
Ok(Err(error)) => warn!(
target: LIVE_AUDIT_LOG_TARGET,
event_name = "codex_live_session_audit_task_failed",
log_type = "ops",
request_id,
error = %error,
"Codex Live session audit task failed"
),
Err(_) => warn!(
target: LIVE_AUDIT_LOG_TARGET,
event_name = "codex_live_session_audit_write_slow",
log_type = "ops",
request_id,
wait_ms = LIVE_AUDIT_WRITE_WAIT.as_millis() as u64,
write_detached = true,
"Codex Live stopped waiting for a slow session audit write"
),
}
}
fn build_terminal_event(self, terminal: LiveSessionTerminal) -> UsageEvent {
let mut data = build_usage_event_data_seed(&self.plan, self.report_context.as_ref());
data.request_type = Some("live".to_string());
data.is_stream = Some(self.transport != LiveAuditTransport::WebRtc);
data.status_code = Some(terminal.status_code);
data.response_time_ms = Some(terminal.elapsed_ms);
data.first_byte_time_ms = terminal.first_upstream_frame_ms;
data.input_tokens = None;
data.output_tokens = None;
data.total_tokens = None;
data.cache_creation_input_tokens = None;
data.cache_creation_ephemeral_5m_input_tokens = None;
data.cache_creation_ephemeral_1h_input_tokens = None;
data.cache_read_input_tokens = None;
data.cache_creation_cost_usd = None;
data.cache_read_cost_usd = None;
data.total_cost_usd = None;
data.actual_total_cost_usd = None;
if terminal.disposition != LiveSessionDisposition::Completed {
data.error_message = Some(terminal.termination.to_string());
data.error_category = Some(terminal.disposition.error_category().to_string());
}
data.request_metadata =
attach_terminal_metadata(data.request_metadata, self.transport, &terminal);
UsageEvent::new(
terminal.disposition.event_type(),
self.plan.request_id,
data,
)
}
}
fn attach_terminal_metadata(
metadata: Option<Value>,
transport: LiveAuditTransport,
terminal: &LiveSessionTerminal,
) -> Option<Value> {
let mut object = match metadata {
Some(Value::Object(object)) => object,
_ => Map::new(),
};
object.insert(USAGE_AVAILABLE_METADATA_KEY.to_string(), Value::Bool(false));
object.insert(
USAGE_PRICING_AVAILABLE_METADATA_KEY.to_string(),
Value::Bool(false),
);
object.insert(
WEBSOCKET_MODE_METADATA_KEY.to_string(),
Value::Bool(transport.websocket_transport().is_some()),
);
if let Some(websocket_transport) = transport.websocket_transport() {
object.insert(
WEBSOCKET_TRANSPORT_METADATA_KEY.to_string(),
Value::String(websocket_transport.to_string()),
);
}
object.insert(
LIVE_SESSION_METADATA_KEY.to_string(),
json!({
"schema_version": LIVE_AUDIT_SCHEMA_VERSION,
"transport": transport.transport(),
"mode": transport.mode(),
"state": terminal.disposition.state(),
"termination": terminal.termination,
"elapsed_ms": terminal.elapsed_ms,
"client_frames": terminal.client_frames,
"client_bytes": terminal.client_bytes,
"upstream_frames": terminal.upstream_frames,
"upstream_bytes": terminal.upstream_bytes,
"first_upstream_frame_ms": terminal.first_upstream_frame_ms,
"usage_state": "unavailable",
}),
);
Some(Value::Object(object))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum LiveSessionDisposition {
Completed,
Failed,
Cancelled,
}
impl LiveSessionDisposition {
const fn event_type(self) -> UsageEventType {
match self {
Self::Completed => UsageEventType::Completed,
Self::Failed => UsageEventType::Failed,
Self::Cancelled => UsageEventType::Cancelled,
}
}
const fn state(self) -> &'static str {
match self {
Self::Completed => "closed",
Self::Failed => "failed",
Self::Cancelled => "cancelled",
}
}
const fn error_category(self) -> &'static str {
match self {
Self::Completed => "none",
Self::Failed => "transport_error",
Self::Cancelled => "client_cancelled",
}
}
}
#[derive(Debug, Clone, Copy)]
pub(super) struct LiveSessionTerminal {
pub(super) disposition: LiveSessionDisposition,
pub(super) status_code: u16,
pub(super) termination: &'static str,
pub(super) elapsed_ms: u64,
pub(super) first_upstream_frame_ms: Option<u64>,
pub(super) client_frames: u64,
pub(super) client_bytes: u64,
pub(super) upstream_frames: u64,
pub(super) upstream_bytes: u64,
}
impl LiveSessionTerminal {
pub(super) const fn failure(
status_code: u16,
termination: &'static str,
elapsed_ms: u64,
) -> Self {
Self {
disposition: LiveSessionDisposition::Failed,
status_code,
termination,
elapsed_ms,
first_upstream_frame_ms: None,
client_frames: 0,
client_bytes: 0,
upstream_frames: 0,
upstream_bytes: 0,
}
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use aether_contracts::{ExecutionTimeouts, RequestBody};
use super::*;
fn sample_attempt() -> AiStreamAttempt {
AiStreamAttempt {
plan: ExecutionPlan {
request_id: "live-request".to_string(),
candidate_id: Some("candidate-live".to_string()),
provider_name: Some("Codex".to_string()),
provider_id: "provider-live".to_string(),
endpoint_id: "endpoint-live".to_string(),
key_id: "key-live".to_string(),
method: "GET".to_string(),
url: "wss://example.test/v1/live".to_string(),
headers: BTreeMap::new(),
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: true,
client_api_format: "codex:live".to_string(),
provider_api_format: "codex:live".to_string(),
model_name: Some("gpt-live".to_string()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts::default()),
},
report_kind: Some("openai_responses_stream".to_string()),
report_context: Some(json!({
"user_id": "user-live",
"api_key_id": "gateway-key-live",
"trace_id": "trace-live"
})),
}
}
#[test]
fn direct_terminal_audit_is_opaque_unmetered_and_void_eligible() {
let audit =
LiveSessionAudit::from_attempt(&sample_attempt(), LiveAuditTransport::DirectWebSocket);
let event = audit.build_terminal_event(LiveSessionTerminal {
disposition: LiveSessionDisposition::Completed,
status_code: 200,
termination: "client_close_frame",
elapsed_ms: 1234,
first_upstream_frame_ms: Some(42),
client_frames: 3,
client_bytes: 128,
upstream_frames: 5,
upstream_bytes: 512,
});
assert_eq!(event.event_type, UsageEventType::Completed);
assert_eq!(event.data.input_tokens, None);
assert_eq!(event.data.total_cost_usd, None);
let metadata = event.data.request_metadata.expect("metadata");
assert_eq!(metadata[USAGE_AVAILABLE_METADATA_KEY], false);
assert_eq!(metadata[USAGE_PRICING_AVAILABLE_METADATA_KEY], false);
assert_eq!(metadata[WEBSOCKET_MODE_METADATA_KEY], true);
assert_eq!(
metadata[WEBSOCKET_TRANSPORT_METADATA_KEY],
"codex_live_direct"
);
assert_eq!(metadata[LIVE_SESSION_METADATA_KEY]["client_frames"], 3);
assert_eq!(
metadata[LIVE_SESSION_METADATA_KEY]["usage_state"],
"unavailable"
);
}
#[test]
fn call_create_is_webrtc_not_websocket() {
let mut context = Some(json!({"trace_id": "trace-live"}));
mark_live_call_create_report_context(&mut context);
let context = context.expect("context");
assert_eq!(context[USAGE_AVAILABLE_METADATA_KEY], false);
assert_eq!(context[USAGE_PRICING_AVAILABLE_METADATA_KEY], false);
assert_eq!(context[WEBSOCKET_MODE_METADATA_KEY], false);
assert!(context.get(WEBSOCKET_TRANSPORT_METADATA_KEY).is_none());
assert_eq!(context[LIVE_SESSION_METADATA_KEY]["transport"], "webrtc");
}
}
@@ -22,6 +22,7 @@ use crate::execution_runtime::execute_execution_runtime_sync_plan_with_report_co
use crate::handlers::proxy::websocket::responses::ResponsesWebSocketTurnAdmission;
use crate::{AppState, GatewayError};
use super::audit::mark_live_call_create_report_context;
use super::live_usage_accounting_is_safe;
use super::planner::{live_call_url, plan_live_candidate, LiveAuthMode, LivePoolLeaseGuard};
use super::protocol::{build_live_multipart, extract_call_id_from_location, parse_live_multipart};
@@ -158,7 +159,7 @@ pub(crate) async fn maybe_handle_live_http(
.to_string(),
);
let Some(attempt) =
let Some(mut attempt) =
build_standard_sync_plan_from_decision(parts, &provider_body_marker, candidate.execution)?
else {
lease.release().await;
@@ -168,6 +169,10 @@ pub(crate) async fn maybe_handle_live_http(
"Codex Live provider request could not be built",
)?));
};
// The synchronous SDP exchange has an ordinary request lifecycle, but it
// does not contain the media leg's token/cost usage. Keep the existing row
// while making that boundary explicit and non-billable.
mark_live_call_create_report_context(&mut attempt.report_context);
if let Some(rejection) = execution_plan_balance_capacity_rejection(
state,
control_decision,
@@ -5,6 +5,7 @@
//! Keeping it in an independent module prevents a `session.update` frame from
//! ever entering the Responses `response.create` state machine.
mod audit;
mod http;
mod planner;
mod protocol;
@@ -1,11 +1,10 @@
//! Candidate planning and provider request shaping for Codex Live.
//!
//! Live deliberately reuses the existing Responses permission and scheduler
//! surface. Only the selected candidate, model alias and transport identity are
//! reused; Responses body normalization and its WebSocket state machine never
//! see a Live protocol frame.
//! Live has its own endpoint and permission surface. Candidate selection,
//! model aliases and transport policy are shared with the ordinary scheduler,
//! but Responses body normalization and its WebSocket state machine never see
//! a Live protocol frame.
use std::collections::BTreeSet;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
@@ -18,8 +17,9 @@ use sha2::{Digest, Sha256};
use url::{form_urlencoded, Url};
use crate::ai_serving::{
build_standard_stream_plan_from_decision, maybe_build_responses_websocket_decision,
AiExecutionDecision, AiStreamAttempt, ResponsesWebSocketPinnedCandidate,
build_standard_stream_plan_from_decision,
maybe_build_pinned_stream_local_same_format_provider_decision_payload, AiExecutionDecision,
AiStreamAttempt, ResponsesWebSocketPinnedCandidate,
};
use crate::control::GatewayControlDecision;
use crate::headers::request_origin_from_headers_and_remote_addr;
@@ -155,30 +155,59 @@ pub(super) async fn plan_live_candidate(
}
let parts = build_live_planning_parts(headers, remote_addr);
let body = json!({"model": client_model, "input": []});
let planned = maybe_build_responses_websocket_decision(
let execution = maybe_build_pinned_stream_local_same_format_provider_decision_payload(
state,
&parts,
trace_id,
decision,
None,
&body,
None::<&BTreeSet<String>>,
None::<&BTreeSet<String>>,
pinned_candidate,
crate::ai_serving::CODEX_LIVE_STREAM_PLAN_KIND,
pinned_candidate
.map(|pinned| (pinned.provider_id(), pinned.endpoint_id(), pinned.key_id())),
)
.await?;
let Some(planned) = planned else {
let Some(mut execution) = execution else {
return Ok(None);
};
if execution
.provider_api_format
.as_deref()
.map(crate::ai_serving::normalize_api_format_alias)
.as_deref()
!= Some("codex:live")
{
crate::orchestration::release_pool_key_lease_from_report_context(
state,
execution.report_context.as_ref(),
)
.await;
return Ok(None);
}
let Some(effective_auth_type) = execution
.report_context
.as_ref()
.and_then(|context| context.get("upstream_credential_mode"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
else {
crate::orchestration::release_pool_key_lease_from_report_context(
state,
execution.report_context.as_ref(),
)
.await;
return Ok(None);
};
let effective_auth_type = planned.effective_auth_type;
let mut execution = planned.execution;
let provider_type = execution
.provider_type
.as_deref()
.map(str::trim)
.unwrap_or_default();
if !provider_type.eq_ignore_ascii_case("codex") && !provider_type.eq_ignore_ascii_case("openai")
{
if !matches!(
provider_type.to_ascii_lowercase().as_str(),
"codex" | "openai" | "custom"
) {
crate::orchestration::release_pool_key_lease_from_report_context(
state,
execution.report_context.as_ref(),
@@ -247,7 +276,7 @@ pub(super) fn direct_live_websocket_url(
if candidate.auth_mode == LiveAuthMode::ChatGptOauth {
return Err(LiveProtocolError::OauthDirectWebSocketUnsupported);
}
replace_responses_suffix(
replace_live_suffix(
candidate.execution.upstream_url.as_deref(),
&["live"],
Some(("model", candidate.provider_model.as_str())),
@@ -257,12 +286,12 @@ pub(super) fn direct_live_websocket_url(
pub(super) fn live_call_url(candidate: &PlannedLiveCandidate) -> Result<String, LiveProtocolError> {
match candidate.auth_mode {
LiveAuthMode::ApiKey => {
replace_responses_suffix(candidate.execution.upstream_url.as_deref(), &["live"], None)
replace_live_suffix(candidate.execution.upstream_url.as_deref(), &["live"], None)
}
LiveAuthMode::ChatGptOauth => {
let source =
validated_official_chatgpt_url(candidate.execution.upstream_url.as_deref())?;
replace_responses_suffix(
replace_live_suffix(
Some(source.as_str()),
&["realtime", "calls"],
Some(("intent", "quicksilver")),
@@ -283,7 +312,7 @@ pub(super) fn live_sideband_url(
) -> Result<String, LiveProtocolError> {
super::protocol::validate_call_id(call_id)?;
match candidate.auth_mode {
LiveAuthMode::ApiKey => replace_responses_suffix(
LiveAuthMode::ApiKey => replace_live_suffix(
candidate.execution.upstream_url.as_deref(),
&["live", call_id],
None,
@@ -353,7 +382,7 @@ fn live_routing_fingerprint(
}
let path = url.path().trim_end_matches('/');
let path_family = path
.strip_suffix("/responses")
.strip_suffix("/live")
.ok_or(LiveProtocolError::InvalidUpstreamUrl)?;
if auth_mode == LiveAuthMode::ChatGptOauth {
validated_official_chatgpt_url(Some(raw_url))?;
@@ -419,15 +448,14 @@ fn validated_official_chatgpt_url(raw: Option<&str>) -> Result<Url, LiveProtocol
&& url.username().is_empty()
&& url.password().is_none()
&& url.fragment().is_none()
&& url.path().trim_end_matches('/').strip_suffix("/responses")
== Some("/backend-api/codex");
&& url.path().trim_end_matches('/').strip_suffix("/live") == Some("/backend-api/codex");
if !official {
return Err(LiveProtocolError::OauthUpstreamUnsupported);
}
Ok(url)
}
fn replace_responses_suffix(
fn replace_live_suffix(
raw: Option<&str>,
suffix: &[&str],
query: Option<(&str, &str)>,
@@ -441,7 +469,7 @@ fn replace_responses_suffix(
{
return Err(LiveProtocolError::InvalidUpstreamUrl);
}
if url.path_segments().and_then(Iterator::last) != Some("responses") {
if url.path_segments().and_then(Iterator::last) != Some("live") {
return Err(LiveProtocolError::InvalidUpstreamUrl);
}
{
@@ -524,8 +552,8 @@ fn build_live_planning_parts(
remote_addr: &SocketAddr,
) -> http::request::Parts {
let mut request = http::Request::builder()
.method(Method::POST)
.uri("/v1/responses")
.method(Method::GET)
.uri("/v1/live")
.body(())
.expect("the fixed Live planning request must be valid");
*request.headers_mut() = sanitize_live_planning_headers(headers.clone());
@@ -658,9 +686,9 @@ mod tests {
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:responses".to_string(),
api_family: Some("openai".to_string()),
endpoint_kind: Some("responses".to_string()),
api_format: "codex:live".to_string(),
api_family: Some("codex".to_string()),
endpoint_kind: Some("live".to_string()),
is_active: true,
base_url: "https://chatgpt.com/backend-api/codex".to_string(),
header_rules: None,
@@ -677,7 +705,7 @@ mod tests {
name: "key".to_string(),
auth_type: default_auth_type.to_string(),
is_active: true,
api_formats: Some(vec!["openai:responses".to_string()]),
api_formats: Some(vec!["codex:live".to_string()]),
auth_type_by_format,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
@@ -701,7 +729,7 @@ mod tests {
session_id: &str,
) -> AiExecutionDecision {
let mut decision = candidate(
"https://chatgpt.com/backend-api/codex/responses",
"https://chatgpt.com/backend-api/codex/live",
LiveAuthMode::ChatGptOauth,
)
.execution;
@@ -717,7 +745,7 @@ mod tests {
#[test]
fn format_auth_override_selects_the_effective_live_auth_mode() {
let overridden =
transport_with_auth_override("oauth", Some(json!({"openai:responses": "bearer"})));
transport_with_auth_override("oauth", Some(json!({"codex:live": "bearer"})));
let effective =
aether_provider_transport::auth::resolve_local_auth_type_for_transport_format(
&overridden,
@@ -744,7 +772,7 @@ mod tests {
#[test]
fn derives_api_key_live_urls_preserves_query_and_replaces_the_mapped_model() {
let mut candidate = candidate(
"https://api.example.test/v1/responses?api-version=2026-08-01&model=stale&MODEL=duplicate",
"https://api.example.test/v1/live?api-version=2026-08-01&model=stale&MODEL=duplicate",
LiveAuthMode::ApiKey,
);
candidate.provider_model = "upstream/model + future".to_string();
@@ -778,7 +806,7 @@ mod tests {
#[test]
fn derives_chatgpt_call_and_official_sideband_urls() {
let candidate = candidate(
"https://chatgpt.com/backend-api/codex/responses?api-version=2026-08-01&intent=stale&INTENT=duplicate&architecture=stale&ARCHITECTURE=duplicate",
"https://chatgpt.com/backend-api/codex/live?api-version=2026-08-01&intent=stale&INTENT=duplicate&architecture=stale&ARCHITECTURE=duplicate",
LiveAuthMode::ChatGptOauth,
);
let call = Url::parse(live_call_url(&candidate).unwrap().as_str()).unwrap();
@@ -816,7 +844,7 @@ mod tests {
#[test]
fn chatgpt_oauth_live_fails_closed_for_custom_backend_origins() {
let candidate = candidate(
"https://relay.example/backend-api/codex/responses",
"https://relay.example/backend-api/codex/live",
LiveAuthMode::ChatGptOauth,
);
assert_eq!(
@@ -861,8 +889,7 @@ mod tests {
#[test]
fn routing_fingerprint_binds_api_key_origin_without_hashing_the_token() {
let mut first =
candidate("https://api-a.example/v1/responses", LiveAuthMode::ApiKey).execution;
let mut first = candidate("https://api-a.example/v1/live", LiveAuthMode::ApiKey).execution;
first.provider_request_headers.extend([
("authorization".to_string(), "Bearer token-1".to_string()),
("x-session-id".to_string(), "session-1".to_string()),
@@ -877,7 +904,7 @@ mod tests {
);
let mut changed_origin =
candidate("https://api-b.example/v1/responses", LiveAuthMode::ApiKey).execution;
candidate("https://api-b.example/v1/live", LiveAuthMode::ApiKey).execution;
changed_origin
.provider_request_headers
.insert("x-session-id".to_string(), "session-1".to_string());
@@ -896,7 +923,7 @@ mod tests {
);
let missing_session =
candidate("https://api-a.example/v1/responses", LiveAuthMode::ApiKey).execution;
candidate("https://api-a.example/v1/live", LiveAuthMode::ApiKey).execution;
assert_eq!(
live_routing_fingerprint(&missing_session, "bearer", LiveAuthMode::ApiKey),
Err(LiveProtocolError::InvalidUpstreamUrl)
@@ -906,7 +933,7 @@ mod tests {
#[test]
fn routing_fingerprint_canonicalizes_safe_query_and_ignores_query_credentials() {
let mut baseline = candidate(
"https://api-a.example/v1/responses?api-version=2026-08-01&deployment=primary&alt=sse&token=secret-1&key=secret-1",
"https://api-a.example/v1/live?api-version=2026-08-01&deployment=primary&alt=sse&token=secret-1&key=secret-1",
LiveAuthMode::ApiKey,
)
.execution;
@@ -917,7 +944,7 @@ mod tests {
live_routing_fingerprint(&baseline, "bearer", LiveAuthMode::ApiKey).unwrap();
let mut reordered = candidate(
"https://api-a.example/v1/responses?key=secret-2&alt=sse&token=secret-2&deployment=primary&api-version=2026-08-01",
"https://api-a.example/v1/live?key=secret-2&alt=sse&token=secret-2&deployment=primary&api-version=2026-08-01",
LiveAuthMode::ApiKey,
)
.execution;
@@ -931,7 +958,7 @@ mod tests {
);
let mut changed_route = candidate(
"https://api-a.example/v1/responses?api-version=2026-08-01&deployment=secondary&alt=sse&token=secret-2&key=secret-2",
"https://api-a.example/v1/live?api-version=2026-08-01&deployment=secondary&alt=sse&token=secret-2&key=secret-2",
LiveAuthMode::ApiKey,
)
.execution;
@@ -989,10 +1016,7 @@ mod tests {
#[test]
fn live_urls_reject_credentials_invalid_suffixes_and_call_ids() {
let credentials = candidate(
"https://[email protected]/v1/responses",
LiveAuthMode::ApiKey,
);
let credentials = candidate("https://[email protected]/v1/live", LiveAuthMode::ApiKey);
assert_eq!(
direct_live_websocket_url(&credentials),
Err(LiveProtocolError::InvalidUpstreamUrl)
@@ -1012,7 +1036,7 @@ mod tests {
);
let fragment = candidate(
"https://api.example.test/v1/responses#not-sent-upstream",
"https://api.example.test/v1/live#not-sent-upstream",
LiveAuthMode::ApiKey,
);
assert_eq!(
File diff suppressed because it is too large Load Diff
@@ -8,6 +8,7 @@
pub(crate) mod ingress;
pub(crate) mod live;
pub(crate) mod realtime;
pub(crate) mod responses;
pub(crate) mod session;
pub(crate) mod transport;
@@ -0,0 +1,430 @@
//! One terminal usage/audit row per OpenAI Realtime WebSocket connection.
//!
//! Realtime exposes authoritative token usage on `response.done`. We preserve
//! those counters when present. A connection that closes without any such
//! usage is still visible as a lifecycle row, but is explicitly marked
//! unavailable and cannot participate in billing or balance materialization.
use std::time::Duration;
use aether_contracts::ExecutionPlan;
use aether_data_contracts::repository::usage::{
REALTIME_SESSION_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY,
USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY,
WEBSOCKET_TRANSPORT_METADATA_KEY,
};
use aether_usage_runtime::build_usage_event_data_seed;
use serde_json::{json, Map, Value};
use tracing::warn;
use crate::usage::{UsageEvent, UsageEventType};
use crate::AppState;
use super::protocol::RealtimeUsageTotals;
const REALTIME_AUDIT_WRITE_WAIT: Duration = Duration::from_secs(5);
const REALTIME_AUDIT_SCHEMA_VERSION: &str = "1";
const REALTIME_AUDIT_LOG_TARGET: &str = "aether_gateway::handlers::proxy::realtime_ws";
const REALTIME_WEBSOCKET_TRANSPORT: &str = "openai_realtime";
pub(super) struct RealtimeSessionAudit {
plan: ExecutionPlan,
report_context: Option<Value>,
}
impl RealtimeSessionAudit {
pub(super) fn new(plan: &ExecutionPlan, report_context: Option<&Value>) -> Self {
Self {
plan: plan.clone(),
report_context: report_context.cloned(),
}
}
/// Persist exactly one terminal row. If the bounded caller wait expires,
/// the spawned write remains alive instead of losing the only session row.
pub(super) async fn finish(self, state: &AppState, terminal: RealtimeSessionTerminal) {
let request_id = self.plan.request_id.clone();
let event = self.build_terminal_event(terminal);
let usage_runtime = std::sync::Arc::clone(&state.usage_runtime);
let usage_data = std::sync::Arc::clone(state.usage_lifecycle_data_state());
let task = tokio::spawn(async move {
usage_runtime
.record_terminal_event_direct(usage_data.as_ref(), event)
.await;
});
match tokio::time::timeout(REALTIME_AUDIT_WRITE_WAIT, task).await {
Ok(Ok(())) => {}
Ok(Err(error)) => warn!(
target: REALTIME_AUDIT_LOG_TARGET,
event_name = "openai_realtime_session_audit_task_failed",
log_type = "ops",
request_id,
error = %error,
"OpenAI Realtime session audit task failed"
),
Err(_) => warn!(
target: REALTIME_AUDIT_LOG_TARGET,
event_name = "openai_realtime_session_audit_write_slow",
log_type = "ops",
request_id,
wait_ms = REALTIME_AUDIT_WRITE_WAIT.as_millis() as u64,
write_detached = true,
"OpenAI Realtime stopped waiting for a slow session audit write"
),
}
}
fn build_terminal_event(self, terminal: RealtimeSessionTerminal) -> UsageEvent {
let usage_available = terminal.usage.responses > 0;
let pricing_available = usage_available
&& terminal.usage.input_audio_tokens == 0
&& terminal.usage.output_audio_tokens == 0;
let mut data = build_usage_event_data_seed(&self.plan, self.report_context.as_ref());
data.request_type = Some("realtime".to_string());
data.is_stream = Some(true);
data.status_code = Some(terminal.status_code);
data.response_time_ms = Some(terminal.elapsed_ms);
data.first_byte_time_ms = terminal.first_upstream_frame_ms;
if usage_available {
data.input_tokens = Some(terminal.usage.input_tokens);
data.output_tokens = Some(terminal.usage.output_tokens);
data.total_tokens = Some(terminal.usage.total_tokens);
data.cache_creation_input_tokens = None;
data.cache_creation_ephemeral_5m_input_tokens = None;
data.cache_creation_ephemeral_1h_input_tokens = None;
data.cache_read_input_tokens = Some(terminal.usage.cached_input_tokens);
} else {
clear_usage_and_cost(&mut data);
}
data.cache_creation_cost_usd = None;
data.cache_read_cost_usd = None;
data.total_cost_usd = None;
data.actual_total_cost_usd = None;
if terminal.disposition != RealtimeSessionDisposition::Completed {
data.error_message = Some(terminal.termination.to_string());
data.error_category = Some(terminal.disposition.error_category().to_string());
}
data.request_metadata = attach_terminal_metadata(
data.request_metadata,
usage_available,
pricing_available,
&terminal,
);
UsageEvent::new(
terminal.disposition.event_type(),
self.plan.request_id,
data,
)
}
}
fn clear_usage_and_cost(data: &mut crate::usage::UsageEventData) {
data.input_tokens = None;
data.output_tokens = None;
data.total_tokens = None;
data.cache_creation_input_tokens = None;
data.cache_creation_ephemeral_5m_input_tokens = None;
data.cache_creation_ephemeral_1h_input_tokens = None;
data.cache_read_input_tokens = None;
data.cache_creation_cost_usd = None;
data.cache_read_cost_usd = None;
data.total_cost_usd = None;
data.actual_total_cost_usd = None;
}
fn attach_terminal_metadata(
metadata: Option<Value>,
usage_available: bool,
pricing_available: bool,
terminal: &RealtimeSessionTerminal,
) -> Option<Value> {
let mut object = match metadata {
Some(Value::Object(object)) => object,
_ => Map::new(),
};
object.insert(
USAGE_AVAILABLE_METADATA_KEY.to_string(),
Value::Bool(usage_available),
);
object.insert(
USAGE_PRICING_AVAILABLE_METADATA_KEY.to_string(),
Value::Bool(pricing_available),
);
object.insert(WEBSOCKET_MODE_METADATA_KEY.to_string(), Value::Bool(true));
object.insert(
WEBSOCKET_TRANSPORT_METADATA_KEY.to_string(),
Value::String(REALTIME_WEBSOCKET_TRANSPORT.to_string()),
);
object.insert(
REALTIME_SESSION_METADATA_KEY.to_string(),
json!({
"schema_version": REALTIME_AUDIT_SCHEMA_VERSION,
"transport": "websocket",
"state": terminal.disposition.state(),
"termination": terminal.termination,
"elapsed_ms": terminal.elapsed_ms,
"client_frames": terminal.client_frames,
"client_bytes": terminal.client_bytes,
"upstream_frames": terminal.upstream_frames,
"upstream_bytes": terminal.upstream_bytes,
"first_upstream_frame_ms": terminal.first_upstream_frame_ms,
"usage_state": if usage_available { "authoritative" } else { "unavailable" },
"pricing_state": if !usage_available {
"usage_unavailable"
} else if pricing_available {
"compatible_text_usage"
} else {
"unsupported_audio_breakdown"
},
"usage_scope": "response_done",
"input_transcription_usage_included": false,
"usage_response_count": terminal.usage.responses,
"cached_input_tokens": terminal.usage.cached_input_tokens,
"input_audio_tokens": terminal.usage.input_audio_tokens,
"output_audio_tokens": terminal.usage.output_audio_tokens,
}),
);
Some(Value::Object(object))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum RealtimeSessionDisposition {
Completed,
Failed,
Cancelled,
}
impl RealtimeSessionDisposition {
const fn event_type(self) -> UsageEventType {
match self {
Self::Completed => UsageEventType::Completed,
Self::Failed => UsageEventType::Failed,
Self::Cancelled => UsageEventType::Cancelled,
}
}
const fn state(self) -> &'static str {
match self {
Self::Completed => "closed",
Self::Failed => "failed",
Self::Cancelled => "cancelled",
}
}
const fn error_category(self) -> &'static str {
match self {
Self::Completed => "none",
Self::Failed => "transport_error",
Self::Cancelled => "client_cancelled",
}
}
}
#[derive(Debug, Clone, Copy)]
pub(super) struct RealtimeSessionTerminal {
pub(super) disposition: RealtimeSessionDisposition,
pub(super) status_code: u16,
pub(super) termination: &'static str,
pub(super) elapsed_ms: u64,
pub(super) first_upstream_frame_ms: Option<u64>,
pub(super) client_frames: u64,
pub(super) client_bytes: u64,
pub(super) upstream_frames: u64,
pub(super) upstream_bytes: u64,
pub(super) usage: RealtimeUsageTotals,
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use aether_contracts::{ExecutionTimeouts, RequestBody};
use super::*;
fn sample_plan() -> ExecutionPlan {
ExecutionPlan {
request_id: "realtime-request".to_string(),
candidate_id: Some("candidate-realtime".to_string()),
provider_name: Some("OpenAI".to_string()),
provider_id: "provider-realtime".to_string(),
endpoint_id: "endpoint-realtime".to_string(),
key_id: "key-realtime".to_string(),
method: "GET".to_string(),
url: "wss://example.test/v1/realtime?model=gpt-realtime".to_string(),
headers: BTreeMap::new(),
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: true,
client_api_format: "openai:realtime".to_string(),
provider_api_format: "openai:realtime".to_string(),
model_name: Some("gpt-realtime".to_string()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts::default()),
}
}
fn terminal(usage: RealtimeUsageTotals) -> RealtimeSessionTerminal {
RealtimeSessionTerminal {
disposition: RealtimeSessionDisposition::Completed,
status_code: 200,
termination: "client_close_frame",
elapsed_ms: 1500,
first_upstream_frame_ms: Some(30),
client_frames: 4,
client_bytes: 600,
upstream_frames: 8,
upstream_bytes: 1200,
usage,
}
}
#[test]
fn response_done_usage_becomes_authoritative_session_usage() {
let event = RealtimeSessionAudit::new(
&sample_plan(),
Some(&json!({"user_id": "user-1", "api_key_id": "api-key-1"})),
)
.build_terminal_event(terminal(RealtimeUsageTotals {
responses: 2,
input_tokens: 120,
output_tokens: 40,
total_tokens: 160,
cached_input_tokens: 30,
input_audio_tokens: 20,
output_audio_tokens: 10,
}));
assert_eq!(event.event_type, UsageEventType::Completed);
assert_eq!(event.data.input_tokens, Some(120));
assert_eq!(event.data.output_tokens, Some(40));
assert_eq!(event.data.total_tokens, Some(160));
assert_eq!(event.data.cache_read_input_tokens, Some(30));
let metadata = event.data.request_metadata.expect("metadata");
assert_eq!(metadata[USAGE_AVAILABLE_METADATA_KEY], true);
assert_eq!(metadata[USAGE_PRICING_AVAILABLE_METADATA_KEY], false);
assert_eq!(metadata[WEBSOCKET_MODE_METADATA_KEY], true);
assert_eq!(
metadata[WEBSOCKET_TRANSPORT_METADATA_KEY],
REALTIME_WEBSOCKET_TRANSPORT
);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["usage_state"],
"authoritative"
);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["pricing_state"],
"unsupported_audio_breakdown"
);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["input_audio_tokens"],
20
);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["usage_scope"],
"response_done"
);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["input_transcription_usage_included"],
false
);
}
#[test]
fn missing_response_done_usage_is_visible_but_unmetered() {
let event = RealtimeSessionAudit::new(&sample_plan(), None)
.build_terminal_event(terminal(RealtimeUsageTotals::default()));
assert_eq!(event.data.input_tokens, None);
assert_eq!(event.data.output_tokens, None);
assert_eq!(event.data.total_tokens, None);
assert_eq!(event.data.total_cost_usd, None);
let metadata = event.data.request_metadata.expect("metadata");
assert_eq!(metadata[USAGE_AVAILABLE_METADATA_KEY], false);
assert_eq!(metadata[USAGE_PRICING_AVAILABLE_METADATA_KEY], false);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["usage_state"],
"unavailable"
);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["pricing_state"],
"usage_unavailable"
);
}
#[test]
fn text_only_response_done_usage_remains_priceable() {
let event = RealtimeSessionAudit::new(&sample_plan(), None).build_terminal_event(terminal(
RealtimeUsageTotals {
responses: 1,
input_tokens: 12,
output_tokens: 4,
total_tokens: 16,
cached_input_tokens: 2,
input_audio_tokens: 0,
output_audio_tokens: 0,
},
));
let metadata = event.data.request_metadata.expect("metadata");
assert_eq!(metadata[USAGE_AVAILABLE_METADATA_KEY], true);
assert_eq!(metadata[USAGE_PRICING_AVAILABLE_METADATA_KEY], true);
assert_eq!(
metadata[REALTIME_SESSION_METADATA_KEY]["pricing_state"],
"compatible_text_usage"
);
}
#[test]
fn failed_session_preserves_authoritative_usage_without_becoming_billable() {
let mut failed = terminal(RealtimeUsageTotals {
responses: 1,
input_tokens: 18,
output_tokens: 3,
total_tokens: 21,
cached_input_tokens: 4,
input_audio_tokens: 7,
output_audio_tokens: 0,
});
failed.disposition = RealtimeSessionDisposition::Failed;
failed.status_code = 502;
failed.termination = "upstream_read_failed";
let event = RealtimeSessionAudit::new(&sample_plan(), None).build_terminal_event(failed);
assert_eq!(event.event_type, UsageEventType::Failed);
assert_eq!(event.data.input_tokens, Some(18));
assert_eq!(event.data.total_tokens, Some(21));
assert_eq!(
event.data.error_category.as_deref(),
Some("transport_error")
);
let metadata = event.data.request_metadata.expect("metadata");
assert_eq!(metadata[USAGE_AVAILABLE_METADATA_KEY], true);
assert_eq!(metadata[USAGE_PRICING_AVAILABLE_METADATA_KEY], false);
assert_eq!(metadata[REALTIME_SESSION_METADATA_KEY]["state"], "failed");
}
#[test]
fn failed_session_without_response_usage_is_explicitly_unavailable() {
let mut failed = terminal(RealtimeUsageTotals::default());
failed.disposition = RealtimeSessionDisposition::Failed;
failed.status_code = 502;
failed.termination = "upstream_closed";
let event = RealtimeSessionAudit::new(&sample_plan(), None).build_terminal_event(failed);
assert_eq!(event.event_type, UsageEventType::Failed);
assert_eq!(event.data.input_tokens, None);
assert_eq!(event.data.total_tokens, None);
let metadata = event.data.request_metadata.expect("metadata");
assert_eq!(metadata[USAGE_AVAILABLE_METADATA_KEY], false);
assert_eq!(metadata[USAGE_PRICING_AVAILABLE_METADATA_KEY], false);
}
}
@@ -0,0 +1,65 @@
//! Public OpenAI Realtime (`/v1/realtime`) WebSocket bridge.
//!
//! This is intentionally separate from Responses WebSocket mode and Codex
//! Frameless `/v1/live`: all three use WebSocket transport but have different
//! event grammars and lifecycle semantics.
mod audit;
mod planner;
mod protocol;
mod session;
use std::net::SocketAddr;
use axum::body::Body;
use axum::extract::ws::WebSocketUpgrade;
use axum::extract::{ConnectInfo, State};
use axum::http::{HeaderMap, Response, Uri};
use crate::handlers::proxy::websocket::ingress::{
prepare_authenticated_ai_websocket, AuthenticatedAiWebSocketUpgradePreparation,
WebSocketIngressSpec,
};
use crate::handlers::proxy::websocket::session::REALTIME_WEBSOCKET_SESSION_LIMITS;
use crate::{AppState, GatewayError};
pub(crate) async fn realtime_websocket(
State(state): State<AppState>,
ConnectInfo(remote_addr): ConnectInfo<SocketAddr>,
ws: WebSocketUpgrade,
headers: HeaderMap,
uri: Uri,
) -> Result<Response<Body>, GatewayError> {
match prepare_authenticated_ai_websocket(
state,
remote_addr,
headers,
uri,
REALTIME_WEBSOCKET_INGRESS_SPEC,
)
.await?
{
AuthenticatedAiWebSocketUpgradePreparation::Rejected(response) => Ok(response),
AuthenticatedAiWebSocketUpgradePreparation::Ready(prepared) => {
let realtime =
match session::prepare_realtime_websocket(prepared.state(), prepared.context())
.await
{
Ok(realtime) => realtime,
Err(rejection) => {
return prepared.rejection_response(rejection.status(), rejection.message())
}
};
Ok(prepared.into_response_with(
ws,
REALTIME_WEBSOCKET_SESSION_LIMITS,
realtime,
session::run_realtime_websocket,
))
}
}
}
const REALTIME_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec {
route_unavailable_message: "OpenAI Realtime WebSocket route is unavailable",
};
@@ -0,0 +1,249 @@
//! Candidate planning for the public OpenAI Realtime WebSocket transport.
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use axum::http::{HeaderValue, Method};
use serde_json::json;
use crate::ai_serving::{
build_standard_stream_plan_from_decision, maybe_build_stream_decision_payload,
AiExecutionDecision,
};
use crate::control::GatewayControlDecision;
use crate::headers::request_origin_from_headers_and_remote_addr;
use crate::privacy::RedactionSessionSlot;
use crate::{AppState, GatewayError};
pub(super) struct PlannedRealtimeCandidate {
pub(super) execution: AiExecutionDecision,
pub(super) admission_plan: aether_contracts::ExecutionPlan,
pub(super) provider_id: String,
pub(super) endpoint_id: String,
pub(super) key_id: String,
pub(super) provider_model: String,
pub(super) pool_lease: RealtimePoolLeaseGuard,
}
pub(super) struct RealtimePoolLeaseGuard {
state: AppState,
report_context: Option<serde_json::Value>,
renewal_task: Option<tokio::task::JoinHandle<()>>,
healthy: Arc<AtomicBool>,
armed: bool,
}
impl RealtimePoolLeaseGuard {
fn new(state: &AppState, decision: &AiExecutionDecision) -> Self {
let report_context = decision.report_context.clone();
let lease = crate::orchestration::local_execution_candidate_metadata_from_report_context(
report_context.as_ref(),
)
.pool_key_lease;
let healthy = Arc::new(AtomicBool::new(true));
let renewal_task = lease.map(|lease| {
let runtime_state = Arc::clone(&state.runtime_state);
let healthy = Arc::clone(&healthy);
tokio::spawn(async move {
let ttl = Duration::from_millis(lease.ttl_ms);
let interval = Duration::from_millis((lease.ttl_ms / 3).max(1));
loop {
tokio::time::sleep(interval).await;
match runtime_state.lock_renew(&lease, ttl).await {
Ok(true) => {}
Ok(false) | Err(_) => {
healthy.store(false, Ordering::Release);
return;
}
}
}
})
});
Self {
state: state.clone(),
report_context,
renewal_task,
healthy,
armed: true,
}
}
pub(super) fn is_healthy(&self) -> bool {
self.healthy.load(Ordering::Acquire)
}
pub(super) async fn release(mut self) {
if let Some(task) = self.renewal_task.take() {
task.abort();
}
crate::orchestration::release_pool_key_lease_from_report_context(
&self.state,
self.report_context.as_ref(),
)
.await;
self.armed = false;
}
}
impl Drop for RealtimePoolLeaseGuard {
fn drop(&mut self) {
if let Some(task) = self.renewal_task.take() {
task.abort();
}
if !self.armed {
return;
}
let state = self.state.clone();
let report_context = self.report_context.take();
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
runtime.spawn(async move {
crate::orchestration::release_pool_key_lease_from_report_context(
&state,
report_context.as_ref(),
)
.await;
});
}
}
}
pub(super) async fn plan_realtime_candidate(
state: &AppState,
context: &crate::handlers::proxy::websocket::ingress::WebSocketRequestContext,
client_model: &str,
) -> Result<Option<PlannedRealtimeCandidate>, GatewayError> {
let parts = realtime_planning_parts(context);
let body = json!({"model": client_model});
let Some(execution) = maybe_build_stream_decision_payload(
state,
&parts,
context.trace_id.as_str(),
&context.decision,
&body,
None,
)
.await?
else {
return Ok(None);
};
if execution
.provider_api_format
.as_deref()
.map(crate::ai_serving::normalize_api_format_alias)
.as_deref()
!= Some("openai:realtime")
{
crate::orchestration::release_pool_key_lease_from_report_context(
state,
execution.report_context.as_ref(),
)
.await;
return Ok(None);
}
let pool_lease = RealtimePoolLeaseGuard::new(state, &execution);
let Some(attempt) =
build_standard_stream_plan_from_decision(&parts, &body, execution.clone(), false)?
else {
pool_lease.release().await;
return Ok(None);
};
let provider_id = execution.provider_id.clone().unwrap_or_default();
let endpoint_id = execution.endpoint_id.clone().unwrap_or_default();
let key_id = execution.key_id.clone().unwrap_or_default();
let provider_model = execution
.mapped_model
.clone()
.or_else(|| execution.model_name.clone())
.unwrap_or_default();
if provider_id.is_empty()
|| endpoint_id.is_empty()
|| key_id.is_empty()
|| provider_model.trim().is_empty()
|| execution.upstream_url.as_deref().is_none_or(str::is_empty)
{
pool_lease.release().await;
return Ok(None);
}
Ok(Some(PlannedRealtimeCandidate {
execution,
admission_plan: attempt.plan,
provider_id,
endpoint_id,
key_id,
provider_model,
pool_lease,
}))
}
fn realtime_planning_parts(
context: &crate::handlers::proxy::websocket::ingress::WebSocketRequestContext,
) -> http::request::Parts {
let mut request = http::Request::builder()
.method(Method::GET)
.uri(context.uri.clone())
.body(())
.expect("the authenticated Realtime URI must remain valid");
*request.headers_mut() = context.headers.clone();
request.headers_mut().insert(
http::header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
request
.extensions_mut()
.insert(request_origin_from_headers_and_remote_addr(
&context.headers,
&context.remote_addr,
));
request
.extensions_mut()
.insert(RedactionSessionSlot::default());
request.into_parts().0
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn realtime_planning_requires_the_explicit_realtime_route() {
let request = http::Request::builder()
.method(Method::GET)
.uri("/v1/realtime?model=gpt-realtime")
.body(())
.unwrap();
let (parts, _) = request.into_parts();
let decision = GatewayControlDecision {
public_path: "/v1/realtime".to_string(),
public_query_string: Some("model=gpt-realtime".to_string()),
route_class: Some("ai_public".to_string()),
route_family: Some("openai".to_string()),
route_kind: Some("realtime".to_string()),
client_surface: None,
api_operation: None,
gateway_credential_carrier: None,
request_auth_channel: None,
auth_context: None,
admin_principal: None,
auth_endpoint_signature: None,
execution_runtime_candidate: true,
local_auth_rejection: None,
model_directive_policy: Default::default(),
};
assert_eq!(
crate::ai_serving::resolve_execution_runtime_stream_plan_kind_with_client_surface(
decision.route_class.as_deref(),
decision.route_family.as_deref(),
decision.route_kind.as_deref(),
decision.client_surface,
decision.request_auth_channel.as_deref(),
&parts.method,
parts.uri.path(),
),
Some(crate::ai_serving::OPENAI_REALTIME_STREAM_PLAN_KIND)
);
}
}
@@ -0,0 +1,236 @@
//! Bounded validation and observation for the public OpenAI Realtime protocol.
//!
//! Realtime events are otherwise relayed as opaque text/binary frames. Keeping
//! this module deliberately small prevents Aether from becoming a schema
//! allowlist for future client and server events.
use std::collections::BTreeSet;
use serde_json::{json, Value};
const MAX_MODEL_BYTES: usize = 256;
const MAX_OBSERVED_RESPONSE_IDS: usize = 1_024;
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
pub(super) enum RealtimeProtocolError {
#[error("invalid Realtime model query")]
InvalidModelQuery,
#[error("invalid Realtime model")]
InvalidModel,
}
impl RealtimeProtocolError {
pub(super) const fn client_message(self) -> &'static str {
match self {
Self::InvalidModelQuery => {
"Realtime WebSocket requires exactly one model query parameter"
}
Self::InvalidModel => {
"Realtime model must be a non-empty identifier no longer than 256 bytes"
}
}
}
}
pub(super) fn model_from_query(query: Option<&str>) -> Result<String, RealtimeProtocolError> {
let mut model = None;
for (name, value) in url::form_urlencoded::parse(query.unwrap_or_default().as_bytes()) {
if name.eq_ignore_ascii_case("model") {
if model.is_some() {
return Err(RealtimeProtocolError::InvalidModelQuery);
}
validate_model(value.as_ref())?;
model = Some(value.into_owned());
} else if query_parameter_is_sensitive(name.as_ref()) {
return Err(RealtimeProtocolError::InvalidModelQuery);
}
}
model.ok_or(RealtimeProtocolError::InvalidModelQuery)
}
fn validate_model(model: &str) -> Result<(), RealtimeProtocolError> {
if model.is_empty()
|| model.len() > MAX_MODEL_BYTES
|| model.trim() != model
|| model.chars().any(char::is_control)
{
return Err(RealtimeProtocolError::InvalidModel);
}
Ok(())
}
fn query_parameter_is_sensitive(name: &str) -> bool {
matches!(
name.to_ascii_lowercase().as_str(),
"key"
| "api_key"
| "api-key"
| "x-api-key"
| "access_token"
| "authorization"
| "token"
| "client_secret"
| "secret_key"
| "signature"
| "sig"
)
}
pub(super) fn error_event(code: &str, message: &str) -> Value {
json!({
"type": "error",
"error": {
"type": "server_error",
"code": code,
"message": message,
}
})
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub(super) struct RealtimeUsageTotals {
pub(super) responses: u64,
pub(super) input_tokens: u64,
pub(super) output_tokens: u64,
pub(super) total_tokens: u64,
pub(super) cached_input_tokens: u64,
pub(super) input_audio_tokens: u64,
pub(super) output_audio_tokens: u64,
}
#[derive(Debug, Default)]
pub(super) struct RealtimeUsageObserver {
totals: RealtimeUsageTotals,
response_ids: BTreeSet<String>,
}
impl RealtimeUsageObserver {
pub(super) fn observe(&mut self, raw: &str) {
let Ok(event) = serde_json::from_str::<Value>(raw) else {
return;
};
if event.get("type").and_then(Value::as_str) != Some("response.done") {
return;
}
let Some(response) = event.get("response").and_then(Value::as_object) else {
return;
};
let Some(usage) = response.get("usage").and_then(Value::as_object) else {
return;
};
if let Some(response_id) = response
.get("id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
if self.response_ids.contains(response_id) {
return;
}
if self.response_ids.len() >= MAX_OBSERVED_RESPONSE_IDS {
return;
}
self.response_ids.insert(response_id.to_string());
}
let input_tokens = json_u64(usage.get("input_tokens"));
let output_tokens = json_u64(usage.get("output_tokens"));
let total_tokens = usage
.get("total_tokens")
.and_then(Value::as_u64)
.unwrap_or_else(|| input_tokens.saturating_add(output_tokens));
self.totals.responses = self.totals.responses.saturating_add(1);
self.totals.input_tokens = self.totals.input_tokens.saturating_add(input_tokens);
self.totals.output_tokens = self.totals.output_tokens.saturating_add(output_tokens);
self.totals.total_tokens = self.totals.total_tokens.saturating_add(total_tokens);
if let Some(details) = usage.get("input_token_details").and_then(Value::as_object) {
self.totals.cached_input_tokens = self
.totals
.cached_input_tokens
.saturating_add(json_u64(details.get("cached_tokens")));
self.totals.input_audio_tokens = self
.totals
.input_audio_tokens
.saturating_add(json_u64(details.get("audio_tokens")));
}
if let Some(details) = usage.get("output_token_details").and_then(Value::as_object) {
self.totals.output_audio_tokens = self
.totals
.output_audio_tokens
.saturating_add(json_u64(details.get("audio_tokens")));
}
}
pub(super) const fn totals(&self) -> RealtimeUsageTotals {
self.totals
}
}
fn json_u64(value: Option<&Value>) -> u64 {
value.and_then(Value::as_u64).unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::{model_from_query, RealtimeUsageObserver};
#[test]
fn model_query_requires_one_bounded_model_and_ignores_safe_hints() {
assert_eq!(
model_from_query(Some("trace=1&model=gpt-realtime-client")),
Ok("gpt-realtime-client".to_string())
);
assert!(model_from_query(None).is_err());
assert!(model_from_query(Some("model=a&MODEL=b")).is_err());
assert!(model_from_query(Some("model=a&key=secret")).is_err());
assert!(model_from_query(Some(format!("model={}", "x".repeat(257)).as_str())).is_err());
}
#[test]
fn response_done_usage_is_observed_once_without_reconstructing_events() {
let event = serde_json::json!({
"type": "response.done",
"future_server_field": {"opaque": true},
"response": {
"id": "resp_1",
"usage": {
"input_tokens": 12,
"output_tokens": 7,
"total_tokens": 19,
"input_token_details": {"cached_tokens": 4, "audio_tokens": 3},
"output_token_details": {"audio_tokens": 2}
}
}
})
.to_string();
let mut observer = RealtimeUsageObserver::default();
observer.observe(event.as_str());
observer.observe(event.as_str());
let totals = observer.totals();
assert_eq!(totals.responses, 1);
assert_eq!(totals.input_tokens, 12);
assert_eq!(totals.output_tokens, 7);
assert_eq!(totals.total_tokens, 19);
assert_eq!(totals.cached_input_tokens, 4);
assert_eq!(totals.input_audio_tokens, 3);
assert_eq!(totals.output_audio_tokens, 2);
}
#[test]
fn missing_total_tokens_uses_authoritative_component_sum() {
let mut observer = RealtimeUsageObserver::default();
observer.observe(
serde_json::json!({
"type": "response.done",
"response": {
"id": "resp_without_total",
"usage": {"input_tokens": 9, "output_tokens": 4}
}
})
.to_string()
.as_str(),
);
assert_eq!(observer.totals().total_tokens, 13);
}
}
@@ -0,0 +1,652 @@
//! Opaque bidirectional relay for the public OpenAI Realtime WebSocket API.
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
use axum::http::StatusCode;
use futures_util::{SinkExt, StreamExt};
use tracing::{info, warn};
use wreq::ws::message::Message as WreqWsMessage;
use crate::control::execution_plan_balance_capacity_rejection;
use crate::handlers::proxy::websocket::ingress::{
WebSocketConnectionLog, WebSocketConnectionLogSpec, WebSocketRequestContext,
};
use crate::handlers::proxy::websocket::responses::ResponsesWebSocketTurnAdmission;
use crate::handlers::proxy::websocket::session::{
CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN,
REALTIME_WEBSOCKET_SESSION_LIMITS, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::{
client_message_to_upstream, close_client_socket, close_upstream_socket,
connect_upstream_websocket, send_client_message, upstream_message_to_client,
websocket_relay_frame_queue, UpstreamWebSocketErrorCodes, WebSocketRelayPumpControl,
WebSocketRelayQueueError, WebSocketWriteError,
};
use crate::{AppState, GatewayError};
use super::audit::{RealtimeSessionAudit, RealtimeSessionDisposition, RealtimeSessionTerminal};
use super::planner::{plan_realtime_candidate, PlannedRealtimeCandidate};
use super::protocol::{error_event, model_from_query, RealtimeUsageObserver};
const REALTIME_LOG_TARGET: &str = "aether_gateway::handlers::proxy::realtime_ws";
const REALTIME_CONNECTION_LOG_SPEC: WebSocketConnectionLogSpec = WebSocketConnectionLogSpec {
opened_event_name: "openai_realtime_websocket_connection_opened",
closed_event_name: "openai_realtime_websocket_connection_closed",
opened_message: "gateway accepted OpenAI Realtime WebSocket connection",
closed_message: "gateway closed OpenAI Realtime WebSocket connection",
execution_path: "openai_realtime_websocket_bridge",
provider_type: "openai_realtime",
};
const REALTIME_UPSTREAM_ERRORS: UpstreamWebSocketErrorCodes = UpstreamWebSocketErrorCodes {
upstream_url_missing: "openai_realtime_upstream_url_missing",
upstream_url_invalid: "openai_realtime_upstream_url_invalid",
frontdoor_self_loop: "openai_realtime_websocket_frontdoor_self_loop",
headers_invalid: "openai_realtime_websocket_headers_invalid",
client_build_failed: "openai_realtime_websocket_client_build_failed",
proxy_invalid: "openai_realtime_websocket_proxy_invalid",
tunnel_proxy_unsupported: "openai_realtime_websocket_tunnel_proxy_unsupported",
handshake_failed: "openai_realtime_websocket_handshake_failed",
upgrade_rejected: "openai_realtime_websocket_upgrade_rejected",
upgrade_failed: "openai_realtime_websocket_upgrade_failed",
};
pub(super) struct PreparedRealtimeWebSocket {
upstream: wreq::ws::WebSocket,
admission: ResponsesWebSocketTurnAdmission,
candidate: PlannedRealtimeCandidate,
}
pub(super) struct RealtimeWebSocketPreflightRejection {
status: StatusCode,
message: String,
}
impl RealtimeWebSocketPreflightRejection {
pub(super) const fn status(&self) -> StatusCode {
self.status
}
pub(super) fn message(&self) -> &str {
self.message.as_str()
}
}
pub(super) async fn prepare_realtime_websocket(
state: &AppState,
context: &WebSocketRequestContext,
) -> Result<PreparedRealtimeWebSocket, RealtimeWebSocketPreflightRejection> {
if !realtime_usage_accounting_is_safe(context) {
return Err(rejection(
StatusCode::NOT_IMPLEMENTED,
"Realtime WebSocket is unavailable for finite-balance keys until session usage settlement is enabled",
));
}
let client_model = model_from_query(context.uri.query())
.map_err(|error| rejection(StatusCode::BAD_REQUEST, error.client_message()))?;
let candidate = plan_realtime_candidate(state, context, client_model.as_str())
.await
.map_err(|error| {
warn!(
target: REALTIME_LOG_TARGET,
event_name = "openai_realtime_planning_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_kind = gateway_error_kind(&error),
"OpenAI Realtime candidate planning failed"
);
rejection(
StatusCode::INTERNAL_SERVER_ERROR,
"Realtime provider planning failed",
)
})?
.ok_or_else(|| {
rejection(
StatusCode::SERVICE_UNAVAILABLE,
"No eligible OpenAI Realtime provider mapping is available",
)
})?;
if execution_plan_balance_capacity_rejection(
state,
&context.decision,
&candidate.admission_plan,
candidate.execution.report_context.as_ref(),
)
.await
.map_err(|_| {
rejection(
StatusCode::INTERNAL_SERVER_ERROR,
"Realtime balance admission failed",
)
})?
.is_some()
{
candidate.pool_lease.release().await;
return Err(rejection(
StatusCode::TOO_MANY_REQUESTS,
"Realtime request capacity is unavailable",
));
}
let admission = match ResponsesWebSocketTurnAdmission::acquire(
state,
&candidate.admission_plan,
context.trace_id.as_str(),
)
.await
{
Ok(admission) => admission,
Err(error) => {
candidate.pool_lease.release().await;
return Err(rejection(
admission_error_status(&error),
"Realtime connection admission failed",
));
}
};
if !candidate.pool_lease.is_healthy() {
admission.release().await;
candidate.pool_lease.release().await;
return Err(rejection(
StatusCode::SERVICE_UNAVAILABLE,
"Realtime provider ownership was lost",
));
}
let mut upstream = match connect_upstream_websocket(
&candidate.execution,
REALTIME_WEBSOCKET_SESSION_LIMITS,
REALTIME_UPSTREAM_ERRORS,
)
.await
{
Ok(connection) => connection.socket,
Err(error_code) => {
warn!(
target: REALTIME_LOG_TARGET,
event_name = "openai_realtime_upstream_connect_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_id = %candidate.provider_id,
endpoint_id = %candidate.endpoint_id,
key_id = %candidate.key_id,
error_code,
"OpenAI Realtime upstream connection failed"
);
admission.release().await;
candidate.pool_lease.release().await;
return Err(rejection(
StatusCode::BAD_GATEWAY,
"Realtime upstream WebSocket connection failed",
));
}
};
// The scheduler lease can expire while the upstream WebSocket handshake
// is in flight. Re-check it after the handshake so an invalid provider
// candidate is rejected before the downstream HTTP 101 is committed.
if !candidate.pool_lease.is_healthy() {
close_upstream_socket(&mut upstream, None).await;
admission.release().await;
candidate.pool_lease.release().await;
return Err(rejection(
StatusCode::SERVICE_UNAVAILABLE,
"Realtime provider ownership was lost",
));
}
Ok(PreparedRealtimeWebSocket {
upstream,
admission,
candidate,
})
}
pub(super) async fn run_realtime_websocket(
mut client_socket: WebSocket,
state: AppState,
context: WebSocketRequestContext,
prepared: PreparedRealtimeWebSocket,
) {
let connection_log = WebSocketConnectionLog::new(&context, REALTIME_CONNECTION_LOG_SPEC);
connection_log.log_opened();
let PreparedRealtimeWebSocket {
mut upstream,
admission,
candidate,
} = prepared;
let audit = RealtimeSessionAudit::new(
&candidate.admission_plan,
candidate.execution.report_context.as_ref(),
);
let terminal = relay_realtime(&mut client_socket, &mut upstream, &context, &candidate).await;
close_upstream_socket(&mut upstream, None).await;
admission.release().await;
candidate.pool_lease.release().await;
if matches!(
terminal.termination,
"connection_duration_limit" | "connection_admission_lost" | "pool_key_lease_lost"
) {
close_client_socket(&mut client_socket, CLOSE_TRY_AGAIN, terminal.termination).await;
}
audit.finish(&state, terminal).await;
}
async fn relay_realtime(
client_socket: &mut WebSocket,
upstream: &mut wreq::ws::WebSocket,
context: &WebSocketRequestContext,
candidate: &PlannedRealtimeCandidate,
) -> RealtimeSessionTerminal {
let started_at = Instant::now();
let connection_deadline =
tokio::time::sleep(REALTIME_WEBSOCKET_SESSION_LIMITS.max_connection_duration);
tokio::pin!(connection_deadline);
let mut lease_health = tokio::time::interval(Duration::from_secs(1));
lease_health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let stats = Arc::new(Mutex::new(RelayStats::default()));
let usage = Arc::new(Mutex::new(RealtimeUsageObserver::default()));
let relay_control = WebSocketRelayPumpControl::new();
let termination = {
let (mut client_write, mut client_read) = (&mut *client_socket).split();
let (mut upstream_write, mut upstream_read) = (&mut *upstream).split();
let client_to_upstream = {
let control = relay_control.clone();
let stats = Arc::clone(&stats);
async move {
let (queue_tx, mut queue_rx) = websocket_relay_frame_queue();
let reader_control = control.clone();
let reader = async move {
loop {
let client = tokio::select! {
biased;
_ = reader_control.cancelled() => return "relay_cancelled",
client = client_read.next() => client,
};
let Some(client) = client else {
return "client_closed";
};
let Ok(client) = client else {
return "client_read_failed";
};
let (bytes, is_close) = client_frame_metadata(&client);
{
let mut stats = stats
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
stats.client_frames = stats.client_frames.saturating_add(1);
stats.client_bytes = stats.client_bytes.saturating_add(bytes as u64);
}
match reader_control
.enqueue(&queue_tx, client_message_to_upstream(client))
.await
{
Ok(()) => {}
Err(WebSocketRelayQueueError::Cancelled) => {
return "relay_cancelled";
}
Err(WebSocketRelayQueueError::Closed) => {
return "upstream_write_failed";
}
}
if is_close {
return "client_close_frame";
}
}
};
let writer_control = control;
let writer = async move {
loop {
let message = tokio::select! {
biased;
_ = writer_control.cancelled() => return None,
message = queue_rx.recv() => message,
};
let Some(message) = message else {
return None;
};
let result = writer_control
.send(async { upstream_write.send(message).await.map_err(|_| ()) })
.await;
match result {
Ok(()) => {}
Err(WebSocketWriteError::Cancelled) => return None,
Err(_) => return Some("upstream_write_failed"),
}
}
};
tokio::pin!(reader, writer);
tokio::select! {
reader_exit = &mut reader => {
writer.await.unwrap_or(reader_exit)
}
writer_exit = &mut writer => {
match writer_exit {
Some(writer_exit) => writer_exit,
None => reader.await,
}
}
}
}
};
let upstream_to_client = {
let control = relay_control.clone();
let stats = Arc::clone(&stats);
let usage = Arc::clone(&usage);
async move {
let (queue_tx, mut queue_rx) = websocket_relay_frame_queue();
let reader_control = control.clone();
let reader = async move {
loop {
let provider = tokio::select! {
biased;
_ = reader_control.cancelled() => return "relay_cancelled",
provider = upstream_read.next() => provider,
};
let Some(provider) = provider else {
return "upstream_closed";
};
let Ok(provider) = provider else {
return "upstream_read_failed";
};
let (bytes, is_close) = upstream_frame_metadata(&provider);
{
let mut stats = stats
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
stats.first_upstream_frame_ms.get_or_insert_with(|| {
started_at.elapsed().as_millis().min(u128::from(u64::MAX)) as u64
});
stats.upstream_frames = stats.upstream_frames.saturating_add(1);
stats.upstream_bytes =
stats.upstream_bytes.saturating_add(bytes as u64);
}
if let WreqWsMessage::Text(text) = &provider {
usage
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.observe(text.as_str());
}
match reader_control
.enqueue(&queue_tx, upstream_message_to_client(provider))
.await
{
Ok(()) => {}
Err(WebSocketRelayQueueError::Cancelled) => {
return "relay_cancelled";
}
Err(WebSocketRelayQueueError::Closed) => {
return "client_write_failed";
}
}
if is_close {
return "upstream_close_frame";
}
}
};
let writer_control = control;
let writer = async move {
loop {
let message = tokio::select! {
biased;
_ = writer_control.cancelled() => return None,
message = queue_rx.recv() => message,
};
let Some(message) = message else {
return None;
};
let result = writer_control
.send(async { client_write.send(message).await.map_err(|_| ()) })
.await;
match result {
Ok(()) => {}
Err(WebSocketWriteError::Cancelled) => return None,
Err(_) => return Some("client_write_failed"),
}
}
};
tokio::pin!(reader, writer);
tokio::select! {
reader_exit = &mut reader => {
writer.await.unwrap_or(reader_exit)
}
writer_exit = &mut writer => {
match writer_exit {
Some(writer_exit) => writer_exit,
None => reader.await,
}
}
}
}
};
tokio::pin!(client_to_upstream, upstream_to_client);
let termination = loop {
tokio::select! {
termination = &mut client_to_upstream => break termination,
termination = &mut upstream_to_client => break termination,
_ = &mut connection_deadline => break "connection_duration_limit",
_ = wait_for_connection_permit_loss(context.websocket_connection_permit.as_ref()) => {
break "connection_admission_lost";
}
_ = lease_health.tick() => {
if !candidate.pool_lease.is_healthy() {
break "pool_key_lease_lost";
}
}
}
};
relay_control.cancel();
termination
};
if termination == "pool_key_lease_lost" {
send_realtime_error(
client_socket,
"openai_realtime_pool_key_lease_lost",
"Realtime provider ownership was lost",
)
.await;
}
let stats = *stats
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let totals = usage
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.totals();
let elapsed_ms = started_at.elapsed().as_millis().min(u128::from(u64::MAX)) as u64;
info!(
target: REALTIME_LOG_TARGET,
event_name = "openai_realtime_relay_finished",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
provider_id = %candidate.provider_id,
endpoint_id = %candidate.endpoint_id,
key_id = %candidate.key_id,
model = %candidate.provider_model,
termination,
client_frames = stats.client_frames,
client_bytes = stats.client_bytes,
upstream_frames = stats.upstream_frames,
upstream_bytes = stats.upstream_bytes,
response_count = totals.responses,
input_tokens = totals.input_tokens,
output_tokens = totals.output_tokens,
total_tokens = totals.total_tokens,
cached_input_tokens = totals.cached_input_tokens,
input_audio_tokens = totals.input_audio_tokens,
output_audio_tokens = totals.output_audio_tokens,
elapsed_ms,
"OpenAI Realtime opaque relay finished"
);
realtime_terminal_from_relay(termination, elapsed_ms, stats, totals)
}
fn realtime_terminal_from_relay(
termination: &'static str,
elapsed_ms: u64,
stats: RelayStats,
usage: super::protocol::RealtimeUsageTotals,
) -> RealtimeSessionTerminal {
let (disposition, status_code) = match termination {
"client_close_frame" | "upstream_close_frame" => {
(RealtimeSessionDisposition::Completed, 200)
}
"client_closed"
| "client_read_failed"
| "client_write_failed"
| "connection_duration_limit" => (RealtimeSessionDisposition::Cancelled, 499),
"pool_key_lease_lost" | "connection_admission_lost" => {
(RealtimeSessionDisposition::Failed, 503)
}
"upstream_closed" | "upstream_read_failed" | "upstream_write_failed" => {
(RealtimeSessionDisposition::Failed, 502)
}
_ => (RealtimeSessionDisposition::Failed, 500),
};
RealtimeSessionTerminal {
disposition,
status_code,
termination,
elapsed_ms,
first_upstream_frame_ms: stats.first_upstream_frame_ms,
client_frames: stats.client_frames,
client_bytes: stats.client_bytes,
upstream_frames: stats.upstream_frames,
upstream_bytes: stats.upstream_bytes,
usage,
}
}
fn realtime_usage_accounting_is_safe(context: &WebSocketRequestContext) -> bool {
context
.decision
.auth_context
.as_ref()
.is_some_and(|auth| auth.balance_remaining.is_none())
}
fn rejection(
status: StatusCode,
message: impl Into<String>,
) -> RealtimeWebSocketPreflightRejection {
RealtimeWebSocketPreflightRejection {
status,
message: message.into(),
}
}
fn admission_error_status(error: &GatewayError) -> StatusCode {
match error {
GatewayError::AdmissionTimeout { .. } => StatusCode::TOO_MANY_REQUESTS,
GatewayError::Client { status, .. } => *status,
GatewayError::LocalExecutionPlanningTimeout { .. } => StatusCode::GATEWAY_TIMEOUT,
_ => StatusCode::INTERNAL_SERVER_ERROR,
}
}
fn gateway_error_kind(error: &GatewayError) -> &'static str {
match error {
GatewayError::UpstreamUnavailable { .. } => "upstream_unavailable",
GatewayError::ControlUnavailable { .. } => "control_unavailable",
GatewayError::LocalExecutionPlanningTimeout { .. } => "planning_timeout",
GatewayError::AdmissionTimeout { .. } => "admission_timeout",
GatewayError::Client { .. } => "client_error",
GatewayError::Internal(_) => "internal_error",
}
}
async fn send_realtime_error(client_socket: &mut WebSocket, code: &str, message: &str) {
let event = error_event(code, message).to_string();
let _ = send_client_message(client_socket, AxumWsMessage::Text(event.into())).await;
}
#[derive(Clone, Copy, Default)]
struct RelayStats {
client_frames: u64,
client_bytes: u64,
upstream_frames: u64,
upstream_bytes: u64,
first_upstream_frame_ms: Option<u64>,
}
fn client_frame_metadata(message: &AxumWsMessage) -> (usize, bool) {
match message {
AxumWsMessage::Text(text) => (text.len(), false),
AxumWsMessage::Binary(data) | AxumWsMessage::Ping(data) | AxumWsMessage::Pong(data) => {
(data.len(), false)
}
AxumWsMessage::Close(frame) => (
frame
.as_ref()
.map_or(0, |frame| 2usize.saturating_add(frame.reason.len())),
true,
),
}
}
fn upstream_frame_metadata(message: &WreqWsMessage) -> (usize, bool) {
match message {
WreqWsMessage::Text(text) => (text.len(), false),
WreqWsMessage::Binary(data) | WreqWsMessage::Ping(data) | WreqWsMessage::Pong(data) => {
(data.len(), false)
}
WreqWsMessage::Close(frame) => (
frame
.as_ref()
.map_or(0, |frame| 2usize.saturating_add(frame.reason.len())),
true,
),
}
}
async fn wait_for_connection_permit_loss(permit: Option<&aether_runtime::AdmissionPermit>) {
let Some(permit) = permit else {
std::future::pending::<()>().await;
return;
};
let mut health = tokio::time::interval(Duration::from_secs(1));
health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
health.tick().await;
if !permit.is_healthy() {
return;
}
}
}
#[cfg(test)]
mod tests {
use super::{client_frame_metadata, upstream_frame_metadata};
use axum::extract::ws::Message as AxumWsMessage;
use wreq::ws::message::Message as WreqWsMessage;
#[test]
fn opaque_frame_accounting_does_not_coalesce_audio_or_json_messages() {
assert_eq!(
client_frame_metadata(&AxumWsMessage::Text("{\"type\":\"session.update\"}".into())),
(25, false)
);
assert_eq!(
client_frame_metadata(&AxumWsMessage::Binary(vec![1, 2, 3].into())),
(3, false)
);
assert_eq!(
upstream_frame_metadata(&WreqWsMessage::Text("delta".into())),
(5, false)
);
}
}
@@ -27,6 +27,14 @@ pub(crate) const LIVE_WEBSOCKET_SESSION_LIMITS: WebSocketSessionLimits = WebSock
max_connection_duration: Duration::from_secs(60 * 60),
};
pub(crate) const REALTIME_WEBSOCKET_SESSION_LIMITS: WebSocketSessionLimits =
WebSocketSessionLimits {
max_frame_size: 16 << 20,
max_message_size: 16 << 20,
initial_message_timeout: Duration::from_secs(60),
max_connection_duration: Duration::from_secs(60 * 60),
};
/// A peer that stops draining its receive window must not be able to pin the
/// relay loop. Session loops await socket writes inside a `tokio::select!`,
/// so an unbounded write also suspends the connection and per-turn deadlines
@@ -15,6 +15,8 @@ use axum::http::header::{
use axum::http::{HeaderMap, HeaderName};
use futures_util::{SinkExt, TryFutureExt};
use serde_json::json;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use url::Url;
use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage};
@@ -255,6 +257,7 @@ pub(crate) fn websocket_timeouts(
pub(crate) enum WebSocketWriteError {
Failed,
TimedOut,
Cancelled,
}
impl WebSocketWriteError {
@@ -262,10 +265,76 @@ impl WebSocketWriteError {
match self {
Self::Failed => "write_failed",
Self::TimedOut => "write_timeout",
Self::Cancelled => "write_cancelled",
}
}
}
/// A small per-direction buffer keeps a slow reader from blocking the opposite
/// WebSocket direction while still applying bounded backpressure. At the Live
/// audio cadence this is deliberately only a short burst buffer, not a place
/// where a session can accumulate unbounded media.
pub(crate) const RELAY_FRAME_QUEUE_CAPACITY: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum WebSocketRelayQueueError {
Closed,
Cancelled,
}
/// Shared cancellation for both read/write halves of a bidirectional relay.
///
/// Queue admission and socket writes both observe this token, so a connection
/// deadline or lease loss can interrupt a full queue and an in-flight slow
/// write immediately instead of waiting for [`RELAY_WRITE_TIMEOUT`].
#[derive(Clone, Default)]
pub(crate) struct WebSocketRelayPumpControl {
cancellation: CancellationToken,
}
impl WebSocketRelayPumpControl {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn cancel(&self) {
self.cancellation.cancel();
}
pub(crate) async fn cancelled(&self) {
self.cancellation.cancelled().await;
}
pub(crate) async fn enqueue<T>(
&self,
sender: &mpsc::Sender<T>,
message: T,
) -> Result<(), WebSocketRelayQueueError> {
tokio::select! {
biased;
_ = self.cancellation.cancelled() => Err(WebSocketRelayQueueError::Cancelled),
result = sender.send(message) => {
result.map_err(|_| WebSocketRelayQueueError::Closed)
}
}
}
pub(crate) async fn send<F>(&self, write: F) -> Result<(), WebSocketWriteError>
where
F: std::future::Future<Output = Result<(), ()>>,
{
tokio::select! {
biased;
_ = self.cancellation.cancelled() => Err(WebSocketWriteError::Cancelled),
result = bounded_send(RELAY_WRITE_TIMEOUT, write) => result,
}
}
}
pub(crate) fn websocket_relay_frame_queue<T>() -> (mpsc::Sender<T>, mpsc::Receiver<T>) {
mpsc::channel(RELAY_FRAME_QUEUE_CAPACITY)
}
/// Relays one frame to the client under [`RELAY_WRITE_TIMEOUT`].
pub(crate) async fn send_client_message(
client_socket: &mut WebSocket,
@@ -500,8 +569,9 @@ mod tests {
use super::{
bounded_send, guarded_websocket_upstream_url, responses_websocket_error_event,
responses_websocket_error_event_with_stream_id, websocket_handshake_headers,
websocket_response_headers, websocket_upstream_url, WebSocketWriteError,
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
websocket_relay_frame_queue, websocket_response_headers, websocket_upstream_url,
WebSocketRelayPumpControl, WebSocketRelayQueueError, WebSocketWriteError,
RELAY_FRAME_QUEUE_CAPACITY, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
};
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
use axum::http::HeaderMap;
@@ -524,6 +594,64 @@ mod tests {
assert_eq!(outcome, Err(WebSocketWriteError::Failed));
assert_eq!(WebSocketWriteError::Failed.as_str(), "write_failed");
assert_eq!(WebSocketWriteError::TimedOut.as_str(), "write_timeout");
assert_eq!(WebSocketWriteError::Cancelled.as_str(), "write_cancelled");
}
#[tokio::test]
async fn relay_frame_queue_is_bounded_and_fifo() {
let (sender, mut receiver) = websocket_relay_frame_queue();
for frame in 0..RELAY_FRAME_QUEUE_CAPACITY {
sender
.try_send(frame)
.expect("the configured burst buffer should accept this frame");
}
assert!(matches!(
sender.try_send(RELAY_FRAME_QUEUE_CAPACITY),
Err(tokio::sync::mpsc::error::TrySendError::Full(_))
));
for expected in 0..RELAY_FRAME_QUEUE_CAPACITY {
assert_eq!(receiver.recv().await, Some(expected));
}
}
#[tokio::test]
async fn relay_cancellation_interrupts_a_full_queue_without_waiting_for_capacity() {
let control = WebSocketRelayPumpControl::new();
let (sender, _receiver) = websocket_relay_frame_queue();
for frame in 0..RELAY_FRAME_QUEUE_CAPACITY {
sender.try_send(frame).expect("queue should fill exactly");
}
let enqueue = control.enqueue(&sender, RELAY_FRAME_QUEUE_CAPACITY);
tokio::pin!(enqueue);
assert!(tokio::time::timeout(Duration::from_millis(5), &mut enqueue)
.await
.is_err());
control.cancel();
assert_eq!(
tokio::time::timeout(Duration::from_millis(100), enqueue)
.await
.expect("cancellation should wake a blocked producer"),
Err(WebSocketRelayQueueError::Cancelled)
);
}
#[tokio::test]
async fn relay_cancellation_interrupts_a_stalled_socket_write() {
let control = WebSocketRelayPumpControl::new();
let write = control.send(std::future::pending::<Result<(), ()>>());
tokio::pin!(write);
assert!(tokio::time::timeout(Duration::from_millis(5), &mut write)
.await
.is_err());
control.cancel();
assert_eq!(
tokio::time::timeout(Duration::from_millis(100), write)
.await
.expect("cancellation should wake a stalled writer"),
Err(WebSocketWriteError::Cancelled)
);
}
#[tokio::test]
@@ -65,6 +65,45 @@ fn parse_users_me_usage_offset(query: Option<&str>) -> Result<usize, String> {
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
struct UsersMeUsageRecordFilter {
statuses: Option<Vec<String>>,
is_stream: Option<bool>,
is_websocket: Option<bool>,
error_only: bool,
}
fn parse_users_me_usage_record_filter(query: Option<&str>) -> UsersMeUsageRecordFilter {
let Some(status) = query_param_value(query, "status")
.map(|value| value.trim().to_ascii_lowercase())
.filter(|value| !value.is_empty())
else {
return UsersMeUsageRecordFilter::default();
};
let mut filter = UsersMeUsageRecordFilter::default();
match status.as_str() {
"stream" => {
filter.is_stream = Some(true);
filter.is_websocket = Some(false);
}
"standard" => {
filter.is_stream = Some(false);
filter.is_websocket = Some(false);
}
"websocket" | "ws" => filter.is_websocket = Some(true),
"error" | "failed" => filter.error_only = true,
"active" => {
filter.statuses = Some(vec!["pending".to_string(), "streaming".to_string()]);
}
"pending" | "streaming" | "completed" | "cancelled" => {
filter.statuses = Some(vec![status]);
}
_ => {}
}
filter
}
fn parse_users_me_usage_hours(query: Option<&str>) -> Result<u32, String> {
match query_param_value(query, "hours") {
Some(value) => parse_bounded_u32("hours", &value, 1, 720),
@@ -480,6 +519,11 @@ fn build_users_me_usage_record_payload(
"first_byte_time_ms": item.first_byte_time_ms,
"is_stream": item.is_stream,
"is_websocket": item.is_websocket(),
"websocket_transport": item.websocket_transport(),
"usage_available": item.usage_available(),
"usage_pricing_available": item.usage_pricing_available(),
"input_audio_tokens": item.realtime_input_audio_tokens(),
"output_audio_tokens": item.realtime_output_audio_tokens(),
"upstream_is_stream": upstream_is_stream,
"client_requested_stream": client_is_stream,
"client_is_stream": client_is_stream,
@@ -564,6 +608,11 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
"endpoint_api_format": item.endpoint_api_format,
"is_stream": item.is_stream,
"is_websocket": item.is_websocket(),
"websocket_transport": item.websocket_transport(),
"usage_available": item.usage_available(),
"usage_pricing_available": item.usage_pricing_available(),
"input_audio_tokens": item.realtime_input_audio_tokens(),
"output_audio_tokens": item.realtime_output_audio_tokens(),
"upstream_is_stream": upstream_is_stream,
"client_requested_stream": client_is_stream,
"client_is_stream": client_is_stream,
@@ -947,6 +996,7 @@ pub(super) async fn handle_users_me_usage_get(
Ok(value) => value,
Err(detail) => return admin_stats_bad_request_response(detail),
};
let record_filter = parse_users_me_usage_record_filter(query);
// When no time range is specified, default to 7 days to avoid full-table scans.
let effective_time_range = time_range.or_else(|| {
@@ -1087,10 +1137,11 @@ pub(super) async fn handle_users_me_usage_get(
api_format: None,
client_family: None,
exclude_unknown_model_or_provider: false,
statuses: None,
statuses: record_filter.statuses.clone(),
exclude_status_codes: Vec::new(),
is_stream: None,
error_only: false,
is_stream: record_filter.is_stream,
is_websocket: record_filter.is_websocket,
error_only: record_filter.error_only,
keywords,
matched_user_ids_by_keyword: Vec::new(),
auth_user_reader_available: false,
@@ -1143,10 +1194,11 @@ pub(super) async fn handle_users_me_usage_get(
api_format: None,
client_family: None,
exclude_unknown_model_or_provider: false,
statuses: None,
statuses: record_filter.statuses.clone(),
exclude_status_codes: Vec::new(),
is_stream: None,
error_only: false,
is_stream: record_filter.is_stream,
is_websocket: record_filter.is_websocket,
error_only: record_filter.error_only,
limit: None,
offset: None,
newest_first: true,
@@ -1172,10 +1224,11 @@ pub(super) async fn handle_users_me_usage_get(
api_format: None,
client_family: None,
exclude_unknown_model_or_provider: false,
statuses: None,
statuses: record_filter.statuses.clone(),
exclude_status_codes: Vec::new(),
is_stream: None,
error_only: false,
is_stream: record_filter.is_stream,
is_websocket: record_filter.is_websocket,
error_only: record_filter.error_only,
limit: Some(limit),
offset: Some(offset),
newest_first: true,
@@ -1314,6 +1367,7 @@ pub(super) async fn handle_users_me_usage_active_get(
statuses: Some(vec!["pending".to_string(), "streaming".to_string()]),
exclude_status_codes: Vec::new(),
is_stream: None,
is_websocket: None,
error_only: false,
limit: Some(50),
offset: None,
@@ -1568,10 +1622,34 @@ mod tests {
use super::{
build_users_me_usage_active_payload, build_users_me_usage_record_payload,
users_me_usage_client_is_stream, users_me_usage_is_failed,
users_me_usage_terminal_candidate_state_override, users_me_usage_upstream_is_stream,
parse_users_me_usage_record_filter, users_me_usage_client_is_stream,
users_me_usage_is_failed, users_me_usage_terminal_candidate_state_override,
users_me_usage_upstream_is_stream,
};
#[test]
fn users_me_usage_transport_statuses_are_disjoint_server_side_filters() {
for status in ["websocket", "ws", "WS"] {
let filter = parse_users_me_usage_record_filter(Some(
format!("limit=20&status={status}").as_str(),
));
assert_eq!(filter.is_websocket, Some(true));
assert_eq!(filter.is_stream, None);
assert_eq!(filter.statuses, None);
assert!(!filter.error_only);
}
for (status, expected_stream) in [("stream", true), ("standard", false)] {
let filter = parse_users_me_usage_record_filter(Some(
format!("limit=20&status={status}").as_str(),
));
assert_eq!(filter.is_stream, Some(expected_stream));
assert_eq!(filter.is_websocket, Some(false));
assert_eq!(filter.statuses, None);
assert!(!filter.error_only);
}
}
fn sample_usage(status: &str) -> StoredRequestUsageAudit {
StoredRequestUsageAudit::new(
"usage-1".to_string(),
@@ -1731,6 +1809,12 @@ mod tests {
request_metadata: Some(json!({
"websocket_mode": true,
"websocket_transport": "responses",
"usage_available": false,
"usage_pricing_available": false,
"realtime_session": {
"input_audio_tokens": 7,
"output_audio_tokens": 3,
},
})),
..sample_usage("completed")
};
@@ -1740,6 +1824,16 @@ mod tests {
assert_eq!(record["is_websocket"], true);
assert_eq!(active["is_websocket"], true);
assert_eq!(record["websocket_transport"], "responses");
assert_eq!(active["websocket_transport"], "responses");
assert_eq!(record["usage_available"], false);
assert_eq!(active["usage_available"], false);
assert_eq!(record["usage_pricing_available"], false);
assert_eq!(active["usage_pricing_available"], false);
assert_eq!(record["input_audio_tokens"], 7);
assert_eq!(active["input_audio_tokens"], 7);
assert_eq!(record["output_audio_tokens"], 3);
assert_eq!(active["output_audio_tokens"], 3);
}
#[test]