Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ uuid = { version = "1.10", features = ["v7", "serde"] }
wasmtime = { version = "41.0", features = ["component-model", "async"] }
wasmtime-wasi = "41.0"
tokio-util = { version = "0.7", features = ["rt"] }
tokio-tungstenite = { version = "0.26", features = ["rustls-tls-native-roots"] }
scopeguard = "1.2"
bitflags = "2.10.0"
schemars = "0.8"
Expand Down
1 change: 1 addition & 0 deletions model_gateway/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@ bytemuck = { workspace = true, features = ["derive"] }
reqwest = { workspace = true, features = ["stream", "blocking", "json", "rustls-tls"] }
serde = { workspace = true, features = ["derive"] }
tokio = { workspace = true, features = ["full"] }
tokio-util.workspace = true
uuid = { workspace = true, features = ["serde"] }

# Workspace crates
Expand Down
4 changes: 3 additions & 1 deletion model_gateway/src/app_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ use crate::{
middleware::TokenBucket,
observability::inflight_tracker::InFlightRequestTracker,
policies::PolicyRegistry,
routers::router_manager::RouterManager,
routers::{openai::realtime::RealtimeRegistry, router_manager::RouterManager},
wasm::{config::WasmRuntimeConfig, module_manager::WasmModuleManager},
};

Expand Down Expand Up @@ -63,6 +63,7 @@ pub struct AppContext {
pub worker_service: Arc<WorkerService>,
pub inflight_tracker: Arc<InFlightRequestTracker>,
pub kv_event_monitor: Option<Arc<KvEventMonitor>>,
pub realtime_registry: Arc<RealtimeRegistry>,
}

impl std::fmt::Debug for AppContext {
Expand Down Expand Up @@ -301,6 +302,7 @@ impl AppContextBuilder {
worker_service,
inflight_tracker: InFlightRequestTracker::new(),
kv_event_monitor: self.kv_event_monitor,
realtime_registry: Arc::new(RealtimeRegistry::new()),
Comment thread
coderabbitai[bot] marked this conversation as resolved.
})
}

Expand Down
1 change: 1 addition & 0 deletions model_gateway/src/routers/openai/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@

mod context;
mod provider;
pub mod realtime;
pub mod responses;
mod router;

Expand Down
11 changes: 11 additions & 0 deletions model_gateway/src/routers/openai/realtime/mod.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,11 @@
//! OpenAI Realtime API gateway implementation.
//!
//! Supports three transport mechanisms:
//! - **WebSocket** (server-to-server): Bidirectional WS proxy with transparent MCP interception
//! - **WebRTC** (browser-to-server): SDP signaling proxy; media + data channel flow directly
//! - **REST**: Ephemeral token generation (`client_secrets`, `sessions`, `transcription_sessions`)

pub mod registry;
pub mod rest;

pub use registry::RealtimeRegistry;
334 changes: 334 additions & 0 deletions model_gateway/src/routers/openai/realtime/registry.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,334 @@
//! In-memory session and call registry for Realtime API connections.

use std::{
sync::{
atomic::{AtomicUsize, Ordering},
Arc,
},
time::{Duration, Instant},
};

use dashmap::DashMap;
use tokio_util::sync::CancellationToken;
use tracing::{debug, info, warn};

/// Connection state for a realtime session.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConnectionState {
/// WebSocket upgrade accepted but upstream not yet connected.
Pending,
/// Bidirectional proxy is active.
Connected,
/// Connection has been closed.
Disconnected,
}

/// A tracked WebSocket session.
#[derive(Debug, Clone)]
pub struct SessionEntry {
pub session_id: String,
pub model: String,
pub worker_url: String,
pub state: ConnectionState,
pub created_at: Instant,
pub cancel_token: CancellationToken,
}

/// A tracked WebRTC call.
#[derive(Debug, Clone)]
pub struct CallEntry {
pub call_id: String,
pub model: String,
pub worker_url: String,
pub state: ConnectionState,
pub created_at: Instant,
pub cancel_token: CancellationToken,
}

