diff --git a/crates/protocols/src/realtime_session.rs b/crates/protocols/src/realtime_session.rs index decb0826b0..78c5be9172 100644 --- a/crates/protocols/src/realtime_session.rs +++ b/crates/protocols/src/realtime_session.rs @@ -5,15 +5,20 @@ use std::collections::HashMap; use serde::{Deserialize, Serialize}; use serde_json::Value; +use validator::{Validate, ValidationError}; -use crate::common::{Redacted, ResponsePrompt, ToolReference}; +use crate::{ + common::{Redacted, ResponsePrompt, ToolReference}, + validated::Normalizable, +}; // ============================================================================ // Session Configuration // ============================================================================ #[serde_with::skip_serializing_none] -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Validate)] +#[validate(schema(function = "validate_session_create_request"))] pub struct RealtimeSessionCreateRequest { #[serde(rename = "type")] pub r#type: RealtimeSessionType, @@ -31,6 +36,18 @@ pub struct RealtimeSessionCreateRequest { pub truncation: Option, } +impl Normalizable for RealtimeSessionCreateRequest {} + +fn validate_session_create_request( + req: &RealtimeSessionCreateRequest, +) -> Result<(), ValidationError> { + let has_model = req.model.as_deref().is_some_and(|m| !m.trim().is_empty()); + if !has_model { + return Err(ValidationError::new("model is required")); + } + Ok(()) +} + // ============================================================================ // Session Object // ============================================================================ @@ -60,12 +77,30 @@ pub struct RealtimeSessionCreateResponse { // ============================================================================ #[serde_with::skip_serializing_none] -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Validate)] +#[validate(schema(function = "validate_transcription_session_create_request"))] pub struct RealtimeTranscriptionSessionCreateRequest { #[serde(rename = "type")] pub r#type: RealtimeTranscriptionSessionType, pub audio: Option, pub include: Option>, + pub model: Option, + pub language: Option, + pub prompt: Option, +} + +impl Normalizable for RealtimeTranscriptionSessionCreateRequest { + // Use default no-op implementation +} + +fn validate_transcription_session_create_request( + req: &RealtimeTranscriptionSessionCreateRequest, +) -> Result<(), ValidationError> { + let has_model = req.model.as_deref().is_some_and(|m| !m.trim().is_empty()); + if !has_model { + return Err(ValidationError::new("model is required")); + } + Ok(()) } // ============================================================================ @@ -446,6 +481,29 @@ pub enum RealtimeTruncation { // Client Secret // ============================================================================ +#[serde_with::skip_serializing_none] +#[derive(Debug, Clone, Serialize, Deserialize, Validate)] +#[validate(schema(function = "validate_client_secret_create_request"))] +pub struct RealtimeClientSecretCreateRequest { + pub session: RealtimeSessionCreateRequest, +} + +impl Normalizable for RealtimeClientSecretCreateRequest {} + +fn validate_client_secret_create_request( + req: &RealtimeClientSecretCreateRequest, +) -> Result<(), ValidationError> { + let has_model = req + .session + .model + .as_deref() + .is_some_and(|m| !m.trim().is_empty()); + if !has_model { + return Err(ValidationError::new("session.model is required")); + } + Ok(()) +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RealtimeSessionClientSecret { pub expires_at: i64, diff --git a/model_gateway/src/observability/metrics.rs b/model_gateway/src/observability/metrics.rs index ea49044832..057c606346 100644 --- a/model_gateway/src/observability/metrics.rs +++ b/model_gateway/src/observability/metrics.rs @@ -389,6 +389,13 @@ pub mod metrics_labels { pub const ENDPOINT_RERANK: &str = "rerank"; pub const ENDPOINT_EMBEDDINGS: &str = "embeddings"; pub const ENDPOINT_CLASSIFY: &str = "classify"; + pub const ENDPOINT_REALTIME: &str = "realtime"; + pub const ENDPOINT_REALTIME_SESSIONS: &str = "realtime_sessions"; + pub const ENDPOINT_REALTIME_CLIENT_SECRETS: &str = "realtime_client_secrets"; + pub const ENDPOINT_REALTIME_TRANSCRIPTION: &str = "realtime_transcription"; + + // Connection modes + pub const CONNECTION_WEBSOCKET: &str = "websocket"; // Worker types pub const WORKER_REGULAR: &str = "regular"; diff --git a/model_gateway/src/routers/mod.rs b/model_gateway/src/routers/mod.rs index 559197912d..5df092e1b6 100644 --- a/model_gateway/src/routers/mod.rs +++ b/model_gateway/src/routers/mod.rs @@ -17,6 +17,10 @@ use openai_protocol::{ generate::GenerateRequest, interactions::InteractionsRequest, messages::CreateMessageRequest, + realtime_session::{ + RealtimeClientSecretCreateRequest, RealtimeSessionCreateRequest, + RealtimeTranscriptionSessionCreateRequest, + }, rerank::RerankRequest, responses::{ResponsesGetParams, ResponsesRequest}, }; @@ -233,6 +237,54 @@ pub trait RouterTrait: Send + Sync + Debug { .into_response() } + /// Route a realtime session create request (/v1/realtime/sessions) + async fn route_realtime_session( + &self, + _headers: Option<&HeaderMap>, + _body: &RealtimeSessionCreateRequest, + ) -> Response { + ( + StatusCode::NOT_IMPLEMENTED, + "Realtime sessions not implemented", + ) + .into_response() + } + + /// Route a realtime client secret create request (/v1/realtime/client_secrets) + async fn route_realtime_client_secret( + &self, + _headers: Option<&HeaderMap>, + _body: &RealtimeClientSecretCreateRequest, + ) -> Response { + ( + StatusCode::NOT_IMPLEMENTED, + "Realtime client secrets not implemented", + ) + .into_response() + } + + /// Route a realtime transcription session create request (/v1/realtime/transcription_sessions) + async fn route_realtime_transcription_session( + &self, + _headers: Option<&HeaderMap>, + _body: &RealtimeTranscriptionSessionCreateRequest, + ) -> Response { + ( + StatusCode::NOT_IMPLEMENTED, + "Realtime transcription sessions not implemented", + ) + .into_response() + } + + /// Route a realtime WebSocket upgrade request + async fn route_realtime_ws(&self, _req: Request, _model_id: &str) -> Response { + ( + StatusCode::NOT_IMPLEMENTED, + "Realtime WebSocket not implemented", + ) + .into_response() + } + /// Get router type name fn router_type(&self) -> &'static str; diff --git a/model_gateway/src/routers/openai/realtime/rest.rs b/model_gateway/src/routers/openai/realtime/rest.rs index 562e60dac8..a4fd537bc2 100644 --- a/model_gateway/src/routers/openai/realtime/rest.rs +++ b/model_gateway/src/routers/openai/realtime/rest.rs @@ -1,127 +1,85 @@ -//! REST handlers for Realtime API token generation endpoints. -//! -//! These endpoints generate ephemeral tokens for browser-safe authentication. -//! They do NOT create sessions — sessions are created implicitly when the -//! client connects (WebSocket or WebRTC). +//! Shared helpers for Realtime API REST proxy responses. -use std::sync::Arc; +use std::{sync::Arc, time::Instant}; use axum::{ - extract::State, http::{HeaderMap, StatusCode}, response::{IntoResponse, Response}, - Json, }; -use serde_json::Value; -use tracing::{debug, error}; +use tracing::error; use crate::{ - core::worker::{RuntimeType, Worker, WorkerLoadGuard}, - routers::{error, header_utils::extract_auth_header}, - server::AppState, + core::{worker::WorkerLoadGuard, Worker}, + observability::metrics::{metrics_labels, Metrics}, + routers::header_utils::extract_auth_header, }; -/// `POST /v1/realtime/client_secrets` — GA ephemeral token generation. +/// Forward a realtime REST request to the upstream worker. /// -/// Generates a short-lived token for browser-safe auth. The session config -/// (model, voice, tools) is included in the request body, pre-configuring -/// the session. MCP tool definitions are injected before forwarding. -pub async fn create_client_secret( - State(state): State>, - headers: HeaderMap, - Json(body): Json, -) -> Response { - let model = match body.pointer("/session/model").and_then(|v| v.as_str()) { - Some(m) => m.to_string(), - None => return error::bad_request("missing_model", "session.model is required"), - }; - - // TODO(Phase 3): Inject MCP tool definitions into body.session.tools - - proxy_realtime_rest( - &state, - &headers, - &body, - &model, - "/v1/realtime/client_secrets", - ) - .await -} - -/// `POST /v1/realtime/sessions` — Legacy ephemeral token generation. -/// -/// Kept for backwards compatibility; prefer `create_client_secret` for new integrations. -pub async fn create_session( - State(state): State>, - headers: HeaderMap, - Json(body): Json, -) -> Response { - let model = match body.get("model").and_then(|v| v.as_str()) { - Some(m) => m.to_string(), - None => return error::bad_request("missing_model", "model is required"), - }; - - // TODO(Phase 3): Inject MCP tool definitions into body.tools - - proxy_realtime_rest(&state, &headers, &body, &model, "/v1/realtime/sessions").await -} - -/// `POST /v1/realtime/transcription_sessions` — Legacy ephemeral token generation for transcription. -/// -/// Kept for backwards compatibility; prefer `create_client_secret` for new integrations. -pub async fn create_transcription_session( - State(state): State>, - headers: HeaderMap, - Json(body): Json, -) -> Response { - let model = match body.get("model").and_then(|v| v.as_str()) { - Some(m) => m.to_string(), - None => return error::bad_request("missing_model", "model is required"), - }; - - // TODO(Phase 3): Inject MCP tool definitions into config - - proxy_realtime_rest( - &state, - &headers, - &body, - &model, - "/v1/realtime/transcription_sessions", - ) - .await -} - -/// Shared proxy logic for all realtime REST endpoints. +/// Shared logic for sessions, client_secrets, and transcription_sessions: +/// auth, load tracking, metrics, and proxy. /// -/// Handles worker selection, load tracking, auth, header forwarding, -/// upstream request, circuit breaker outcome recording, and response proxying. -async fn proxy_realtime_rest( - state: &AppState, - headers: &HeaderMap, - body: &Value, +/// The caller is responsible for worker selection; this function receives +/// the pre-selected worker (or an error response). +pub(crate) async fn forward_realtime_rest( + client: &reqwest::Client, + worker: Result, Response>, + headers: Option<&HeaderMap>, + body: &(impl serde::Serialize + Sync), model: &str, - path: &str, + endpoint: &str, + endpoint_label: &'static str, ) -> Response { - let worker = match select_worker(state, model) { - Some(w) => w, - None => { - error!(model, path, "No available worker for realtime model"); - return StatusCode::SERVICE_UNAVAILABLE.into_response(); + let start = Instant::now(); + + Metrics::record_router_request( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_HTTP, + model, + endpoint_label, + "false", + ); + + let worker = match worker { + Ok(w) => w, + Err(response) => { + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_HTTP, + model, + endpoint_label, + metrics_labels::ERROR_NO_WORKERS, + ); + return response; } }; - let auth = match extract_auth_header(Some(headers), worker.api_key()) { + let auth = match extract_auth_header(headers, worker.api_key()) { Some(v) => v, - None => return StatusCode::UNAUTHORIZED.into_response(), + None => { + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_HTTP, + model, + endpoint_label, + metrics_labels::ERROR_VALIDATION, + ); + return StatusCode::UNAUTHORIZED.into_response(); + } }; // Track load for the duration of the upstream request - let _guard = WorkerLoadGuard::new(worker.clone(), Some(headers)); + let _guard = WorkerLoadGuard::new(worker.clone(), headers); - let upstream_url = format!("{}{path}", worker.url().trim_end_matches('/')); - debug!(model, upstream_url, "Forwarding realtime REST request"); + let upstream_url = format!("{}{endpoint}", worker.url().trim_end_matches('/')); - let result = forward_post(&state.context.client, &upstream_url, &auth, body) + let result = client + .post(&upstream_url) + .header("Authorization", &auth) + .json(body) .send() .await; @@ -130,47 +88,46 @@ async fn proxy_realtime_rest( let success = resp.status().is_success(); let response = proxy_response(resp).await; worker.record_outcome(success); + if success { + Metrics::record_router_duration( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_HTTP, + model, + endpoint_label, + start.elapsed(), + ); + } else { + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_HTTP, + model, + endpoint_label, + metrics_labels::ERROR_BACKEND, + ); + } response } Err(e) => { - error!(error = %e, path, "Failed to forward realtime REST request"); + error!(error = %e, endpoint, "Failed to forward realtime REST request"); worker.record_outcome(false); + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_HTTP, + model, + endpoint_label, + metrics_labels::ERROR_BACKEND, + ); StatusCode::BAD_GATEWAY.into_response() } } } -/// Build a POST request to upstream with the caller's auth token. -fn forward_post( - client: &reqwest::Client, - upstream_url: &str, - auth: &http::HeaderValue, - body: &Value, -) -> reqwest::RequestBuilder { - client - .post(upstream_url) - .header("Authorization", auth) - .json(body) -} - -/// Select the best available worker for a realtime model. -/// -/// Uses `supports_model()` which checks `models_override` (populated by lazy -/// discovery), unlike `get_by_model()` which relies on `model_index` (not -/// populated for lazily-discovered workers). -pub(super) fn select_worker(state: &AppState, model: &str) -> Option> { - state - .context - .worker_registry - .get_workers_filtered(None, None, None, Some(RuntimeType::External), true) - .into_iter() - .filter(|w| w.supports_model(model) && w.circuit_breaker().can_execute()) - .min_by_key(|w| w.load()) -} - /// Convert an upstream reqwest Response into an axum Response, /// preserving status code and body. -pub(super) async fn proxy_response(resp: reqwest::Response) -> Response { +pub(crate) async fn proxy_response(resp: reqwest::Response) -> Response { let status = StatusCode::from_u16(resp.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY); let content_type = resp .headers() diff --git a/model_gateway/src/routers/openai/realtime/ws.rs b/model_gateway/src/routers/openai/realtime/ws.rs index 597ff403c0..833cb28f2b 100644 --- a/model_gateway/src/routers/openai/realtime/ws.rs +++ b/model_gateway/src/routers/openai/realtime/ws.rs @@ -1,112 +1,137 @@ -//! WebSocket transport handler for `/v1/realtime`. -//! -//! Server-to-server transport. Client connects with an API key, -//! session is created implicitly when the connection is established. +//! WebSocket transport helpers for `/v1/realtime`. use std::sync::Arc; use axum::{ extract::{ ws::{WebSocket, WebSocketUpgrade}, - Query, State, + FromRequestParts, }, - http::HeaderMap, + http::{request::Parts, HeaderValue, StatusCode}, response::{IntoResponse, Response}, }; use serde::Deserialize; -use tracing::{debug, error, info, warn}; -use super::{proxy, rest::select_worker}; -use crate::{routers::header_utils::extract_auth_header, server::AppState}; +use super::{proxy::run_ws_proxy, RealtimeRegistry}; +use crate::{ + core::Worker, + observability::metrics::{metrics_labels, Metrics}, + routers::header_utils::extract_auth_header, +}; #[derive(Debug, Deserialize)] pub struct RealtimeQueryParams { pub model: Option, } -/// Handler for `GET /v1/realtime` — WebSocket upgrade. +/// Handle a realtime WebSocket upgrade request. /// -/// 1. Extract model from query params, auth from headers -/// 2. Select worker via WorkerRegistry (model routing) -/// 3. Upgrade to WebSocket -/// 4. Delegate to bidirectional WS proxy -pub async fn ws_handler( - State(state): State>, - Query(params): Query, - headers: HeaderMap, - ws: WebSocketUpgrade, +/// The caller is responsible for extracting the model from the query string +/// and selecting a worker. This function handles WS upgrade, auth resolution, +/// and proxy spawning. +pub(crate) async fn handle_realtime_ws( + mut parts: Parts, + model: String, + worker: Result, Response>, + auth_header: Option, + realtime_registry: Arc, ) -> Response { - let model = match params.model { - Some(m) => m, - None => { - warn!("Missing required 'model' query parameter"); - return ( - http::StatusCode::BAD_REQUEST, - "Missing required 'model' query parameter", - ) - .into_response(); - } - }; - - let worker = match select_worker(&state, &model) { - Some(w) => w, - None => { - warn!(model, "No available worker for realtime model"); - return http::StatusCode::SERVICE_UNAVAILABLE.into_response(); + let worker = match worker { + Ok(w) => w, + Err(response) => { + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_WEBSOCKET, + &model, + metrics_labels::ENDPOINT_REALTIME, + metrics_labels::ERROR_NO_WORKERS, + ); + return response; } }; let worker_url = worker.url().to_string(); - let upstream_ws_url = build_upstream_ws_url(&worker_url, &model); - let auth_header_value = extract_auth_header(Some(&headers), worker.api_key()); - let auth_str = match auth_header_value { + // Use user auth if available, fall back to worker's API key + let effective_auth = auth_header.or_else(|| extract_auth_header(None, worker.api_key())); + let auth_str = match effective_auth { Some(v) => match v.to_str() { Ok(s) => s.to_string(), Err(_) => { - warn!("Authorization header contains invalid UTF-8 characters"); + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_WEBSOCKET, + &model, + metrics_labels::ENDPOINT_REALTIME, + metrics_labels::ERROR_VALIDATION, + ); return ( - http::StatusCode::BAD_REQUEST, + StatusCode::BAD_REQUEST, "Authorization header contains invalid UTF-8 characters", ) .into_response(); } }, None => { - error!("No authorization available for upstream realtime connection"); - return http::StatusCode::UNAUTHORIZED.into_response(); + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_WEBSOCKET, + &model, + metrics_labels::ENDPOINT_REALTIME, + metrics_labels::ERROR_VALIDATION, + ); + return StatusCode::UNAUTHORIZED.into_response(); } }; - let registry = Arc::clone(&state.context.realtime_registry); + let ws = match WebSocketUpgrade::from_request_parts(&mut parts, &()).await { + Ok(ws) => ws, + Err(e) => { + Metrics::record_router_error( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_WEBSOCKET, + &model, + metrics_labels::ENDPOINT_REALTIME, + metrics_labels::ERROR_VALIDATION, + ); + return e.into_response(); + } + }; let session_id = uuid::Uuid::now_v7().to_string(); - let entry = registry.register_session(session_id.clone(), model.clone(), worker_url.clone()); + let entry = + realtime_registry.register_session(session_id.clone(), model.clone(), worker_url.clone()); let cancel_token = entry.cancel_token.clone(); - info!( + tracing::info!( session_id, - model, worker_url, "Upgrading to realtime WebSocket" + model, + worker_url, + "Upgrading to realtime WebSocket" ); ws.on_upgrade(move |socket: WebSocket| async move { - if let Err(e) = proxy::run_ws_proxy( + if let Err(e) = run_ws_proxy( socket, &upstream_ws_url, &auth_str, - registry.clone(), + realtime_registry.clone(), session_id.clone(), cancel_token, ) .await { - error!(session_id, error = %e, "Realtime WebSocket proxy error"); + tracing::error!(session_id, error = %e, "Realtime WebSocket proxy error"); } // Cleanup: remove session on disconnect - registry.remove_session(&session_id); - debug!(session_id, "Realtime session cleaned up"); + realtime_registry.remove_session(&session_id); + tracing::debug!(session_id, "Realtime session cleaned up"); }) } @@ -114,7 +139,7 @@ pub async fn ws_handler( /// /// Worker URLs use `http(s)://` but tungstenite requires `ws(s)://`, /// e.g. `https://api.openai.com` → `wss://api.openai.com/v1/realtime?model=gpt-4o-realtime-preview` -fn build_upstream_ws_url(worker_url: &str, model: &str) -> String { +pub(crate) fn build_upstream_ws_url(worker_url: &str, model: &str) -> String { let base = worker_url.trim_end_matches('/'); let ws_base = if let Some(rest) = base.strip_prefix("https://") { format!("wss://{rest}") diff --git a/model_gateway/src/routers/openai/router.rs b/model_gateway/src/routers/openai/router.rs index c4fc8afbbf..bc6bbb2ad3 100644 --- a/model_gateway/src/routers/openai/router.rs +++ b/model_gateway/src/routers/openai/router.rs @@ -15,6 +15,10 @@ use axum::{ use futures_util::{future::join_all, StreamExt}; use openai_protocol::{ chat::ChatCompletionRequest, + realtime_session::{ + RealtimeClientSecretCreateRequest, RealtimeSessionCreateRequest, + RealtimeTranscriptionSessionCreateRequest, + }, responses::{ generate_id, ResponseContentPart, ResponseInput, ResponseInputOutputItem, ResponsesGetParams, ResponsesRequest, @@ -42,7 +46,10 @@ use crate::{ WorkerRegistry, }, observability::metrics::{bool_to_static_str, metrics_labels, Metrics}, - routers::header_utils::{apply_provider_headers, extract_auth_header}, + routers::{ + header_utils::{apply_provider_headers, extract_auth_header}, + openai::realtime::{rest::forward_realtime_rest, ws::handle_realtime_ws, RealtimeRegistry}, + }, }; pub struct OpenAIRouter { @@ -52,6 +59,7 @@ pub struct OpenAIRouter { shared_components: Arc, responses_components: Arc, retry_config: RetryConfig, + realtime_registry: Arc, } impl std::fmt::Debug for OpenAIRouter { @@ -121,6 +129,7 @@ impl OpenAIRouter { shared_components, responses_components, retry_config: ctx.router_config.effective_retry_config(), + realtime_registry: ctx.realtime_registry.clone(), }) } @@ -1068,6 +1077,95 @@ impl crate::routers::RouterTrait for OpenAIRouter { } } + async fn route_realtime_session( + &self, + headers: Option<&HeaderMap>, + body: &RealtimeSessionCreateRequest, + ) -> Response { + // TODO(Phase 3): Inject MCP tool definitions into body.tools + let model = body.model.as_deref().unwrap_or_default(); + let auth = extract_auth_header(headers, None); + let worker = self.select_worker_for_model(model, auth.as_ref()).await; + forward_realtime_rest( + &self.shared_components.client, + worker, + headers, + body, + model, + "/v1/realtime/sessions", + metrics_labels::ENDPOINT_REALTIME_SESSIONS, + ) + .await + } + + async fn route_realtime_client_secret( + &self, + headers: Option<&HeaderMap>, + body: &RealtimeClientSecretCreateRequest, + ) -> Response { + // TODO(Phase 3): Inject MCP tool definitions into body.session.tools + let model = body.session.model.as_deref().unwrap_or_default(); + let auth = extract_auth_header(headers, None); + let worker = self.select_worker_for_model(model, auth.as_ref()).await; + forward_realtime_rest( + &self.shared_components.client, + worker, + headers, + body, + model, + "/v1/realtime/client_secrets", + metrics_labels::ENDPOINT_REALTIME_CLIENT_SECRETS, + ) + .await + } + + async fn route_realtime_transcription_session( + &self, + headers: Option<&HeaderMap>, + body: &RealtimeTranscriptionSessionCreateRequest, + ) -> Response { + let model = body.model.as_deref().unwrap_or_default(); + let auth = extract_auth_header(headers, None); + let worker = self.select_worker_for_model(model, auth.as_ref()).await; + forward_realtime_rest( + &self.shared_components.client, + worker, + headers, + body, + model, + "/v1/realtime/transcription_sessions", + metrics_labels::ENDPOINT_REALTIME_TRANSCRIPTION, + ) + .await + } + + async fn route_realtime_ws(&self, req: Request, model: &str) -> Response { + let (parts, _body) = req.into_parts(); + + Metrics::record_router_request( + metrics_labels::ROUTER_OPENAI, + metrics_labels::BACKEND_EXTERNAL, + metrics_labels::CONNECTION_WEBSOCKET, + model, + metrics_labels::ENDPOINT_REALTIME, + "false", + ); + + let auth_header = extract_auth_header(Some(&parts.headers), None); + let worker = self + .select_worker_for_model(model, auth_header.as_ref()) + .await; + + handle_realtime_ws( + parts, + model.to_owned(), + worker, + auth_header, + Arc::clone(&self.realtime_registry), + ) + .await + } + fn router_type(&self) -> &'static str { "openai" } diff --git a/model_gateway/src/routers/router_manager.rs b/model_gateway/src/routers/router_manager.rs index 0453248ac4..dc257d5179 100644 --- a/model_gateway/src/routers/router_manager.rs +++ b/model_gateway/src/routers/router_manager.rs @@ -23,6 +23,10 @@ use openai_protocol::{ generate::GenerateRequest, interactions::InteractionsRequest, messages::CreateMessageRequest, + realtime_session::{ + RealtimeClientSecretCreateRequest, RealtimeSessionCreateRequest, + RealtimeTranscriptionSessionCreateRequest, + }, rerank::RerankRequest, responses::{ResponsesGetParams, ResponsesRequest}, }; @@ -748,6 +752,75 @@ impl RouterTrait for RouterManager { } } + async fn route_realtime_session( + &self, + headers: Option<&HeaderMap>, + body: &RealtimeSessionCreateRequest, + ) -> Response { + let model = body.model.as_deref(); + let router = self.select_router_for_request(headers, model); + if let Some(router) = router { + router.route_realtime_session(headers, body).await + } else { + ( + StatusCode::NOT_FOUND, + "No router available for realtime session request", + ) + .into_response() + } + } + + async fn route_realtime_client_secret( + &self, + headers: Option<&HeaderMap>, + body: &RealtimeClientSecretCreateRequest, + ) -> Response { + let model = body.session.model.as_deref(); + let router = self.select_router_for_request(headers, model); + if let Some(router) = router { + router.route_realtime_client_secret(headers, body).await + } else { + ( + StatusCode::NOT_FOUND, + "No router available for realtime client secret request", + ) + .into_response() + } + } + + async fn route_realtime_transcription_session( + &self, + headers: Option<&HeaderMap>, + body: &RealtimeTranscriptionSessionCreateRequest, + ) -> Response { + let model = body.model.as_deref(); + let router = self.select_router_for_request(headers, model); + if let Some(router) = router { + router + .route_realtime_transcription_session(headers, body) + .await + } else { + ( + StatusCode::NOT_FOUND, + "No router available for realtime transcription request", + ) + .into_response() + } + } + + async fn route_realtime_ws(&self, req: Request, model: &str) -> Response { + let router = self.select_router_for_request(None, Some(model)); + if let Some(router) = router { + router.route_realtime_ws(req, model).await + } else { + ( + StatusCode::NOT_FOUND, + "No router available for realtime WebSocket request", + ) + .into_response() + } + } + fn router_type(&self) -> &'static str { "manager" } diff --git a/model_gateway/src/server.rs b/model_gateway/src/server.rs index 6c5c1f4ce0..ee5832aae9 100644 --- a/model_gateway/src/server.rs +++ b/model_gateway/src/server.rs @@ -23,6 +23,10 @@ use openai_protocol::{ interactions::InteractionsRequest, messages::CreateMessageRequest, parser::{ParseFunctionCallRequest, SeparateReasoningRequest}, + realtime_session::{ + RealtimeClientSecretCreateRequest, RealtimeSessionCreateRequest, + RealtimeTranscriptionSessionCreateRequest, + }, rerank::{RerankRequest, V1RerankReqInput}, responses::{ResponsesGetParams, ResponsesRequest}, tokenize::{AddTokenizerRequest, DetokenizeRequest, TokenizeRequest}, @@ -60,7 +64,7 @@ use crate::{ get_mesh_health, get_policy_state, get_policy_states, get_worker_state, get_worker_states, set_global_rate_limit, trigger_graceful_shutdown, update_app_config, }, - openai::realtime::{rest as realtime_rest, ws as realtime_ws}, + openai::realtime::ws::RealtimeQueryParams, parse, router_manager::RouterManager, tokenize, RouterTrait, @@ -434,6 +438,57 @@ async fn v1_conversations_delete_item( .await } +async fn v1_realtime_ws( + State(state): State>, + Query(params): Query, + req: Request, +) -> Response { + let model = match params.model { + Some(m) if !m.trim().is_empty() => m, + _ => { + return ( + StatusCode::BAD_REQUEST, + "Missing required 'model' query parameter", + ) + .into_response(); + } + }; + state.router.route_realtime_ws(req, &model).await +} + +async fn v1_realtime_session( + State(state): State>, + headers: http::HeaderMap, + ValidatedJson(body): ValidatedJson, +) -> Response { + state + .router + .route_realtime_session(Some(&headers), &body) + .await +} + +async fn v1_realtime_client_secret( + State(state): State>, + headers: http::HeaderMap, + ValidatedJson(body): ValidatedJson, +) -> Response { + state + .router + .route_realtime_client_secret(Some(&headers), &body) + .await +} + +async fn v1_realtime_transcription_session( + State(state): State>, + headers: http::HeaderMap, + ValidatedJson(body): ValidatedJson, +) -> Response { + state + .router + .route_realtime_transcription_session(Some(&headers), &body) + .await +} + async fn flush_cache(State(state): State>, _req: Request) -> Response { WorkerManager::flush_cache_all(&state.context.worker_registry, &state.context.client) .await @@ -624,6 +679,16 @@ pub fn build_app( // Tokenize / Detokenize endpoints .route("/v1/tokenize", post(v1_tokenize)) .route("/v1/detokenize", post(v1_detokenize)) + // Realtime REST endpoints (same middleware as other protected routes) + .route("/v1/realtime/sessions", post(v1_realtime_session)) + .route( + "/v1/realtime/client_secrets", + post(v1_realtime_client_secret), + ) + .route( + "/v1/realtime/transcription_sessions", + post(v1_realtime_transcription_session), + ) .route_layer(axum::middleware::from_fn_with_state( app_state.clone(), middleware::concurrency_limit_middleware, @@ -637,17 +702,11 @@ pub fn build_app( middleware::wasm_middleware, )); - let realtime_routes = Router::new() - .route("/v1/realtime", get(realtime_ws::ws_handler)) - .route("/v1/realtime/sessions", post(realtime_rest::create_session)) - .route( - "/v1/realtime/client_secrets", - post(realtime_rest::create_client_secret), - ) - .route( - "/v1/realtime/transcription_sessions", - post(realtime_rest::create_transcription_session), - ) + // WebSocket route: auth + concurrency but NO WASM middleware. + // WASM OnResponse reconstructs the response from status/headers/body, + // dropping the response extensions that carry the WebSocket upgrade future. + let ws_routes = Router::new() + .route("/v1/realtime", get(v1_realtime_ws)) .route_layer(axum::middleware::from_fn_with_state( app_state.clone(), middleware::concurrency_limit_middleware, @@ -736,7 +795,7 @@ pub fn build_app( Router::new() .merge(protected_routes) - .merge(realtime_routes) + .merge(ws_routes) .merge(public_routes) .merge(admin_routes) .merge(worker_routes)