use crate::config::RealtimeConfig;
use crate::error::{RealtimeError, Result};
use crate::events::ServerEvent;
use crate::openai::protocol::OpenAITransportLink;
use async_trait::async_trait;
use futures::{Sink, SinkExt, StreamExt};
use serde_json::Value;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::sync::{Mutex, mpsc, oneshot};
use tokio::task::JoinHandle;
use tokio_tungstenite::{
connect_async,
tungstenite::{
Message,
http::{Request, Uri},
},
};
type WsStream =
tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>;
type WsSource = futures::stream::SplitStream<WsStream>;
const OUTBOUND_CAPACITY: usize = 64;
const CLOSE_GRACE: Duration = Duration::from_secs(5);
pub struct OpenAIRealtimeSession {
session_id: String,
connected: Arc<AtomicBool>,
outbound: mpsc::Sender<Message>,
close: Mutex<Option<oneshot::Sender<()>>>,
writer: Mutex<Option<JoinHandle<()>>>,
receiver: Arc<Mutex<WsSource>>,
}
fn payload_digest(text: &str) -> String {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
text.hash(&mut hasher);
format!("{:016x}", hasher.finish())
}
fn redacted_payload(text: &str) -> String {
if cfg!(feature = "record-payloads") {
let end = text
.char_indices()
.map(|(index, _)| index)
.chain(std::iter::once(text.len()))
.take_while(|index| *index <= 300)
.last()
.unwrap_or(0);
return text[..end].to_string();
}
"<redacted>".to_string()
}
async fn writer_loop<S>(
mut sink: S,
mut outbound: mpsc::Receiver<Message>,
mut close: oneshot::Receiver<()>,
connected: Arc<AtomicBool>,
) where
S: Sink<Message> + Unpin,
S::Error: std::fmt::Display,
{
let mut peer_is_gone = false;
loop {
let message = tokio::select! {
biased;
_ = &mut close => break,
message = outbound.recv() => match message {
Some(message) => message,
None => break,
},
};
if let Err(error) = sink.send(message).await {
tracing::warn!(error = %error, "openai realtime write failed; connection is gone");
peer_is_gone = true;
break;
}
}
if !peer_is_gone {
let _ = sink.send(Message::Close(None)).await;
}
connected.store(false, Ordering::SeqCst);
}
async fn shutdown_writer(
close: &Mutex<Option<oneshot::Sender<()>>>,
writer: &Mutex<Option<JoinHandle<()>>>,
grace: Duration,
) -> Result<()> {
if let Some(signal) = close.lock().await.take() {
let _ = signal.send(());
}
let handle = writer.lock().await.take();
if let Some(mut handle) = handle
&& tokio::time::timeout(grace, &mut handle).await.is_err()
{
tracing::warn!(
grace_secs = grace.as_secs_f64(),
"openai realtime writer did not finish; abandoning it"
);
handle.abort();
return Err(RealtimeError::connection("Close timed out; writer abandoned"));
}
Ok(())
}
impl OpenAIRealtimeSession {
pub async fn connect(url: &str, api_key: &str, config: RealtimeConfig) -> Result<Self> {
let uri: Uri =
url.parse().map_err(|e| RealtimeError::connection(format!("Invalid URL: {}", e)))?;
let host = uri.host().unwrap_or("api.openai.com");
let request = Request::builder()
.uri(url)
.header("Host", host)
.header("Authorization", format!("Bearer {}", api_key))
.header("Sec-WebSocket-Key", generate_ws_key())
.header("Sec-WebSocket-Version", "13")
.header("Connection", "Upgrade")
.header("Upgrade", "websocket")
.body(())
.map_err(|e| RealtimeError::connection(format!("Request build error: {}", e)))?;
let (ws_stream, _response) = connect_async(request)
.await
.map_err(|e| RealtimeError::connection(format!("WebSocket connect error: {}", e)))?;
let (sink, source) = ws_stream.split();
let connected = Arc::new(AtomicBool::new(true));
let (outbound, outbound_rx) = mpsc::channel(OUTBOUND_CAPACITY);
let (close, close_rx) = oneshot::channel();
let writer = tokio::spawn(writer_loop(sink, outbound_rx, close_rx, Arc::clone(&connected)));
let session_id = uuid::Uuid::new_v4().to_string();
let session = Self {
session_id,
connected,
outbound,
close: Mutex::new(Some(close)),
writer: Mutex::new(Some(writer)),
receiver: Arc::new(Mutex::new(source)),
};
session.configure_session(config).await?;
Ok(session)
}
}
#[async_trait]
impl OpenAITransportLink for OpenAIRealtimeSession {
fn session_id(&self) -> &str {
&self.session_id
}
fn is_connected(&self) -> bool {
self.connected.load(Ordering::SeqCst)
}
async fn send_raw(&self, value: &Value) -> Result<()> {
let msg = serde_json::to_string(value)
.map_err(|e| RealtimeError::protocol(format!("JSON serialize error: {}", e)))?;
self.outbound
.send(Message::Text(msg.into()))
.await
.map_err(|e| RealtimeError::connection(format!("Send error: {}", e)))?;
Ok(())
}
async fn receive_raw(&self) -> Option<Result<ServerEvent>> {
let mut receiver = self.receiver.lock().await;
match receiver.next().await {
Some(Ok(Message::Text(text))) => {
let event_type = serde_json::from_str::<serde_json::Value>(&text)
.ok()
.and_then(|v| v.get("type").and_then(|t| t.as_str()).map(String::from))
.unwrap_or_else(|| "unknown".to_string());
match serde_json::from_str::<ServerEvent>(&text) {
Ok(ServerEvent::Unknown) => {
tracing::debug!(
event_type = %event_type,
"unmodeled realtime event, ignored"
);
Some(Ok(ServerEvent::Unknown))
}
Ok(event) => Some(Ok(event)),
Err(e) => {
tracing::warn!(
event_type = %event_type,
error = %e,
payload.bytes = text.len(),
payload.digest = %payload_digest(&text),
payload.raw = %redacted_payload(&text),
"recognized realtime event failed to parse (schema drift?)"
);
Some(Ok(ServerEvent::Unknown))
}
}
}
Some(Ok(Message::Close(_))) => {
self.connected.store(false, Ordering::SeqCst);
None
}
Some(Ok(_)) => {
Some(Ok(ServerEvent::Unknown))
}
Some(Err(e)) => {
self.connected.store(false, Ordering::SeqCst);
Some(Err(RealtimeError::connection(format!("Receive error: {}", e))))
}
None => {
self.connected.store(false, Ordering::SeqCst);
None
}
}
}
async fn close(&self) -> Result<()> {
self.connected.store(false, Ordering::SeqCst);
shutdown_writer(&self.close, &self.writer, CLOSE_GRACE).await
}
}
impl Drop for OpenAIRealtimeSession {
fn drop(&mut self) {
if let Some(handle) = self.writer.get_mut().take() {
handle.abort();
}
}
}
fn generate_ws_key() -> String {
use base64::Engine;
let mut key = [0u8; 16];
getrandom::fill(&mut key).unwrap_or_default();
base64::engine::general_purpose::STANDARD.encode(key)
}
impl std::fmt::Debug for OpenAIRealtimeSession {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OpenAIRealtimeSession")
.field("session_id", &self.session_id)
.field("connected", &self.connected.load(Ordering::SeqCst))
.finish()
}
}
#[cfg(test)]
mod writer_tests {
use super::{CLOSE_GRACE, OUTBOUND_CAPACITY, shutdown_writer, writer_loop};
use futures::Sink;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use tokio::sync::{Mutex, mpsc, oneshot};
use tokio_tungstenite::tungstenite::Message;
struct StalledSink;
impl Sink<Message> for StalledSink {
type Error = std::io::Error;
fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
fn start_send(self: Pin<&mut Self>, _: Message) -> Result<(), Self::Error> {
unreachable!("poll_ready never resolves, so start_send is never reached")
}
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Pending
}
}
#[derive(Clone, Default)]
struct RecordingSink(Arc<std::sync::Mutex<Vec<Message>>>);
impl Sink<Message> for RecordingSink {
type Error = std::io::Error;
fn poll_ready(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
self.0.lock().unwrap().push(item);
Ok(())
}
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
}
#[allow(clippy::type_complexity)]
fn spawn_writer<S>(
sink: S,
) -> (
mpsc::Sender<Message>,
Mutex<Option<oneshot::Sender<()>>>,
Mutex<Option<tokio::task::JoinHandle<()>>>,
Arc<AtomicBool>,
)
where
S: Sink<Message> + Unpin + Send + 'static,
S::Error: std::fmt::Display,
{
let connected = Arc::new(AtomicBool::new(true));
let (outbound, outbound_rx) = mpsc::channel(OUTBOUND_CAPACITY);
let (close, close_rx) = oneshot::channel();
let writer = tokio::spawn(writer_loop(sink, outbound_rx, close_rx, Arc::clone(&connected)));
(outbound, Mutex::new(Some(close)), Mutex::new(Some(writer)), connected)
}
#[tokio::test]
async fn teardown_reports_and_terminates_while_the_peer_is_still_stalled() {
let (outbound, close, writer, _connected) = spawn_writer(StalledSink);
outbound.send(Message::Text("in flight".into())).await.expect("queued");
tokio::task::yield_now().await;
let grace = Duration::from_millis(200);
let started = Instant::now();
let outcome =
tokio::time::timeout(Duration::from_secs(5), shutdown_writer(&close, &writer, grace))
.await
.expect("teardown must terminate even though the peer never drains");
assert!(outcome.is_err(), "an abandoned writer must be reported, not swallowed");
assert!(started.elapsed() >= grace, "the grace period must actually be honoured");
assert!(writer.lock().await.is_none(), "the writer handle must be released");
}
#[tokio::test]
async fn close_preempts_a_completely_full_queue() {
let sink = RecordingSink::default();
let (outbound, close, writer, _connected) = spawn_writer(sink.clone());
for frame in 0..OUTBOUND_CAPACITY {
outbound
.try_send(Message::Text(format!("frame {frame}").into()))
.expect("queue accepts up to capacity");
}
assert!(
outbound.try_send(Message::Text("overflow".into())).is_err(),
"the queue must be bounded — an unbounded queue would hide the very stall \
this design exists to survive"
);
tokio::time::timeout(Duration::from_secs(5), shutdown_writer(&close, &writer, CLOSE_GRACE))
.await
.expect("teardown must not wait out the backlog")
.expect("a healthy peer must close cleanly");
let written = sink.0.lock().unwrap().clone();
assert!(
matches!(written.last(), Some(Message::Close(_))),
"the close frame must be written: {written:?}"
);
assert!(
written.len() < OUTBOUND_CAPACITY,
"close must preempt the backlog rather than drain all {OUTBOUND_CAPACITY} \
queued frames first; wrote {}",
written.len()
);
}
#[tokio::test]
async fn a_healthy_peer_receives_the_close_frame_promptly() {
let sink = RecordingSink::default();
let (outbound, close, writer, connected) = spawn_writer(sink.clone());
outbound.send(Message::Text("hello".into())).await.expect("queued");
tokio::time::timeout(
Duration::from_millis(500),
shutdown_writer(&close, &writer, Duration::from_secs(30)),
)
.await
.expect("a healthy peer must close promptly, not merely within the grace period")
.expect("a healthy peer must close cleanly");
let written = sink.0.lock().unwrap().clone();
assert!(
matches!(written.last(), Some(Message::Close(_))),
"the close frame must reach a peer that is reading: {written:?}"
);
assert!(!connected.load(Ordering::SeqCst), "the writer marks the session disconnected");
}
}
#[cfg(test)]
mod redaction_tests {
use super::{payload_digest, redacted_payload};
#[cfg(not(feature = "record-payloads"))]
#[derive(Clone, Default)]
struct CapturedLogs(std::sync::Arc<std::sync::Mutex<Vec<u8>>>);
#[cfg(not(feature = "record-payloads"))]
impl std::io::Write for CapturedLogs {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
#[cfg(not(feature = "record-payloads"))]
impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for CapturedLogs {
type Writer = Self;
fn make_writer(&'a self) -> Self::Writer {
self.clone()
}
}
fn drifting_frame() -> String {
serde_json::json!({
"type": "conversation.item.input_audio_transcription.completed",
"transcript": "my card is 4111-1111-1111-1111",
"item_id": 12345
})
.to_string()
}
#[cfg(not(feature = "record-payloads"))]
#[test]
fn the_frame_is_withheld_by_default() {
assert_eq!(redacted_payload(&drifting_frame()), "<redacted>");
}
#[cfg(feature = "record-payloads")]
#[test]
fn the_frame_is_recorded_when_explicitly_enabled() {
let recorded = redacted_payload(&drifting_frame());
assert!(recorded.contains("transcript"), "the frame must be recorded under the feature");
assert!(recorded.len() <= 300, "recording stays bounded: {}", recorded.len());
}
#[cfg(not(feature = "record-payloads"))]
#[test]
fn a_drift_warning_logs_a_summary_and_no_content() {
let frame = drifting_frame();
let logs = CapturedLogs::default();
let layer = tracing_subscriber::fmt::layer()
.with_writer(logs.clone())
.with_ansi(false)
.with_target(false);
let collector = {
use tracing_subscriber::layer::SubscriberExt;
tracing_subscriber::registry().with(layer)
};
tracing::subscriber::with_default(collector, || {
tracing::warn!(
event_type = "conversation.item.input_audio_transcription.completed",
payload.bytes = frame.len(),
payload.digest = %payload_digest(&frame),
payload.raw = %redacted_payload(&frame),
"recognized realtime event failed to parse (schema drift?)"
);
});
let captured = String::from_utf8_lossy(&logs.0.lock().unwrap()).to_string();
assert!(
!captured.contains("4111-1111-1111-1111"),
"the frame's content reached the log: {captured}"
);
assert!(captured.contains("<redacted>"), "the log must say the frame was withheld");
assert!(captured.contains("payload.bytes"), "the size must remain available");
assert!(captured.contains("payload.digest"), "a digest must remain available");
}
#[test]
fn the_digest_correlates_repeats_and_separates_shapes() {
let frame = r#"{"type":"response.done","response":{"unexpected":true}}"#;
assert_eq!(payload_digest(frame), payload_digest(frame));
assert_ne!(payload_digest(frame), payload_digest(r#"{"type":"response.done"}"#));
}
}