feat(gateway): add Codex Live transport

This commit is contained in:
ZheFox
2026-08-20 21:59:05 +08:00
parent 6916e9da76
commit 4185ad1b1e
35 changed files with 6683 additions and 85 deletions
+24 -2
View File
@@ -9,6 +9,7 @@ use self::body_buffer::{
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::responses::responses_websocket;
use super::internal::resolve_local_proxy_execution_path;
pub(crate) use super::public::matches_model_mapping_for_models;
@@ -106,6 +107,7 @@ const AUTH_API_KEY_CONCURRENCY_LIMIT_REACHED_DETAIL: &str =
const LOCAL_EXECUTION_PLANNING_TIMEOUT_DETAIL: &str =
"当前 AI 请求在本地执行规划阶段超时,请稍后重试";
const EXECUTION_PATH_TUNNEL_AFFINITY_FORWARD: &str = "tunnel_affinity_forward";
const EXECUTION_PATH_CODEX_LIVE_CALL: &str = "codex_live_call";
const MANAGEMENT_TOKEN_PREFIX: &str = "ae-";
const LEGACY_MANAGEMENT_TOKEN_PREFIX: &str = "ae_";
fn finalize_request_body_buffer_rejection(
@@ -939,11 +941,11 @@ pub(crate) async fn proxy_request(
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
request: Request,
) -> Result<Response<Body>, GatewayError> {
crate::request_diagnostics::scope_request_diagnostics(proxy_request_inner(
crate::request_diagnostics::scope_request_diagnostics(Box::pin(proxy_request_inner(
state,
remote_addr,
request,
))
)))
.await
}
@@ -1567,6 +1569,26 @@ async fn proxy_request_inner(
));
}
if let Some(response) = Box::pin(maybe_handle_live_http(
&state,
&request_context,
&parts,
buffered_body.as_ref(),
&remote_addr,
))
.await?
{
return Ok(finalize_gateway_response_with_context(
&state,
response,
&remote_addr,
&request_context,
EXECUTION_PATH_CODEX_LIVE_CALL,
&started_at,
request_permit.take(),
));
}
let local_ai_public_started_at = Instant::now();
let local_ai_public_response = super::public::maybe_build_local_ai_public_response(
&state,
@@ -50,6 +50,82 @@ pub(crate) struct WebSocketIngressSpec {
pub(crate) route_unavailable_message: &'static str,
}
pub(crate) enum AuthenticatedAiWebSocketUpgradePreparation {
Ready(AuthenticatedAiWebSocketUpgrade),
Rejected(Response<Body>),
}
/// Authenticated HTTP Upgrade state retained while an adapter performs any
/// protocol-specific preflight that must complete before status 101 is sent.
pub(crate) struct AuthenticatedAiWebSocketUpgrade {
state: AppState,
context: WebSocketRequestContext,
request_permit: Option<aether_runtime::AdmissionPermit>,
}
impl AuthenticatedAiWebSocketUpgrade {
pub(crate) fn state(&self) -> &AppState {
&self.state
}
pub(crate) fn context(&self) -> &WebSocketRequestContext {
&self.context
}
pub(crate) fn rejection_response(
&self,
status: StatusCode,
message: &str,
) -> Result<Response<Body>, GatewayError> {
build_local_http_error_response(
self.context.trace_id.as_str(),
Some(&self.context.decision),
status,
message,
)
}
pub(crate) fn into_response<F, Fut>(
self,
ws: WebSocketUpgrade,
limits: WebSocketSessionLimits,
run_session: F,
) -> Response<Body>
where
F: FnOnce(WebSocket, AppState, WebSocketRequestContext) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
self.into_response_with(ws, limits, (), move |socket, state, context, ()| {
run_session(socket, state, context)
})
}
pub(crate) fn into_response_with<P, F, Fut>(
self,
ws: WebSocketUpgrade,
limits: WebSocketSessionLimits,
prepared: P,
run_session: F,
) -> Response<Body>
where
P: Send + 'static,
F: FnOnce(WebSocket, AppState, WebSocketRequestContext, P) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
let Self {
state,
context,
request_permit,
} = self;
ws.max_frame_size(limits.max_frame_size)
.max_message_size(limits.max_message_size)
.on_upgrade(move |socket| async move {
drop(request_permit);
run_session(socket, state, context, prepared).await;
})
}
}
/// Performs the HTTP-only part of an AI WebSocket request.
///
/// The ordinary request permit covers only the HTTP Upgrade window. A
@@ -69,6 +145,21 @@ where
F: FnOnce(WebSocket, AppState, WebSocketRequestContext) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
match prepare_authenticated_ai_websocket(state, remote_addr, headers, uri, spec).await? {
AuthenticatedAiWebSocketUpgradePreparation::Ready(prepared) => {
Ok(prepared.into_response(ws, limits, run_session))
}
AuthenticatedAiWebSocketUpgradePreparation::Rejected(response) => Ok(response),
}
}
pub(crate) async fn prepare_authenticated_ai_websocket(
state: AppState,
remote_addr: SocketAddr,
headers: HeaderMap,
uri: Uri,
spec: WebSocketIngressSpec,
) -> Result<AuthenticatedAiWebSocketUpgradePreparation, GatewayError> {
let trace_id = extract_or_generate_trace_id(&headers);
let client_ip = effective_client_ip(&headers, &remote_addr);
if state.admin_security_ip_blacklisted(client_ip).await? {
@@ -77,7 +168,8 @@ where
None,
StatusCode::FORBIDDEN,
"当前 IP 已被禁止访问",
);
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
}
let request_context = crate::control::resolve_public_request_context(
@@ -94,10 +186,12 @@ where
None,
StatusCode::NOT_FOUND,
spec.route_unavailable_message,
);
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
};
if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) {
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection);
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
}
// Browsers attach cookies to WebSocket handshakes automatically and the
// WebSocket API does not let callers add an Authorization header. A
@@ -119,14 +213,16 @@ where
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
}
let Some(auth_context) = decision.auth_context.as_ref() else {
return build_local_auth_rejection_response(
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
};
if !auth_context.access_allowed
|| auth_context.user_id.trim().is_empty()
@@ -136,7 +232,8 @@ where
&trace_id,
Some(&decision),
&GatewayLocalAuthRejection::InvalidApiKey,
);
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
}
if !ip_rules_allow(auth_context.ip_rules.as_deref(), client_ip) {
return build_local_auth_rejection_response(
@@ -145,7 +242,8 @@ where
&GatewayLocalAuthRejection::IpNotAllowed {
remote_ip: client_ip.to_string(),
},
);
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected);
}
let request_permit = match state.try_acquire_request_permit().await {
@@ -157,6 +255,7 @@ where
Some(uri.path()),
error,
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected)
}
};
let websocket_connection_permit = match state.try_acquire_websocket_connection_permit().await {
@@ -168,6 +267,7 @@ where
Some(uri.path()),
error,
)
.map(AuthenticatedAiWebSocketUpgradePreparation::Rejected)
}
};
@@ -188,13 +288,13 @@ where
decision,
websocket_connection_permit,
};
Ok(ws
.max_frame_size(limits.max_frame_size)
.max_message_size(limits.max_message_size)
.on_upgrade(move |socket| async move {
drop(request_permit);
run_session(socket, state, context).await;
}))
Ok(AuthenticatedAiWebSocketUpgradePreparation::Ready(
AuthenticatedAiWebSocketUpgrade {
state,
context,
request_permit,
},
))
}
fn websocket_credential_carrier_is_allowed(carrier: Option<GatewayCredentialCarrier>) -> bool {
@@ -353,7 +453,7 @@ impl WebSocketConnectionLog {
spec,
trace_id: context.trace_id.clone(),
remote_addr: context.remote_addr,
path: context.uri.path().to_string(),
path: websocket_log_path(context.uri.path()),
route_class: context
.decision
.route_class
@@ -392,6 +492,17 @@ impl WebSocketConnectionLog {
}
}
fn websocket_log_path(path: &str) -> String {
if path
.strip_prefix("/v1/live/")
.is_some_and(|call_id| !call_id.is_empty())
{
"/v1/live/{call_id}".to_string()
} else {
path.to_string()
}
}
impl Drop for WebSocketConnectionLog {
fn drop(&mut self) {
info!(
@@ -424,7 +535,8 @@ mod tests {
use axum::http::{HeaderMap, HeaderValue, Uri};
use super::{
websocket_credential_carrier_is_allowed, websocket_planning_headers, websocket_planning_uri,
websocket_credential_carrier_is_allowed, websocket_log_path, websocket_planning_headers,
websocket_planning_uri,
};
use crate::control::GatewayCredentialCarrier;
@@ -524,4 +636,13 @@ mod tests {
assert!(websocket_credential_carrier_is_allowed(carrier));
}
}
#[test]
fn live_sideband_access_logs_do_not_retain_the_opaque_call_id() {
assert_eq!(websocket_log_path("/v1/live"), "/v1/live");
assert_eq!(
websocket_log_path("/v1/live/rtc_secret_opaque"),
"/v1/live/{call_id}"
);
}
}
@@ -0,0 +1,731 @@
//! Authenticated WebRTC call creation for Codex Live.
use std::collections::BTreeMap;
use std::net::SocketAddr;
use aether_contracts::{
ExecutionPlan, ExecutionResponseBodyMode, ExecutionResult, EXECUTION_RESPONSE_BODY_MODE_HEADER,
};
use axum::body::{Body, Bytes};
use axum::http::{HeaderValue, Response, StatusCode};
use base64::Engine as _;
use serde_json::json;
use tracing::{info, warn};
use crate::ai_serving::build_standard_sync_plan_from_decision;
use crate::api::response::{
build_client_response_from_parts, build_client_response_from_parts_with_mutator,
build_local_auth_rejection_response, build_local_http_error_response_with_request_path,
};
use crate::control::{execution_plan_balance_capacity_rejection, GatewayPublicRequestContext};
use crate::execution_runtime::execute_execution_runtime_sync_plan_with_report_context;
use crate::handlers::proxy::websocket::responses::ResponsesWebSocketTurnAdmission;
use crate::{AppState, GatewayError};
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};
use super::registry::{LiveCallBinding, LiveCallRegistry};
const MAX_LIVE_HTTP_BODY_BYTES: usize = 1024 * 1024;
pub(crate) async fn maybe_handle_live_http(
state: &AppState,
request_context: &GatewayPublicRequestContext,
parts: &http::request::Parts,
body: Option<&Bytes>,
remote_addr: &SocketAddr,
) -> Result<Option<Response<Body>>, GatewayError> {
if parts.method != http::Method::POST || request_context.request_path != "/v1/live" {
return Ok(None);
}
let Some(control_decision) = request_context.control_decision.as_ref() else {
return Ok(Some(local_live_error(
request_context,
StatusCode::NOT_FOUND,
"Codex Live route is unavailable",
)?));
};
if !live_usage_accounting_is_safe(control_decision) {
return Ok(Some(local_live_error(
request_context,
StatusCode::NOT_IMPLEMENTED,
"Codex Live is unavailable for finite-balance keys until Frameless usage settlement is supported",
)?));
}
let Some(body) = body else {
return Ok(Some(local_live_error(
request_context,
StatusCode::BAD_REQUEST,
"Codex Live requires a multipart WebRTC offer",
)?));
};
if body.len() > MAX_LIVE_HTTP_BODY_BYTES {
return Ok(Some(local_live_error(
request_context,
StatusCode::PAYLOAD_TOO_LARGE,
"Codex Live WebRTC offer exceeds the 1 MiB limit",
)?));
}
let content_type = parts
.headers
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default();
let offer = match parse_live_multipart(content_type, body.as_ref()) {
Ok(offer) => offer,
Err(error) => {
return Ok(Some(local_live_error(
request_context,
error.status_code(),
error.client_message(),
)?))
}
};
let Some(client_model) = offer
.session
.get("model")
.and_then(serde_json::Value::as_str)
else {
return Ok(Some(local_live_error(
request_context,
StatusCode::BAD_REQUEST,
"Codex Live session.model must be a non-empty model identifier",
)?));
};
let Some(mut candidate) = plan_live_candidate(
state,
request_context.trace_id.as_str(),
control_decision,
&parts.headers,
remote_addr,
client_model,
None,
)
.await?
else {
return Ok(Some(local_live_error(
request_context,
StatusCode::BAD_GATEWAY,
"No eligible Codex Live provider mapping is available",
)?));
};
let lease = LivePoolLeaseGuard::new(state, &candidate);
let binding = LiveCallBinding::from_candidate(&candidate);
let mut provider_session = offer.session.clone();
provider_session
.as_object_mut()
.expect("validated Live session is a JSON object")
.insert(
"model".to_string(),
serde_json::Value::String(candidate.provider_model.clone()),
);
let upstream_url = match live_call_url(&candidate) {
Ok(url) => url,
Err(error) => {
lease.release().await;
return Ok(Some(local_live_error(
request_context,
error.status_code(),
error.client_message(),
)?));
}
};
let (provider_content_type, provider_body_base64) =
build_live_call_provider_body(candidate.auth_mode, offer.sdp.as_str(), &provider_session)?;
// The standard plan builder requires a JSON body marker even when the exact wire body is
// carried as bytes. Keep only the mapped model here: retaining the SDP/session projection in
// the decision would unnecessarily widen the surface for future logging or report changes.
let provider_body_marker = json!({"model": candidate.provider_model.clone()});
candidate.execution.upstream_url = Some(upstream_url);
candidate.execution.provider_request_method = Some("POST".to_string());
candidate.execution.provider_request_body = Some(provider_body_marker.clone());
candidate.execution.provider_request_body_base64 = Some(provider_body_base64);
candidate.execution.content_type = Some(provider_content_type.clone());
candidate.execution.content_encoding = None;
candidate.execution.request_gzip = None;
candidate.execution.upstream_is_stream = false;
prepare_live_call_request_headers(
&mut candidate.execution.provider_request_headers,
provider_content_type.as_str(),
);
candidate.execution.provider_request_headers.insert(
EXECUTION_RESPONSE_BODY_MODE_HEADER.to_string(),
ExecutionResponseBodyMode::PreserveBytes
.as_str()
.to_string(),
);
let Some(attempt) =
build_standard_sync_plan_from_decision(parts, &provider_body_marker, candidate.execution)?
else {
lease.release().await;
return Ok(Some(local_live_error(
request_context,
StatusCode::BAD_GATEWAY,
"Codex Live provider request could not be built",
)?));
};
if let Some(rejection) = execution_plan_balance_capacity_rejection(
state,
control_decision,
&attempt.plan,
attempt.report_context.as_ref(),
)
.await?
{
lease.release().await;
return Ok(Some(build_local_auth_rejection_response(
request_context.trace_id.as_str(),
Some(control_decision),
&rejection,
)?));
}
let admission = ResponsesWebSocketTurnAdmission::acquire(
state,
&attempt.plan,
request_context.trace_id.as_str(),
)
.await?;
let result = execute_execution_runtime_sync_plan_with_report_context(
state,
Some(request_context.trace_id.as_str()),
&attempt.plan,
attempt.report_context.as_ref(),
)
.await;
// These guards intentionally cover only the synchronous call-creation exchange. The
// WebRTC media leg bypasses Aether, and neither the two-hour routing binding nor a sideband
// attachment proves that media is still alive. Holding either guard for a guessed lifetime
// would leak capacity or release it early without an authoritative upstream close signal.
admission.release().await;
let pool_lease_healthy = lease.is_healthy();
lease.release().await;
let result = result?;
if !(200..300).contains(&result.status_code) {
let response_body = execution_result_body(&result)?;
let downstream_headers =
sanitized_live_response_headers(&result.headers, response_body.preserves_wire_encoding);
warn!(
event_name = "codex_live_call_upstream_failed",
log_type = "ops",
trace_id = %request_context.trace_id,
provider_id = %attempt.plan.provider_id,
endpoint_id = %attempt.plan.endpoint_id,
key_id = %attempt.plan.key_id,
status_code = result.status_code,
elapsed_ms = result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
"Codex Live call creation failed upstream"
);
return Ok(Some(build_client_response_from_parts(
result.status_code,
&downstream_headers,
Body::from(response_body.bytes),
request_context.trace_id.as_str(),
Some(control_decision),
)?));
}
if !pool_lease_healthy {
warn_live_call_orphaned(
request_context,
&attempt.plan,
&result,
"pool_lease_lost",
None,
);
return Ok(Some(local_live_error(
request_context,
StatusCode::SERVICE_UNAVAILABLE,
"Codex Live provider lease expired during call creation",
)?));
}
let response_body = match execution_result_body(&result) {
Ok(body) => body,
Err(error) => {
warn_live_call_orphaned(
request_context,
&attempt.plan,
&result,
"response_body_unavailable",
None,
);
return Err(error);
}
};
let Some(location) = header_value(&result.headers, "location") else {
warn_live_call_orphaned(
request_context,
&attempt.plan,
&result,
"location_missing",
None,
);
return Ok(Some(local_live_error(
request_context,
StatusCode::BAD_GATEWAY,
"Codex Live upstream response did not include a call location",
)?));
};
let call_id = match extract_call_id_from_location(location) {
Ok(call_id) => call_id,
Err(error) => {
warn_live_call_orphaned(
request_context,
&attempt.plan,
&result,
"location_invalid",
Some(error.code()),
);
return Ok(Some(local_live_error(
request_context,
StatusCode::BAD_GATEWAY,
error.client_message(),
)?));
}
};
let Some(auth_context) = control_decision.auth_context.as_ref() else {
warn_live_call_orphaned(
request_context,
&attempt.plan,
&result,
"auth_context_missing",
None,
);
return Ok(Some(local_live_error(
request_context,
StatusCode::UNAUTHORIZED,
"Codex Live requires an authenticated gateway API key",
)?));
};
let registry = LiveCallRegistry::new(std::sync::Arc::clone(&state.runtime_state));
if let Err(error) = registry
.register(
auth_context.user_id.as_str(),
auth_context.api_key_id.as_str(),
call_id.as_str(),
&binding,
)
.await
{
warn_live_call_orphaned(
request_context,
&attempt.plan,
&result,
"binding_failed",
Some(error.kind()),
);
return Ok(Some(local_live_error(
request_context,
StatusCode::SERVICE_UNAVAILABLE,
"Codex Live sideband binding is temporarily unavailable",
)?));
}
info!(
event_name = "codex_live_call_created",
log_type = "event",
trace_id = %request_context.trace_id,
provider_id = %attempt.plan.provider_id,
endpoint_id = %attempt.plan.endpoint_id,
key_id = %attempt.plan.key_id,
client_model = %binding.client_model(),
status_code = result.status_code,
elapsed_ms = result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
usage_unavailable = true,
"Codex Live created a bound WebRTC call"
);
let downstream_location = format!("/v1/live/{call_id}");
let downstream_headers =
sanitized_live_response_headers(&result.headers, response_body.preserves_wire_encoding);
Ok(Some(build_client_response_from_parts_with_mutator(
result.status_code,
&downstream_headers,
Body::from(response_body.bytes),
request_context.trace_id.as_str(),
Some(control_decision),
|headers| {
headers.insert(
http::header::LOCATION,
HeaderValue::from_str(downstream_location.as_str())
.map_err(|error| GatewayError::Internal(error.to_string()))?,
);
Ok(())
},
)?))
}
fn local_live_error(
request_context: &GatewayPublicRequestContext,
status: StatusCode,
message: &str,
) -> Result<Response<Body>, GatewayError> {
build_local_http_error_response_with_request_path(
request_context.trace_id.as_str(),
request_context.control_decision.as_ref(),
Some("/v1/live"),
status,
message,
)
}
fn remove_headers(headers: &mut BTreeMap<String, String>, names: &[&str]) {
headers.retain(|candidate, _| {
!names
.iter()
.any(|name| candidate.eq_ignore_ascii_case(name))
});
}
fn header_value<'a>(headers: &'a BTreeMap<String, String>, name: &str) -> Option<&'a str> {
headers
.iter()
.find(|(candidate, _)| candidate.eq_ignore_ascii_case(name))
.map(|(_, value)| value.as_str())
}
fn prepare_live_call_request_headers(headers: &mut BTreeMap<String, String>, content_type: &str) {
remove_headers(
headers,
&[
"content-type",
"content-length",
"content-encoding",
"accept",
"accept-encoding",
],
);
headers.insert("content-type".to_string(), content_type.to_string());
headers.insert("accept".to_string(), "application/sdp".to_string());
headers.insert("accept-encoding".to_string(), "identity".to_string());
}
fn build_live_call_provider_body(
auth_mode: LiveAuthMode,
sdp: &str,
session: &serde_json::Value,
) -> Result<(String, String), GatewayError> {
let (content_type, bytes) = match auth_mode {
LiveAuthMode::ApiKey => {
let (content_type, bytes) = build_live_multipart(sdp, session);
(content_type, bytes)
}
LiveAuthMode::ChatGptOauth => {
let bytes = serde_json::to_vec(&json!({"sdp": sdp, "session": session}))
.map_err(|error| GatewayError::Internal(error.to_string()))?;
("application/json".to_string(), bytes)
}
};
Ok((
content_type,
base64::engine::general_purpose::STANDARD.encode(bytes),
))
}
fn sanitized_live_response_headers(
headers: &BTreeMap<String, String>,
preserves_wire_encoding: bool,
) -> BTreeMap<String, String> {
let mut sanitized = headers.clone();
remove_headers(&mut sanitized, &["location", "set-cookie", "set-cookie2"]);
if !preserves_wire_encoding {
remove_headers(&mut sanitized, &["content-length", "content-encoding"]);
}
sanitized
}
fn warn_live_call_orphaned(
request_context: &GatewayPublicRequestContext,
plan: &ExecutionPlan,
result: &ExecutionResult,
reason: &'static str,
error_kind: Option<&'static str>,
) {
warn!(
event_name = "codex_live_call_orphaned",
log_type = "ops",
trace_id = %request_context.trace_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
status_code = result.status_code,
elapsed_ms = result.telemetry.as_ref().and_then(|value| value.elapsed_ms),
reason,
error_kind = error_kind.unwrap_or("none"),
"Codex Live upstream call succeeded but could not be safely exposed downstream"
);
}
struct LiveResponseBody {
bytes: Vec<u8>,
preserves_wire_encoding: bool,
}
fn execution_result_body(result: &ExecutionResult) -> Result<LiveResponseBody, GatewayError> {
let Some(body) = result.body.as_ref() else {
return Ok(LiveResponseBody {
bytes: Vec::new(),
preserves_wire_encoding: false,
});
};
if let Some(encoded) = body.body_bytes_b64.as_deref() {
let bytes = base64::engine::general_purpose::STANDARD
.decode(encoded)
.map_err(|error| GatewayError::Internal(error.to_string()))?;
return Ok(LiveResponseBody {
bytes,
preserves_wire_encoding: true,
});
}
let bytes = body
.json_body
.as_ref()
.map(serde_json::to_vec)
.transpose()
.map(|body| body.unwrap_or_default())
.map_err(|error| GatewayError::Internal(error.to_string()))?;
Ok(LiveResponseBody {
bytes,
preserves_wire_encoding: false,
})
}
#[cfg(test)]
mod tests {
use aether_contracts::{ExecutionPlan, ExecutionResult, ResponseBody};
use axum::body::to_bytes;
use crate::control::{GatewayControlAuthContext, GatewayControlDecision};
use super::*;
#[test]
fn preserved_wire_bytes_win_over_the_json_projection() {
let result = ExecutionResult {
request_id: "request".to_string(),
candidate_id: None,
status_code: 201,
headers: Default::default(),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(json!({"projected": true})),
body_bytes_b64: Some(
base64::engine::general_purpose::STANDARD.encode(b"raw-sdp-answer"),
),
}),
telemetry: None,
error: None,
};
let body = execution_result_body(&result).unwrap();
assert_eq!(body.bytes, b"raw-sdp-answer");
assert!(body.preserves_wire_encoding);
}
#[test]
fn live_call_bodies_are_bytes_only_and_do_not_enter_report_context() {
let session = json!({
"model": "provider-live-model",
"instructions": "opaque private instructions",
"future_capability": {"enabled": true}
});
let sdp = "v=0\r\no=private-live-offer";
let provider_body_marker = json!({"model": "provider-live-model"});
for auth_mode in [LiveAuthMode::ApiKey, LiveAuthMode::ChatGptOauth] {
let (content_type, encoded) =
build_live_call_provider_body(auth_mode, sdp, &session).unwrap();
let report_context = aether_ai_serving::augment_sync_report_context(
Some(json!({"trace_id": "trace-live"})),
&BTreeMap::new(),
&provider_body_marker,
)
.unwrap()
.unwrap();
assert!(report_context.get("provider_request_body").is_none());
assert!(!report_context.to_string().contains("private-live-offer"));
assert!(!report_context
.to_string()
.contains("opaque private instructions"));
let plan_body = aether_ai_serving::resolve_ai_passthrough_sync_request_body(
Some(provider_body_marker.clone()),
Some(encoded.clone()),
);
assert!(plan_body.json_body.is_none());
assert_eq!(plan_body.body_bytes_b64.as_deref(), Some(encoded.as_str()));
let usage_plan = ExecutionPlan {
request_id: "trace-live".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: "POST".to_string(),
url: "https://api.openai.com/v1/live".to_string(),
headers: BTreeMap::new(),
content_type: Some(content_type.clone()),
content_encoding: None,
body: plan_body,
stream: false,
client_api_format: "openai:responses".to_string(),
provider_api_format: "openai:responses".to_string(),
model_name: Some("provider-live-model".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
};
let usage_seed = aether_usage_runtime::build_terminal_usage_context_seed(
&usage_plan,
Some(&report_context),
);
assert!(usage_seed.provider_request.is_none());
assert!(!usage_seed
.request_metadata
.as_ref()
.map(ToString::to_string)
.unwrap_or_default()
.contains("private-live-offer"));
let wire = base64::engine::general_purpose::STANDARD
.decode(encoded)
.unwrap();
match auth_mode {
LiveAuthMode::ApiKey => {
let parsed = parse_live_multipart(content_type.as_str(), wire.as_slice())
.expect("API-key multipart should round-trip");
assert_eq!(parsed.sdp, sdp);
assert_eq!(parsed.session, session);
}
LiveAuthMode::ChatGptOauth => {
assert_eq!(content_type, "application/json");
let decoded: serde_json::Value = serde_json::from_slice(wire.as_slice())
.expect("OAuth JSON should round-trip");
assert_eq!(decoded["sdp"], sdp);
assert_eq!(decoded["session"], session);
assert_eq!(decoded["session"]["model"], "provider-live-model");
assert_eq!(
decoded["session"]["future_capability"],
json!({"enabled": true})
);
}
}
}
}
#[test]
fn live_call_request_headers_replace_stale_body_and_encoding_metadata() {
let mut headers = BTreeMap::from([
("Content-Type".to_string(), "stale".to_string()),
("CONTENT-LENGTH".to_string(), "42".to_string()),
("Content-Encoding".to_string(), "gzip".to_string()),
("Accept".to_string(), "application/json".to_string()),
("ACCEPT-ENCODING".to_string(), "br, gzip".to_string()),
("x-future".to_string(), "opaque".to_string()),
]);
prepare_live_call_request_headers(&mut headers, "multipart/form-data; boundary=live-test");
assert_eq!(
header_value(&headers, "content-type"),
Some("multipart/form-data; boundary=live-test")
);
assert_eq!(header_value(&headers, "accept"), Some("application/sdp"));
assert_eq!(header_value(&headers, "accept-encoding"), Some("identity"));
assert_eq!(header_value(&headers, "content-length"), None);
assert_eq!(header_value(&headers, "content-encoding"), None);
assert_eq!(headers.get("x-future").map(String::as_str), Some("opaque"));
}
#[test]
fn live_response_headers_never_expose_upstream_location_or_cookies() {
let headers = BTreeMap::from([
(
"Location".to_string(),
"https://upstream/v1/live/secret".to_string(),
),
("SET-COOKIE".to_string(), "session=secret".to_string()),
("Set-Cookie2".to_string(), "legacy=secret".to_string()),
("Content-Length".to_string(), "128".to_string()),
("Content-Encoding".to_string(), "gzip".to_string()),
("x-future".to_string(), "opaque".to_string()),
]);
let sanitized = sanitized_live_response_headers(&headers, true);
assert_eq!(header_value(&sanitized, "location"), None);
assert_eq!(header_value(&sanitized, "set-cookie"), None);
assert_eq!(header_value(&sanitized, "set-cookie2"), None);
assert_eq!(header_value(&sanitized, "content-length"), Some("128"));
assert_eq!(header_value(&sanitized, "content-encoding"), Some("gzip"));
assert_eq!(header_value(&sanitized, "x-future"), Some("opaque"));
}
#[test]
fn rebuilt_live_response_body_drops_stale_length_and_encoding() {
let headers = BTreeMap::from([
("content-length".to_string(), "128".to_string()),
("content-encoding".to_string(), "gzip".to_string()),
("content-type".to_string(), "application/json".to_string()),
]);
let sanitized = sanitized_live_response_headers(&headers, false);
assert_eq!(header_value(&sanitized, "content-length"), None);
assert_eq!(header_value(&sanitized, "content-encoding"), None);
assert_eq!(
header_value(&sanitized, "content-type"),
Some("application/json")
);
}
#[tokio::test]
async fn finite_balance_post_live_fails_before_parsing_or_upstream_execution() {
let mut decision = GatewayControlDecision::synthetic(
"/v1/live",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("codex_live".to_string()),
Some("openai:responses".to_string()),
);
decision.auth_context = Some(GatewayControlAuthContext {
user_id: "user-finite".to_string(),
api_key_id: "key-finite".to_string(),
username: Some("finite".to_string()),
api_key_name: Some("finite".to_string()),
balance_remaining: Some(1.25),
access_allowed: true,
user_rate_limit: None,
api_key_rate_limit: None,
api_key_is_standalone: false,
admin_bypass_limits: false,
local_rejection: None,
allowed_models: None,
ip_rules: None,
});
let request_context = GatewayPublicRequestContext {
trace_id: "trace-live-finite".to_string(),
request_method: http::Method::POST,
request_path: "/v1/live".to_string(),
request_query_string: None,
request_content_type: None,
host_header: None,
control_decision: Some(decision),
};
let (parts, _) = http::Request::builder()
.method(http::Method::POST)
.uri("/v1/live")
.body(())
.unwrap()
.into_parts();
let response = maybe_handle_live_http(
&AppState::new().expect("gateway state should build"),
&request_context,
&parts,
None,
&"127.0.0.1:65000".parse().unwrap(),
)
.await
.unwrap()
.expect("Live HTTP route must produce a local rejection");
assert_eq!(response.status(), StatusCode::NOT_IMPLEMENTED);
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
assert!(String::from_utf8_lossy(body.as_ref()).contains("finite-balance"));
}
}
@@ -0,0 +1,80 @@
//! Experimental Codex Frameless Bidi V3 (`/v1/live`) bridge.
//!
//! The public OpenAI Realtime and Responses WebSocket protocols are related
//! transport families, but this Codex protocol has a distinct event grammar.
//! Keeping it in an independent module prevents a `session.update` frame from
//! ever entering the Responses `response.create` state machine.
mod http;
mod planner;
mod protocol;
mod registry;
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::control::GatewayControlDecision;
use crate::handlers::proxy::websocket::ingress::{
prepare_authenticated_ai_websocket, AuthenticatedAiWebSocketUpgradePreparation,
WebSocketIngressSpec,
};
use crate::handlers::proxy::websocket::session::LIVE_WEBSOCKET_SESSION_LIMITS;
use crate::{AppState, GatewayError};
pub(crate) use http::maybe_handle_live_http;
/// Frameless Bidi currently exposes no stable token/cost usage object that can
/// be fed into Aether's settlement pipeline. Fail closed for finite-balance
/// principals instead of silently serving unmetered traffic. Standalone or
/// shared keys backed by an unlimited/no-wallet policy resolve without a
/// finite `balance_remaining` and remain eligible.
fn live_usage_accounting_is_safe(decision: &GatewayControlDecision) -> bool {
decision
.auth_context
.as_ref()
.is_some_and(|auth| auth.balance_remaining.is_none())
}
pub(crate) async fn live_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,
LIVE_WEBSOCKET_INGRESS_SPEC,
)
.await?
{
AuthenticatedAiWebSocketUpgradePreparation::Rejected(response) => Ok(response),
AuthenticatedAiWebSocketUpgradePreparation::Ready(prepared) => {
let live =
match session::prepare_live_websocket(prepared.state(), prepared.context()).await {
Ok(live) => live,
Err(rejection) => {
return prepared.rejection_response(rejection.status(), rejection.message())
}
};
Ok(prepared.into_response_with(
ws,
LIVE_WEBSOCKET_SESSION_LIMITS,
live,
session::run_live_websocket,
))
}
}
}
const LIVE_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec {
route_unavailable_message: "Codex Live WebSocket route is unavailable",
};
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,766 @@
//! Minimal validation at the Codex Live trust boundary.
//!
//! Live events remain opaque after the initial discriminator. This module owns
//! only bounded identifiers, multipart framing and the first `session.update`
//! check; it intentionally does not copy the evolving Codex event schema.
use axum::http::StatusCode;
use serde_json::Value;
const MAX_MODEL_BYTES: usize = 256;
const MAX_CALL_ID_BYTES: usize = 256;
const MAX_BOUNDARY_BYTES: usize = 70;
const MAX_MULTIPART_BODY_BYTES: usize = 1024 * 1024;
const MAX_SDP_BYTES: usize = 512 * 1024;
const MAX_SESSION_BYTES: usize = 256 * 1024;
const MAX_PART_HEADERS_BYTES: usize = 8 * 1024;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub(super) enum LiveProtocolError {
#[error("missing Live upstream URL")]
MissingUpstreamUrl,
#[error("invalid Live upstream URL")]
InvalidUpstreamUrl,
#[error("ChatGPT OAuth does not support direct Codex Live WebSocket")]
OauthDirectWebSocketUnsupported,
#[error("ChatGPT OAuth Codex Live requires the official backend origin")]
OauthUpstreamUnsupported,
#[error("invalid Live model query")]
InvalidModelQuery,
#[error("invalid Live model")]
InvalidModel,
#[error("invalid Live call ID")]
InvalidCallId,
#[error("unsupported Live media type")]
UnsupportedMediaType,
#[error("invalid Live multipart boundary")]
InvalidBoundary,
#[error("Live multipart body is too large")]
MultipartBodyTooLarge,
#[error("malformed Live multipart body")]
MalformedMultipart,
#[error("unexpected Live multipart part")]
UnexpectedMultipartPart,
#[error("duplicate Live multipart part")]
DuplicateMultipartPart,
#[error("missing Live SDP part")]
MissingSdp,
#[error("invalid Live SDP part")]
InvalidSdp,
#[error("Live SDP is too large")]
SdpTooLarge,
#[error("missing Live session part")]
MissingSession,
#[error("invalid Live session JSON")]
InvalidSession,
#[error("Live session is too large")]
SessionTooLarge,
#[error("invalid initial Live JSON event")]
InvalidInitialEvent,
#[error("initial Live event must be session.update")]
ExpectedSessionUpdate,
#[error("initial Live event must be text")]
InitialEventMustBeText,
#[error("initial Live client read failed")]
InitialClientReadFailed,
#[error("timed out waiting for initial Live session.update")]
InitialSessionUpdateTimeout,
#[error("invalid Live call location")]
InvalidCallLocation,
}
impl LiveProtocolError {
pub(super) const fn status_code(&self) -> StatusCode {
match self {
Self::MultipartBodyTooLarge | Self::SdpTooLarge | Self::SessionTooLarge => {
StatusCode::PAYLOAD_TOO_LARGE
}
Self::UnsupportedMediaType => StatusCode::UNSUPPORTED_MEDIA_TYPE,
Self::MissingUpstreamUrl
| Self::InvalidUpstreamUrl
| Self::OauthUpstreamUnsupported
| Self::InvalidCallLocation => StatusCode::BAD_GATEWAY,
Self::InitialSessionUpdateTimeout => StatusCode::REQUEST_TIMEOUT,
_ => StatusCode::BAD_REQUEST,
}
}
pub(super) const fn code(&self) -> &'static str {
match self {
Self::MissingUpstreamUrl => "codex_live_upstream_url_missing",
Self::InvalidUpstreamUrl => "codex_live_upstream_url_invalid",
Self::OauthDirectWebSocketUnsupported => "codex_live_oauth_direct_unsupported",
Self::OauthUpstreamUnsupported => "codex_live_oauth_upstream_unsupported",
Self::InvalidModelQuery => "codex_live_model_query_invalid",
Self::InvalidModel => "codex_live_model_invalid",
Self::InvalidCallId => "codex_live_call_id_invalid",
Self::UnsupportedMediaType => "codex_live_media_type_unsupported",
Self::InvalidBoundary => "codex_live_boundary_invalid",
Self::MultipartBodyTooLarge => "codex_live_body_too_large",
Self::MalformedMultipart => "codex_live_multipart_invalid",
Self::UnexpectedMultipartPart => "codex_live_multipart_part_unexpected",
Self::DuplicateMultipartPart => "codex_live_multipart_part_duplicate",
Self::MissingSdp => "codex_live_sdp_missing",
Self::InvalidSdp => "codex_live_sdp_invalid",
Self::SdpTooLarge => "codex_live_sdp_too_large",
Self::MissingSession => "codex_live_session_missing",
Self::InvalidSession => "codex_live_session_invalid",
Self::SessionTooLarge => "codex_live_session_too_large",
Self::InvalidInitialEvent => "codex_live_initial_event_invalid",
Self::ExpectedSessionUpdate => "codex_live_expected_session_update",
Self::InitialEventMustBeText => "codex_live_initial_event_must_be_text",
Self::InitialClientReadFailed => "codex_live_initial_client_read_failed",
Self::InitialSessionUpdateTimeout => "codex_live_initial_session_update_timeout",
Self::InvalidCallLocation => "codex_live_call_location_invalid",
}
}
pub(super) const fn client_message(&self) -> &'static str {
match self {
Self::MissingUpstreamUrl | Self::InvalidUpstreamUrl => {
"Codex Live provider URL is invalid"
}
Self::OauthDirectWebSocketUnsupported => {
"Direct Codex Live WebSocket requires an API-key provider; use WebRTC for ChatGPT OAuth"
}
Self::OauthUpstreamUnsupported => {
"ChatGPT OAuth Codex Live requires the official ChatGPT backend"
}
Self::InvalidModelQuery => {
"Codex Live WebSocket requires exactly one model query parameter"
}
Self::InvalidModel => {
"Codex Live model must be a non-empty identifier no longer than 256 bytes"
}
Self::InvalidCallId => "Codex Live call ID is invalid",
Self::UnsupportedMediaType => {
"Codex Live WebRTC call creation requires multipart/form-data"
}
Self::InvalidBoundary => "Codex Live multipart boundary is invalid",
Self::MultipartBodyTooLarge => "Codex Live WebRTC offer exceeds the 1 MiB limit",
Self::MalformedMultipart => "Codex Live multipart body is malformed",
Self::UnexpectedMultipartPart => {
"Codex Live multipart body may contain only sdp and session parts"
}
Self::DuplicateMultipartPart => "Codex Live multipart part is duplicated",
Self::MissingSdp => "Codex Live multipart body is missing the sdp part",
Self::InvalidSdp => "Codex Live sdp part must be non-empty UTF-8",
Self::SdpTooLarge => "Codex Live sdp part exceeds the 512 KiB limit",
Self::MissingSession => "Codex Live multipart body is missing the session part",
Self::InvalidSession => "Codex Live session part must be a JSON object",
Self::SessionTooLarge => "Codex Live session exceeds the 256 KiB limit",
Self::InvalidInitialEvent => {
"The initial Codex Live WebSocket text message must be a JSON object"
}
Self::ExpectedSessionUpdate => {
"Codex Live WebSocket must start with a session.update event"
}
Self::InitialEventMustBeText => {
"The initial Codex Live session.update event must be a text message"
}
Self::InitialClientReadFailed => {
"Failed to read the initial Codex Live session.update event"
}
Self::InitialSessionUpdateTimeout => {
"Timed out waiting for the initial Codex Live session.update event"
}
Self::InvalidCallLocation => {
"Codex Live upstream returned an invalid call location"
}
}
}
pub(super) const fn is_timeout(&self) -> bool {
matches!(self, Self::InitialSessionUpdateTimeout)
}
}
#[derive(Debug, Clone, PartialEq)]
pub(super) struct LiveMultipart {
pub(super) sdp: String,
pub(super) session: Value,
}
pub(super) fn validate_model(model: &str) -> Result<(), LiveProtocolError> {
if model.is_empty()
|| model.len() > MAX_MODEL_BYTES
|| model.trim() != model
|| model.chars().any(char::is_control)
{
return Err(LiveProtocolError::InvalidModel);
}
Ok(())
}
pub(super) fn validate_call_id(call_id: &str) -> Result<(), LiveProtocolError> {
if call_id.is_empty()
|| call_id.len() > MAX_CALL_ID_BYTES
|| matches!(call_id, "." | "..")
|| !call_id
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-' | b'.'))
{
return Err(LiveProtocolError::InvalidCallId);
}
Ok(())
}
fn is_live_call_id_segment(call_id: &str) -> bool {
if validate_call_id(call_id).is_err() {
return false;
}
if call_id.starts_with("rtc_") && call_id.len() > "rtc_".len() {
return true;
}
call_id.len() == 36
&& call_id
.bytes()
.enumerate()
.all(|(index, byte)| match index {
8 | 13 | 18 | 23 => byte == b'-',
_ => byte.is_ascii_hexdigit(),
})
}
pub(super) fn direct_model_from_query(query: Option<&str>) -> Result<String, LiveProtocolError> {
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(LiveProtocolError::InvalidModelQuery);
}
validate_model(value.as_ref())?;
model = Some(value.into_owned());
continue;
}
// Latest Codex preserves provider query parameters while rewriting
// `/v1/realtime` to `/v1/live`. They are downstream transport hints,
// not Aether routing authority, so accept and ignore non-credential
// parameters. Credential values should already have been consumed by
// ingress; rejecting them here keeps this parser safe in isolation.
if live_query_parameter_is_sensitive(name.as_ref()) {
return Err(LiveProtocolError::InvalidModelQuery);
}
}
model.ok_or(LiveProtocolError::InvalidModelQuery)
}
fn live_query_parameter_is_sensitive(name: &str) -> bool {
matches!(
name.to_ascii_lowercase().as_str(),
"key"
| "api_key"
| "api-key"
| "x-api-key"
| "x-goog-api-key"
| "access_token"
| "authorization"
| "token"
| "oauth_token"
| "client_secret"
| "secret_key"
| "signature"
| "sig"
)
}
pub(super) fn call_id_from_path(path: &str) -> Result<String, LiveProtocolError> {
let call_id = path
.strip_prefix("/v1/live/")
.ok_or(LiveProtocolError::InvalidCallId)?;
if call_id.contains('/') {
return Err(LiveProtocolError::InvalidCallId);
}
validate_call_id(call_id)?;
Ok(call_id.to_string())
}
pub(super) fn validate_initial_session_update(raw: &str) -> Result<(), LiveProtocolError> {
if raw.len() > MAX_SESSION_BYTES {
return Err(LiveProtocolError::SessionTooLarge);
}
let value: Value =
serde_json::from_str(raw).map_err(|_| LiveProtocolError::InvalidInitialEvent)?;
if !value.is_object() {
return Err(LiveProtocolError::InvalidInitialEvent);
}
if value.get("type").and_then(Value::as_str) != Some("session.update") {
return Err(LiveProtocolError::ExpectedSessionUpdate);
}
Ok(())
}
pub(super) fn event_type(raw: &str) -> Option<String> {
serde_json::from_str::<Value>(raw)
.ok()?
.get("type")?
.as_str()
.map(str::to_string)
}
pub(super) fn parse_live_multipart(
content_type: &str,
body: &[u8],
) -> Result<LiveMultipart, LiveProtocolError> {
if body.len() > MAX_MULTIPART_BODY_BYTES {
return Err(LiveProtocolError::MultipartBodyTooLarge);
}
let boundary = multipart_boundary(content_type)?;
let parts = parse_multipart_parts(body, boundary.as_bytes())?;
let mut sdp = None;
let mut session = None;
for part in parts {
match part.name.as_str() {
"sdp" => {
if sdp.is_some() {
return Err(LiveProtocolError::DuplicateMultipartPart);
}
if part.body.len() > MAX_SDP_BYTES {
return Err(LiveProtocolError::SdpTooLarge);
}
let value =
std::str::from_utf8(part.body).map_err(|_| LiveProtocolError::InvalidSdp)?;
if value.trim().is_empty() {
return Err(LiveProtocolError::InvalidSdp);
}
sdp = Some(value.to_string());
}
"session" => {
if session.is_some() {
return Err(LiveProtocolError::DuplicateMultipartPart);
}
if part.body.len() > MAX_SESSION_BYTES {
return Err(LiveProtocolError::SessionTooLarge);
}
let value: Value = serde_json::from_slice(part.body)
.map_err(|_| LiveProtocolError::InvalidSession)?;
if !value.is_object() {
return Err(LiveProtocolError::InvalidSession);
}
session = Some(value);
}
_ => return Err(LiveProtocolError::UnexpectedMultipartPart),
}
}
Ok(LiveMultipart {
sdp: sdp.ok_or(LiveProtocolError::MissingSdp)?,
session: session.ok_or(LiveProtocolError::MissingSession)?,
})
}
pub(super) fn build_live_multipart(sdp: &str, session: &Value) -> (String, Vec<u8>) {
let boundary = format!("aether-live-{}", uuid::Uuid::new_v4().simple());
let session = serde_json::to_vec(session).expect("a JSON value must serialize");
let mut body = Vec::with_capacity(sdp.len() + session.len() + 320);
append_part(
&mut body,
boundary.as_str(),
"sdp",
"application/sdp",
sdp.as_bytes(),
);
append_part(
&mut body,
boundary.as_str(),
"session",
"application/json",
session.as_slice(),
);
body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes());
(format!("multipart/form-data; boundary={boundary}"), body)
}
pub(super) fn extract_call_id_from_location(location: &str) -> Result<String, LiveProtocolError> {
let location = location.trim();
if location.is_empty() {
return Err(LiveProtocolError::InvalidCallLocation);
}
let path = if let Ok(url) = url::Url::parse(location) {
url.path().to_string()
} else {
location
.split_once('?')
.map_or(location, |(path, _)| path)
.to_string()
};
let call_id = path
.trim_end_matches('/')
.rsplit('/')
.next()
.filter(|value| !value.is_empty())
.ok_or(LiveProtocolError::InvalidCallLocation)?;
if !is_live_call_id_segment(call_id) {
return Err(LiveProtocolError::InvalidCallLocation);
}
Ok(call_id.to_string())
}
struct MultipartPart<'a> {
name: String,
body: &'a [u8],
}
fn multipart_boundary(content_type: &str) -> Result<String, LiveProtocolError> {
let mut values = content_type.split(';');
if !values
.next()
.is_some_and(|value| value.trim().eq_ignore_ascii_case("multipart/form-data"))
{
return Err(LiveProtocolError::UnsupportedMediaType);
}
let mut boundary = None;
for parameter in values {
let Some((name, value)) = parameter.trim().split_once('=') else {
continue;
};
if !name.trim().eq_ignore_ascii_case("boundary") {
continue;
}
if boundary.is_some() {
return Err(LiveProtocolError::InvalidBoundary);
}
let value = value.trim();
let value = if value.starts_with('"') && value.ends_with('"') && value.len() >= 2 {
&value[1..value.len() - 1]
} else {
value
};
boundary = Some(value.to_string());
}
let boundary = boundary.ok_or(LiveProtocolError::InvalidBoundary)?;
if boundary.is_empty()
|| boundary.len() > MAX_BOUNDARY_BYTES
|| !boundary.bytes().all(|byte| {
byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'\''
| b'('
| b')'
| b'+'
| b'_'
| b','
| b'-'
| b'.'
| b'/'
| b':'
| b'='
| b'?'
)
})
{
return Err(LiveProtocolError::InvalidBoundary);
}
Ok(boundary)
}
fn parse_multipart_parts<'a>(
body: &'a [u8],
boundary: &[u8],
) -> Result<Vec<MultipartPart<'a>>, LiveProtocolError> {
let delimiter = [b"--".as_slice(), boundary].concat();
if !body.starts_with(delimiter.as_slice()) {
return Err(LiveProtocolError::MalformedMultipart);
}
let mut cursor = delimiter.len();
let mut parts = Vec::new();
loop {
if body.get(cursor..cursor + 2) == Some(b"--") {
cursor += 2;
if body
.get(cursor..)
.is_some_and(|tail| tail.is_empty() || tail == b"\r\n")
{
return Ok(parts);
}
return Err(LiveProtocolError::MalformedMultipart);
}
if body.get(cursor..cursor + 2) != Some(b"\r\n") {
return Err(LiveProtocolError::MalformedMultipart);
}
cursor += 2;
let header_end = find_bytes(&body[cursor..], b"\r\n\r\n")
.ok_or(LiveProtocolError::MalformedMultipart)?;
if header_end > MAX_PART_HEADERS_BYTES {
return Err(LiveProtocolError::MalformedMultipart);
}
let headers = &body[cursor..cursor + header_end];
cursor += header_end + 4;
let marker = [b"\r\n--".as_slice(), boundary].concat();
let body_end = find_bytes(&body[cursor..], marker.as_slice())
.ok_or(LiveProtocolError::MalformedMultipart)?;
let part_body = &body[cursor..cursor + body_end];
let name = multipart_part_name(headers)?;
parts.push(MultipartPart {
name,
body: part_body,
});
if parts.len() > 2 {
return Err(LiveProtocolError::UnexpectedMultipartPart);
}
cursor += body_end + 2 + delimiter.len();
}
}
fn multipart_part_name(headers: &[u8]) -> Result<String, LiveProtocolError> {
let headers =
std::str::from_utf8(headers).map_err(|_| LiveProtocolError::MalformedMultipart)?;
let mut disposition = None;
for line in headers.split("\r\n") {
let Some((name, value)) = line.split_once(':') else {
return Err(LiveProtocolError::MalformedMultipart);
};
if name.trim().eq_ignore_ascii_case("content-disposition") {
if disposition.is_some() {
return Err(LiveProtocolError::MalformedMultipart);
}
disposition = Some(value.trim());
}
}
let disposition = disposition.ok_or(LiveProtocolError::MalformedMultipart)?;
let mut parameters = disposition.split(';');
if !parameters
.next()
.is_some_and(|value| value.trim().eq_ignore_ascii_case("form-data"))
{
return Err(LiveProtocolError::MalformedMultipart);
}
let mut part_name = None;
for parameter in parameters {
let Some((name, value)) = parameter.trim().split_once('=') else {
return Err(LiveProtocolError::MalformedMultipart);
};
if name.trim().eq_ignore_ascii_case("filename") {
return Err(LiveProtocolError::UnexpectedMultipartPart);
}
if name.trim().eq_ignore_ascii_case("name") {
if part_name.is_some() {
return Err(LiveProtocolError::MalformedMultipart);
}
let value = value.trim();
if !(value.starts_with('"') && value.ends_with('"') && value.len() >= 2) {
return Err(LiveProtocolError::MalformedMultipart);
}
part_name = Some(value[1..value.len() - 1].to_string());
}
}
part_name.ok_or(LiveProtocolError::MalformedMultipart)
}
fn find_bytes(haystack: &[u8], needle: &[u8]) -> Option<usize> {
(!needle.is_empty())
.then(|| {
haystack
.windows(needle.len())
.position(|value| value == needle)
})
.flatten()
}
fn append_part(body: &mut Vec<u8>, boundary: &str, name: &str, content_type: &str, value: &[u8]) {
body.extend_from_slice(format!("--{boundary}\r\n").as_bytes());
body.extend_from_slice(
format!("Content-Disposition: form-data; name=\"{name}\"\r\n").as_bytes(),
);
body.extend_from_slice(format!("Content-Type: {content_type}\r\n\r\n").as_bytes());
body.extend_from_slice(value);
body.extend_from_slice(b"\r\n");
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
#[test]
fn direct_query_requires_one_bounded_model() {
assert_eq!(
direct_model_from_query(Some("model=gpt-live%2Ffuture")).unwrap(),
"gpt-live/future"
);
assert_eq!(
direct_model_from_query(Some("foo=bar&model=gpt-live&trace=1")).unwrap(),
"gpt-live"
);
assert_eq!(
direct_model_from_query(Some("model=a&MODEL=b")),
Err(LiveProtocolError::InvalidModelQuery)
);
assert_eq!(
direct_model_from_query(Some("model=a&token=secret")),
Err(LiveProtocolError::InvalidModelQuery)
);
}
#[test]
fn direct_protocol_starts_with_session_update_not_response_create() {
let opaque = r#"{"type":"session.update","session":{"future_capability":{"version":2}},"future_event_field":[1,2,3]}"#;
validate_initial_session_update(opaque).unwrap();
assert_eq!(
serde_json::from_str::<Value>(opaque).unwrap()["future_event_field"],
json!([1, 2, 3])
);
assert_eq!(
validate_initial_session_update(r#"{"type":"response.create","model":"gpt"}"#),
Err(LiveProtocolError::ExpectedSessionUpdate)
);
assert_eq!(
validate_initial_session_update(r#"["session.update"]"#),
Err(LiveProtocolError::InvalidInitialEvent)
);
assert_eq!(
validate_initial_session_update("not-json"),
Err(LiveProtocolError::InvalidInitialEvent)
);
}
#[test]
fn direct_session_update_uses_the_bounded_session_limit() {
let oversized = format!(
r#"{{"type":"session.update","session":{{"future":"{}"}}}}"#,
"x".repeat(MAX_SESSION_BYTES)
);
assert_eq!(
validate_initial_session_update(oversized.as_str()),
Err(LiveProtocolError::SessionTooLarge)
);
}
#[test]
fn multipart_round_trip_preserves_unknown_session_fields() {
let session = json!({
"model": "gpt-future-live",
"instructions": "opaque",
"future_capability": {"enabled": true},
"audio": {"input": {"format": "pcm16"}}
});
let (content_type, body) = build_live_multipart("v=0\r\no=test", &session);
let parsed = parse_live_multipart(content_type.as_str(), body.as_slice()).unwrap();
assert_eq!(parsed.sdp, "v=0\r\no=test");
assert_eq!(parsed.session, session);
}
#[test]
fn multipart_rejects_duplicate_or_unknown_parts() {
let boundary = "test-boundary";
let duplicate = format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"sdp\"\r\n\r\nv=0\r\n--{boundary}\r\nContent-Disposition: form-data; name=\"sdp\"\r\n\r\nv=1\r\n--{boundary}--\r\n"
);
assert_eq!(
parse_live_multipart(
format!("multipart/form-data; boundary={boundary}").as_str(),
duplicate.as_bytes(),
),
Err(LiveProtocolError::DuplicateMultipartPart)
);
let unknown = format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"sdp\"\r\n\r\nv=0\r\n--{boundary}\r\nContent-Disposition: form-data; name=\"credentials\"\r\n\r\nsecret\r\n--{boundary}--\r\n"
);
assert_eq!(
parse_live_multipart(
format!("multipart/form-data; boundary={boundary}").as_str(),
unknown.as_bytes(),
),
Err(LiveProtocolError::UnexpectedMultipartPart)
);
}
#[test]
fn multipart_enforces_total_sdp_and_session_limits() {
let oversized_body = vec![b'x'; MAX_MULTIPART_BODY_BYTES + 1];
assert_eq!(
parse_live_multipart(
"multipart/form-data; boundary=limit",
oversized_body.as_slice(),
),
Err(LiveProtocolError::MultipartBodyTooLarge)
);
let oversized_sdp = "x".repeat(MAX_SDP_BYTES + 1);
let sdp_body = format!(
"--limit\r\nContent-Disposition: form-data; name=\"sdp\"\r\n\r\n{oversized_sdp}\r\n--limit\r\nContent-Disposition: form-data; name=\"session\"\r\n\r\n{{}}\r\n--limit--\r\n"
);
assert_eq!(
parse_live_multipart("multipart/form-data; boundary=limit", sdp_body.as_bytes()),
Err(LiveProtocolError::SdpTooLarge)
);
let oversized_session = format!(r#"{{"future":"{}"}}"#, "x".repeat(MAX_SESSION_BYTES));
let session_body = format!(
"--limit\r\nContent-Disposition: form-data; name=\"sdp\"\r\n\r\nv=0\r\n--limit\r\nContent-Disposition: form-data; name=\"session\"\r\n\r\n{oversized_session}\r\n--limit--\r\n"
);
assert_eq!(
parse_live_multipart(
"multipart/form-data; boundary=limit",
session_body.as_bytes(),
),
Err(LiveProtocolError::SessionTooLarge)
);
}
#[test]
fn multipart_rejects_unbounded_or_ambiguous_boundaries() {
let valid_body =
b"--safe\r\nContent-Disposition: form-data; name=\"sdp\"\r\n\r\nv=0\r\n--safe--\r\n";
assert_eq!(
parse_live_multipart("application/json", valid_body),
Err(LiveProtocolError::UnsupportedMediaType)
);
assert_eq!(
parse_live_multipart(
format!(
"multipart/form-data; boundary={}",
"x".repeat(MAX_BOUNDARY_BYTES + 1)
)
.as_str(),
valid_body,
),
Err(LiveProtocolError::InvalidBoundary)
);
assert_eq!(
parse_live_multipart(
"multipart/form-data; boundary=safe; boundary=other",
valid_body,
),
Err(LiveProtocolError::InvalidBoundary)
);
}
#[test]
fn extracts_only_realtime_call_ids_from_location() {
for location in [
"https://api.openai.com/v1/live/rtc_abc-123",
"/v1/live/550e8400-e29b-41d4-a716-446655440000",
] {
assert!(extract_call_id_from_location(location).is_ok());
}
assert_eq!(
extract_call_id_from_location("/v1/live/rtc%2Fescape"),
Err(LiveProtocolError::InvalidCallLocation)
);
for location in ["/v1/live", "/v1/live/not-a-call-id"] {
assert_eq!(
extract_call_id_from_location(location),
Err(LiveProtocolError::InvalidCallLocation)
);
}
for dot_segment in [".", ".."] {
assert_eq!(
validate_call_id(dot_segment),
Err(LiveProtocolError::InvalidCallId)
);
}
}
#[test]
fn opaque_event_discriminator_does_not_project_unknown_fields() {
let raw = r#"{"type":"delegation.created","unknown":{"nested":[1,2,3]}}"#;
assert_eq!(event_type(raw).as_deref(), Some("delegation.created"));
assert_eq!(
serde_json::from_str::<Value>(raw).unwrap()["unknown"]["nested"][2],
3
);
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -7,6 +7,7 @@
//! decisions.
pub(crate) mod ingress;
pub(crate) mod live;
pub(crate) mod responses;
pub(crate) mod session;
pub(crate) mod transport;
@@ -134,6 +134,7 @@ const STANDARD_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes =
UpstreamWebSocketErrorCodes {
upstream_url_missing: "responses_upstream_url_missing",
upstream_url_invalid: "responses_upstream_url_invalid",
frontdoor_self_loop: "responses_websocket_frontdoor_self_loop",
headers_invalid: "responses_websocket_headers_invalid",
client_build_failed: "responses_websocket_client_build_failed",
proxy_invalid: "responses_websocket_proxy_invalid",
@@ -215,6 +216,10 @@ mod tests {
adapter.upstream_errors().handshake_failed,
"responses_websocket_handshake_failed"
);
assert_eq!(
adapter.upstream_errors().frontdoor_self_loop,
"responses_websocket_frontdoor_self_loop"
);
}
#[test]
@@ -24,6 +24,7 @@ const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_
const CODEX_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes = UpstreamWebSocketErrorCodes {
upstream_url_missing: "codex_upstream_url_missing",
upstream_url_invalid: "codex_upstream_url_invalid",
frontdoor_self_loop: "codex_websocket_frontdoor_self_loop",
headers_invalid: "codex_websocket_headers_invalid",
client_build_failed: "codex_websocket_client_build_failed",
proxy_invalid: "codex_websocket_proxy_invalid",
@@ -259,6 +260,16 @@ mod tests {
ResponsesWebSocketRebindSafety, ResponsesWebSocketRelayDirective,
};
#[test]
fn codex_adapter_has_a_distinct_frontdoor_self_loop_error() {
let adapter = CodexResponsesWebSocketAdapter;
assert_eq!(
adapter.upstream_errors().frontdoor_self_loop,
"codex_websocket_frontdoor_self_loop"
);
}
#[test]
fn codex_rate_limit_chunk_is_kept_for_the_terminal_report() {
let adapter = CodexResponsesWebSocketAdapter;
@@ -16,7 +16,7 @@ use crate::provider_pool_demand::{
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
use crate::{AppState, GatewayError};
pub(super) struct ResponsesWebSocketTurnAdmission {
pub(crate) struct ResponsesWebSocketTurnAdmission {
upstream_execution: Option<aether_runtime::ConcurrencyPermit>,
upstream_target: Option<UpstreamTargetAdmissionPermit>,
provider_pool: Option<ProviderPoolInFlightGuard>,
@@ -24,7 +24,7 @@ pub(super) struct ResponsesWebSocketTurnAdmission {
}
impl ResponsesWebSocketTurnAdmission {
pub(super) async fn acquire(
pub(crate) async fn acquire(
state: &AppState,
plan: &ExecutionPlan,
trace_id: &str,
@@ -61,7 +61,7 @@ impl ResponsesWebSocketTurnAdmission {
/// Release the distributed provider token before the turn's persistence
/// work. The remaining permits are local RAII guards and are dropped with
/// this value.
pub(super) async fn release(mut self) {
pub(crate) async fn release(mut self) {
if let Some(provider_pool) = self.provider_pool.take() {
provider_pool.release().await;
}
@@ -29,6 +29,8 @@ mod turn;
mod turn_state;
mod upstream;
pub(crate) use admission::ResponsesWebSocketTurnAdmission;
use std::net::SocketAddr;
use axum::body::Body;
@@ -20,6 +20,13 @@ pub(crate) const RESPONSES_WEBSOCKET_SESSION_LIMITS: WebSocketSessionLimits =
max_connection_duration: Duration::from_secs(60 * 60),
};
pub(crate) const LIVE_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
@@ -22,6 +22,7 @@ use crate::ai_serving::AiExecutionDecision;
use crate::execution_runtime::transport::{
build_browser_wreq_client, build_request_headers, ExecutionTransportControls,
};
use crate::frontdoor_loop_guard::gateway_frontdoor_self_loop_guard_error;
use crate::handlers::proxy::websocket::session::{
WebSocketSessionLimits, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
};
@@ -30,6 +31,7 @@ use crate::handlers::proxy::websocket::session::{
pub(crate) struct UpstreamWebSocketErrorCodes {
pub(crate) upstream_url_missing: &'static str,
pub(crate) upstream_url_invalid: &'static str,
pub(crate) frontdoor_self_loop: &'static str,
pub(crate) headers_invalid: &'static str,
pub(crate) client_build_failed: &'static str,
pub(crate) proxy_invalid: &'static str,
@@ -53,7 +55,11 @@ pub(crate) async fn connect_upstream_websocket(
.upstream_url
.as_deref()
.ok_or(errors.upstream_url_missing)?;
let upstream_url = websocket_upstream_url(upstream_url, errors.upstream_url_invalid)?;
let upstream_url = guarded_websocket_upstream_url(
upstream_url,
errors.upstream_url_invalid,
errors.frontdoor_self_loop,
)?;
let headers =
websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?;
let client = build_websocket_client(decision, errors)?;
@@ -79,6 +85,18 @@ pub(crate) async fn connect_upstream_websocket(
})
}
fn guarded_websocket_upstream_url(
raw: &str,
invalid_code: &'static str,
frontdoor_self_loop_code: &'static str,
) -> Result<Url, &'static str> {
let upstream_url = websocket_upstream_url(raw, invalid_code)?;
if gateway_frontdoor_self_loop_guard_error(upstream_url.as_str()).is_some() {
return Err(frontdoor_self_loop_code);
}
Ok(upstream_url)
}
fn websocket_response_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
headers
.iter()
@@ -310,6 +328,19 @@ pub(crate) fn upstream_message_to_client(message: WreqWsMessage) -> AxumWsMessag
}
}
pub(crate) fn client_message_to_upstream(message: AxumWsMessage) -> WreqWsMessage {
match message {
AxumWsMessage::Text(text) => WreqWsMessage::Text(text.to_string().into()),
AxumWsMessage::Binary(data) => WreqWsMessage::Binary(data),
AxumWsMessage::Ping(data) => WreqWsMessage::Ping(data),
AxumWsMessage::Pong(data) => WreqWsMessage::Pong(data),
AxumWsMessage::Close(frame) => WreqWsMessage::Close(frame.map(|frame| WreqCloseFrame {
code: frame.code.into(),
reason: frame.reason.to_string().into(),
})),
}
}
/// Builds a Responses WebSocket error event in the shape understood by the
/// official client implementations. The status is part of the event body,
/// not the WebSocket handshake, because the connection is already upgraded.
@@ -467,11 +498,12 @@ pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16
#[cfg(test)]
mod tests {
use super::{
bounded_send, responses_websocket_error_event,
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,
};
use crate::frontdoor_loop_guard::configured_gateway_frontdoor_base_url;
use axum::http::HeaderMap;
use std::collections::BTreeMap;
use std::time::Duration;
@@ -587,6 +619,39 @@ mod tests {
assert!(websocket_upstream_url("https://[email protected]/responses", "invalid").is_err());
}
#[test]
fn rejects_responses_websocket_frontdoor_self_loop_before_connecting() {
let base_url = configured_gateway_frontdoor_base_url();
let raw_url = format!("{base_url}/v1/responses");
assert_eq!(
guarded_websocket_upstream_url(
raw_url.as_str(),
"responses_upstream_url_invalid",
"responses_websocket_frontdoor_self_loop",
),
Err("responses_websocket_frontdoor_self_loop")
);
}
#[test]
fn rejects_live_direct_and_sideband_frontdoor_self_loops_before_connecting() {
let base_url = configured_gateway_frontdoor_base_url();
for path in ["/v1/live", "/v1/live/rtc_test"] {
let raw_url = format!("{base_url}{path}");
assert_eq!(
guarded_websocket_upstream_url(
raw_url.as_str(),
"codex_live_upstream_url_invalid",
"codex_live_websocket_frontdoor_self_loop",
),
Err("codex_live_websocket_frontdoor_self_loop"),
"{path} must be rejected before an upstream handshake"
);
}
}
#[test]
fn upstream_handshake_keeps_provider_auth_but_drops_transport_managed_headers() {
let provider_headers = BTreeMap::from([