const DEFAULT_MAX_SESSIONS: usize = 10_000;
const DEFAULT_MAX_CALLS: usize = 10_000;

/// DashMap-backed registry for realtime sessions and WebRTC calls.
///
/// Uses atomic counters for capacity enforcement to avoid TOCTOU races
/// between the length check and the DashMap insert.
#[derive(Debug)]
Comment thread
pallasathena92 marked this conversation as resolved.
pub struct RealtimeRegistry {
sessions: DashMap<String, SessionEntry>,
calls: DashMap<String, CallEntry>,
session_count: AtomicUsize,
call_count: AtomicUsize,
max_sessions: usize,
max_calls: usize,
}

impl RealtimeRegistry {
pub fn new() -> Self {
Self {
sessions: DashMap::new(),
calls: DashMap::new(),
session_count: AtomicUsize::new(0),
call_count: AtomicUsize::new(0),
max_sessions: DEFAULT_MAX_SESSIONS,
max_calls: DEFAULT_MAX_CALLS,
}
}

pub fn with_capacity(max_sessions: usize, max_calls: usize) -> Self {
Self {
sessions: DashMap::new(),
calls: DashMap::new(),
session_count: AtomicUsize::new(0),
call_count: AtomicUsize::new(0),
max_sessions,
max_calls,
}
}

// ---- Session methods ----

pub fn register_session(
&self,
session_id: String,
model: String,
worker_url: String,
) -> Option<SessionEntry> {
if !self.try_reserve_session() {
warn!(
max = self.max_sessions,
"Session registry at capacity, rejecting registration"
);
return None;
}
let entry = SessionEntry {
session_id: session_id.clone(),
model,
worker_url,
state: ConnectionState::Pending,
created_at: Instant::now(),
cancel_token: CancellationToken::new(),
};
if let Some(old) = self.sessions.insert(session_id, entry.clone()) {
// Replaced an existing entry — cancel its token so awaiting tasks
// are notified, and undo the extra reservation.
old.cancel_token.cancel();
self.session_count.fetch_sub(1, Ordering::Relaxed);
}
Some(entry)
}

pub fn set_session_state(&self, session_id: &str, state: ConnectionState) {
if let Some(mut entry) = self.sessions.get_mut(session_id) {
entry.state = state;
}
}

pub fn get_session(&self, session_id: &str) -> Option<SessionEntry> {
self.sessions.get(session_id).map(|e| e.clone())
}

pub fn remove_session(&self, session_id: &str) -> Option<SessionEntry> {
self.sessions.remove(session_id).map(|(_, e)| {
e.cancel_token.cancel();
self.session_count.fetch_sub(1, Ordering::Relaxed);
e
})
}

Comment thread
coderabbitai[bot] marked this conversation as resolved.
// ---- Call methods ----

pub fn register_call(
&self,
call_id: String,
model: String,
worker_url: String,
) -> Option<CallEntry> {
if !self.try_reserve_call() {
warn!(
max = self.max_calls,
"Call registry at capacity, rejecting registration"
);
return None;
}
let entry = CallEntry {
call_id: call_id.clone(),
model,
worker_url,
state: ConnectionState::Pending,
created_at: Instant::now(),
cancel_token: CancellationToken::new(),
};
if let Some(old) = self.calls.insert(call_id, entry.clone()) {
// Replaced an existing entry — cancel its token so awaiting tasks
// are notified, and undo the extra reservation.
old.cancel_token.cancel();
self.call_count.fetch_sub(1, Ordering::Relaxed);
}
Some(entry)
}

pub fn get_call(&self, call_id: &str) -> Option<CallEntry> {
self.calls.get(call_id).map(|e| e.clone())
}

pub fn set_call_state(&self, call_id: &str, state: ConnectionState) {
if let Some(mut entry) = self.calls.get_mut(call_id) {
entry.state = state;
}
}

pub fn remove_call(&self, call_id: &str) -> Option<CallEntry> {
self.calls.remove(call_id).map(|(_, e)| {
e.cancel_token.cancel();
self.call_count.fetch_sub(1, Ordering::Relaxed);
e
})
}

// ---- Atomic reservation helpers ----

/// Atomically reserve a session slot. Returns `true` if a slot was
/// successfully claimed, `false` if at capacity.
fn try_reserve_session(&self) -> bool {
self.try_reserve(&self.session_count, self.max_sessions)
}

/// Atomically reserve a call slot.
fn try_reserve_call(&self) -> bool {
self.try_reserve(&self.call_count, self.max_calls)
}

/// CAS loop: increment `counter` only if it is below `max`.
#[expect(
clippy::unused_self,
reason = "method for API consistency with other registry methods"
)]
fn try_reserve(&self, counter: &AtomicUsize, max: usize) -> bool {
loop {
let current = counter.load(Ordering::Relaxed);
if current >= max {
return false;
}
if counter
.compare_exchange_weak(current, current + 1, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
return true;
}
}
}

// ---- Reaper ----

/// Start a background task that evicts stale entries.
///
/// Returns a `CancellationToken` that stops the reaper when cancelled.
pub fn start_reaper(
self: &Arc<Self>,
max_age: Duration,
interval: Duration,
) -> CancellationToken {
let shutdown = CancellationToken::new();
let token = shutdown.clone();
let registry = Arc::clone(self);
#[expect(
clippy::disallowed_methods,
reason = "reaper task cancelled via returned token"
)]
tokio::spawn(async move {
let mut tick = tokio::time::interval(interval);
loop {
tokio::select! {
_ = tick.tick() => {}
() = shutdown.cancelled() => {
info!("Realtime registry reaper shutting down");
return;
}
}
let now = Instant::now();

// Reap stale, non-active entries atomically using remove_if
// to avoid TOCTOU races with concurrent re-registration.
// Active (Connected) sessions are never reaped — they will
// be cleaned up when the connection closes normally.
let stale_session_ids: Vec<String> = registry
.sessions
.iter()
.filter(|e| {
e.state != ConnectionState::Connected
&& now.duration_since(e.created_at) > max_age
})
.map(|e| e.session_id.clone())
.collect();

let mut sessions_reaped = 0usize;
for id in &stale_session_ids {
if let Some((_, entry)) = registry.sessions.remove_if(id, |_, e| {
e.state != ConnectionState::Connected
&& now.duration_since(e.created_at) > max_age
}) {
entry.cancel_token.cancel();
sessions_reaped += 1;
}
}

let stale_call_ids: Vec<String> = registry
.calls
.iter()
.filter(|e| {
e.state != ConnectionState::Connected
&& now.duration_since(e.created_at) > max_age
})
.map(|e| e.call_id.clone())
.collect();

let mut calls_reaped = 0usize;
for id in &stale_call_ids {
if let Some((_, entry)) = registry.calls.remove_if(id, |_, e| {
e.state != ConnectionState::Connected
&& now.duration_since(e.created_at) > max_age
}) {
entry.cancel_token.cancel();
calls_reaped += 1;
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

// Sync atomic counters with actual removal counts.
if sessions_reaped > 0 {
registry
.session_count
.fetch_sub(sessions_reaped, Ordering::Relaxed);
}
if calls_reaped > 0 {
registry
.call_count
.fetch_sub(calls_reaped, Ordering::Relaxed);
}

if sessions_reaped > 0 || calls_reaped > 0 {
debug!(
sessions_reaped,
calls_reaped, "Realtime registry reaper cycle"
);
}
}
});
info!("Realtime registry reaper started (max_age={max_age:?}, interval={interval:?})");
token
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.

/// Stats for observability.
pub fn session_count(&self) -> usize {
self.session_count.load(Ordering::Relaxed)
}

pub fn call_count(&self) -> usize {
self.call_count.load(Ordering::Relaxed)
}
}

impl Default for RealtimeRegistry {
fn default() -> Self {
Self::new()
}
}
Loading