Skip to main content

zai_rs/realtime/
session.rs

1//! [`RealtimeSession`] — an active realtime conversation over WebSocket.
2//!
3//! A session owns a background event-loop task that pumps client events onto
4//! the socket and fans server events (and decoded audio) out to subscribers.
5//! Callers drive it via command methods (`send_audio`, `send_text`, …) and
6//! consume the two streams: [`RealtimeSession::events`] and
7//! [`RealtimeSession::audio_stream`].
8
9use std::{
10    collections::HashSet,
11    pin::Pin,
12    sync::{Arc, Mutex},
13    time::Duration,
14};
15
16use bytes::Bytes;
17use futures_util::{Stream, stream};
18use tokio::{
19    sync::{broadcast, mpsc, watch},
20    task::JoinHandle,
21};
22use tracing::{debug, warn};
23
24use super::{
25    audio::{
26        InputAudioFormat, OutputAudioFormat, decode_base64, encode_base64,
27        encode_jpeg_frame_base64, encode_wav_pcm_base64,
28    },
29    client::AuthMode,
30    events::{ClientEvent, ServerEvent},
31    jwt,
32    protocol::{
33        ChatMode, GreetingConfig, InputAudioNoiseReduction, NoiseReductionType, RealtimeModality,
34        RealtimeTool, RealtimeVoice, SessionConfig, TurnDetectionType,
35    },
36    transport::{RealtimeTransport, WsMessage},
37};
38use crate::{
39    ZaiResult,
40    client::{
41        error::RealtimeErrorKind,
42        secret::ApiSecret,
43        transport::limits::{REALTIME_AUDIO_FRAME_MAX, WS_MESSAGE_MAX},
44    },
45};
46
47/// The server documents an application heartbeat approximately every 30
48/// seconds. Three missed heartbeats is treated as a dead/half-open session.
49const INBOUND_IDLE_TIMEOUT: Duration = Duration::from_secs(90);
50/// Frozen realtime contract: each per-session queue is deliberately small so
51/// backpressure or consumer lag is detected before large audio/event backlogs
52/// accumulate in memory.
53const SESSION_CHANNEL_CAPACITY: usize = 8;
54
55/// One decoded `response.audio.delta` payload with its correlation metadata.
56#[derive(Debug, Clone)]
57pub struct RealtimeAudioChunk {
58    /// Id of the response this audio belongs to.
59    pub response_id: String,
60    /// Id of the output item this audio belongs to.
61    pub item_id: String,
62    /// Index of the output item within the response, when supplied.
63    pub output_index: Option<u64>,
64    /// Index of the content part within the output item, when supplied.
65    pub content_index: Option<u64>,
66    /// Decoded raw 24 kHz, mono, 16-bit PCM bytes.
67    pub data: Bytes,
68}
69
70/// Builder for an [`RealtimeSession`].
71///
72/// Produced by [`super::client::RealtimeClient::session`]. Configure the
73/// session defaults, then [`SessionBuilder::build`] opens the WebSocket and
74/// sends the initial `session.update`.
75pub struct SessionBuilder {
76    api_key: Arc<ApiSecret>,
77    auth: AuthMode,
78    realtime_url: String,
79    model_name: String,
80    session_config: SessionConfig,
81}
82
83impl SessionBuilder {
84    pub(super) fn new(
85        api_key: Arc<ApiSecret>,
86        auth: AuthMode,
87        realtime_url: String,
88        model_name: String,
89    ) -> Self {
90        Self {
91            api_key,
92            auth,
93            realtime_url,
94            model_name,
95            session_config: SessionConfig::default(),
96        }
97    }
98
99    /// System instructions guiding the model.
100    pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
101        self.session_config.instructions = Some(instructions.into());
102        self
103    }
104
105    /// VAD strategy (defaults to client-VAD).
106    pub fn turn_detection(mut self, vad: TurnDetectionType) -> Self {
107        self.session_config.turn_detection.type_ = vad;
108        self
109    }
110
111    /// Configure whether server VAD automatically creates a response at the
112    /// end of a detected speech turn.
113    pub fn create_response_on_vad(mut self, enabled: bool) -> Self {
114        self.session_config.turn_detection.create_response = Some(enabled);
115        self
116    }
117
118    /// Configure whether server VAD interrupts an in-progress response when
119    /// new speech begins.
120    pub fn interrupt_response_on_vad(mut self, enabled: bool) -> Self {
121        self.session_config.turn_detection.interrupt_response = Some(enabled);
122        self
123    }
124
125    /// Configure the server-VAD activation threshold (`0.0..=1.0`).
126    pub fn vad_threshold(mut self, threshold: f64) -> Self {
127        self.session_config.turn_detection.threshold = Some(threshold);
128        self
129    }
130
131    /// Configure how much audio before detected speech is retained.
132    pub fn vad_prefix_padding_ms(mut self, milliseconds: u32) -> Self {
133        self.session_config.turn_detection.prefix_padding_ms = Some(milliseconds);
134        self
135    }
136
137    /// Configure how much silence ends a server-VAD turn.
138    pub fn vad_silence_duration_ms(mut self, milliseconds: u32) -> Self {
139        self.session_config.turn_detection.silence_duration_ms = Some(milliseconds);
140        self
141    }
142
143    /// Input audio format (defaults to 16 kHz WAV).
144    pub fn input_audio_format(mut self, format: InputAudioFormat) -> Self {
145        self.session_config.input_audio_format = format;
146        self
147    }
148
149    /// Output audio format (defaults to PCM).
150    pub fn output_audio_format(mut self, format: OutputAudioFormat) -> Self {
151        self.session_config.output_audio_format = format;
152        self
153    }
154
155    /// Output modalities (text, audio, or both).
156    pub fn modalities(mut self, modalities: impl IntoIterator<Item = RealtimeModality>) -> Self {
157        self.session_config.modalities = modalities.into_iter().collect();
158        self
159    }
160
161    /// Voice used for generated audio.
162    pub fn voice(mut self, voice: RealtimeVoice) -> Self {
163        self.session_config.voice = Some(voice);
164        self
165    }
166
167    /// Sampling temperature. Values outside `0.0..=1.0` are rejected by
168    /// [`Self::build`] before a connection is opened.
169    pub fn temperature(mut self, temperature: f64) -> Self {
170        self.session_config.temperature = Some(temperature);
171        self
172    }
173
174    /// Maximum response text-token count. Values above 1024 are rejected by
175    /// [`Self::build`] before a connection is opened.
176    pub fn max_response_output_tokens(mut self, tokens: u16) -> Self {
177        self.session_config.max_response_output_tokens = Some(tokens);
178        self
179    }
180
181    /// Configure input-audio noise reduction for the microphone placement.
182    pub fn input_audio_noise_reduction(mut self, profile: NoiseReductionType) -> Self {
183        self.session_config.input_audio_noise_reduction =
184            Some(InputAudioNoiseReduction::new(profile));
185        self
186    }
187
188    /// Conversation mode under `beta_fields.chat_mode`.
189    pub fn chat_mode(mut self, mode: ChatMode) -> Self {
190        self.session_config.beta_fields.chat_mode = Some(mode);
191        self
192    }
193
194    /// Enable/disable the server-side built-in web search.
195    pub fn auto_search(mut self, enabled: bool) -> Self {
196        self.session_config.beta_fields.auto_search = Some(enabled);
197        self
198    }
199
200    /// Configure an optional server-generated greeting.
201    pub fn greeting_config(mut self, greeting: GreetingConfig) -> Self {
202        self.session_config.greeting_config = Some(greeting);
203        self
204    }
205
206    /// Register function tools.
207    pub fn tools(mut self, tools: Vec<RealtimeTool>) -> Self {
208        self.session_config.tools = tools;
209        self
210    }
211
212    /// Override the entire session config.
213    ///
214    /// The model is still taken from [`RealtimeClient::session`](super::RealtimeClient::session)
215    /// and replaces `config.model` during [`Self::build`].
216    pub fn session_config(mut self, config: SessionConfig) -> Self {
217        self.session_config = config;
218        self
219    }
220
221    /// Open the WebSocket, send `session.update`, and spawn the event loop.
222    #[tracing::instrument(name = "realtime.session.build", skip_all, fields(model = %self.model_name))]
223    pub async fn build(self) -> ZaiResult<RealtimeSession> {
224        let Self {
225            api_key,
226            auth,
227            realtime_url,
228            model_name,
229            mut session_config,
230        } = self;
231
232        // The selected type-safe model is part of the session.update wire
233        // contract. It takes precedence over an arbitrary value supplied via
234        // `session_config` so the marker-trait guarantee cannot be bypassed.
235        session_config.model = Some(model_name.clone());
236        validate_session_config(&session_config)?;
237        let input_audio_format = session_config.input_audio_format;
238        let init = ClientEvent::SessionUpdate {
239            event_id: Some(new_event_id()),
240            session: session_config,
241        };
242        // Serialize and enforce the message limit before opening a socket so a
243        // locally invalid configuration cannot cause network side effects.
244        let init = serialize_event(&init)?;
245
246        let jwt_ttl = match auth {
247            AuthMode::Bearer => None,
248            AuthMode::Jwt { ttl_seconds } => Some(ttl_seconds),
249        };
250        let authorization = jwt::authorization_header(api_key.expose(), jwt_ttl)?;
251
252        let mut transport =
253            super::transport::TungsteniteTransport::connect(&realtime_url, &authorization).await?;
254
255        if let Err(error) = transport.send(init).await {
256            let _ = transport.close().await;
257            return Err(error);
258        }
259        debug!(model = %model_name, "Realtime session opened");
260
261        let (cmd_tx, cmd_rx) = mpsc::channel::<String>(SESSION_CHANNEL_CAPACITY);
262        let (shutdown_tx, shutdown_rx) = watch::channel(false);
263        let (events_tx, _) = broadcast::channel::<ServerEvent>(SESSION_CHANNEL_CAPACITY);
264        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(SESSION_CHANNEL_CAPACITY);
265        // Subscribe before the event loop starts so session-created events and
266        // greeting audio cannot race ahead of the caller's first subscription.
267        let initial_events_rx = events_tx.subscribe();
268        let initial_audio_rx = audio_tx.subscribe();
269
270        let (completion_tx, completion_rx) = watch::channel(None);
271        let loop_events_tx = events_tx.clone();
272        let loop_audio_tx = audio_tx.clone();
273        let join = tokio::spawn(async move {
274            let result = run_loop(
275                transport,
276                cmd_rx,
277                shutdown_rx,
278                loop_events_tx,
279                loop_audio_tx,
280            )
281            .await;
282            completion_tx.send_replace(Some(result.clone()));
283            result
284        });
285
286        Ok(RealtimeSession {
287            cmd_tx,
288            shutdown_tx,
289            events_tx,
290            audio_tx,
291            initial_events_rx: Mutex::new(Some(initial_events_rx)),
292            initial_audio_rx: Mutex::new(Some(initial_audio_rx)),
293            completion_rx,
294            model_name,
295            input_audio_format,
296            join,
297        })
298    }
299}
300
301/// Background event-loop body: drains commands onto the socket and fans server
302/// messages out to the broadcast channels. Generic over the transport so a mock
303/// can be substituted in tests.
304async fn run_loop<T: RealtimeTransport>(
305    mut transport: T,
306    mut cmd_rx: mpsc::Receiver<String>,
307    mut shutdown_rx: watch::Receiver<bool>,
308    events_tx: broadcast::Sender<ServerEvent>,
309    audio_tx: broadcast::Sender<RealtimeAudioChunk>,
310) -> ZaiResult<()> {
311    let idle_deadline = tokio::time::sleep(INBOUND_IDLE_TIMEOUT);
312    tokio::pin!(idle_deadline);
313
314    loop {
315        tokio::select! {
316            biased;
317            changed = shutdown_rx.changed() => {
318                if changed.is_err() || *shutdown_rx.borrow() {
319                    debug!("Realtime session closed (client requested)");
320                    return transport.close().await;
321                }
322            },
323            // Poll inbound traffic before the deadline so a heartbeat arriving
324            // exactly at the boundary refreshes the session instead of racing
325            // with a false timeout. One queued command is drained after every
326            // inbound frame, preventing either direction from starving the
327            // other under sustained traffic.
328            msg = transport.recv() => match msg {
329                Ok(Some(WsMessage::Text(text))) => {
330                    idle_deadline
331                        .as_mut()
332                        .reset(tokio::time::Instant::now() + INBOUND_IDLE_TIMEOUT);
333                    match decode_server_frame(&text) {
334                        Ok(DecodedServerFrame::Event(event)) => {
335                            if let ServerEvent::Error { error } = event.as_ref() {
336                                // The free-form message may echo caller content, so
337                                // only the machine-readable code enters logs.
338                                warn!(code = ?error.code, "Realtime server error event");
339                            }
340                            let _ = events_tx.send(*event);
341                        },
342                        Ok(DecodedServerFrame::Audio(bytes)) => {
343                            let _ = audio_tx.send(bytes);
344                        },
345                        Ok(DecodedServerFrame::Unknown) => {
346                            warn!(bytes = text.len(), "Ignoring unknown realtime event");
347                        },
348                        Err(error) => {
349                            warn!(bytes = text.len(), "Closing session after malformed realtime event");
350                            let _ = transport.close().await;
351                            return Err(error);
352                        },
353                    }
354
355                    match cmd_rx.try_recv() {
356                        Ok(command) => {
357                            if !handle_outbound(&mut transport, Some(command)).await? {
358                                return Ok(());
359                            }
360                        },
361                        Err(mpsc::error::TryRecvError::Disconnected) => {
362                            handle_outbound(&mut transport, None).await?;
363                            return Ok(());
364                        },
365                        Err(mpsc::error::TryRecvError::Empty) => {},
366                    }
367                },
368                Ok(Some(WsMessage::Binary(bytes))) => {
369                    warn!(bytes = bytes.len(), "Closing session after unexpected realtime binary frame");
370                    let _ = transport.close().await;
371                    return Err(protocol_error(
372                        "unexpected binary frame in realtime JSON protocol",
373                    ));
374                },
375                Ok(None) => {
376                    debug!("Realtime session closed (peer disconnected)");
377                    return Ok(());
378                },
379                Err(error) => {
380                    // Avoid logging the error source: handshake/transport errors
381                    // can contain endpoint details or server-provided text.
382                    warn!("Realtime event loop terminated due to transport error");
383                    let _ = transport.close().await;
384                    return Err(error);
385                },
386            },
387            _ = &mut idle_deadline => {
388                warn!(
389                    timeout_seconds = INBOUND_IDLE_TIMEOUT.as_secs(),
390                    "Realtime session timed out waiting for inbound traffic"
391                );
392                let _ = transport.close().await;
393                return Err(RealtimeErrorKind::Timeout {
394                    operation: "Realtime inbound heartbeat",
395                }
396                .into());
397            },
398            cmd = cmd_rx.recv() => {
399                if !handle_outbound(&mut transport, cmd).await? {
400                    return Ok(());
401                }
402            },
403        }
404    }
405}
406
407enum DecodedServerFrame {
408    Event(Box<ServerEvent>),
409    Audio(RealtimeAudioChunk),
410    Unknown,
411}
412
413fn decode_server_frame(text: &str) -> ZaiResult<DecodedServerFrame> {
414    let event = serde_json::from_str::<ServerEvent>(text)
415        .map_err(|_| protocol_error("malformed realtime server event"))?;
416    match event {
417        ServerEvent::ResponseAudioDelta {
418            response_id,
419            item_id,
420            output_index,
421            content_index,
422            delta,
423        } => {
424            let bytes = decode_base64(&delta)?;
425            if bytes.len() as u64 > REALTIME_AUDIO_FRAME_MAX {
426                return Err(protocol_error(format!(
427                    "realtime audio delta exceeds {REALTIME_AUDIO_FRAME_MAX} bytes"
428                )));
429            }
430            Ok(DecodedServerFrame::Audio(RealtimeAudioChunk {
431                response_id,
432                item_id,
433                output_index,
434                content_index,
435                data: Bytes::from(bytes),
436            }))
437        },
438        ServerEvent::Unknown => Ok(DecodedServerFrame::Unknown),
439        event => Ok(DecodedServerFrame::Event(Box::new(event))),
440    }
441}
442
443async fn handle_outbound<T: RealtimeTransport>(
444    transport: &mut T,
445    message: Option<String>,
446) -> ZaiResult<bool> {
447    match message {
448        Some(json) => {
449            if let Err(error) = transport.send(json).await {
450                let _ = transport.close().await;
451                return Err(error);
452            }
453            Ok(true)
454        },
455        None => {
456            debug!("Realtime session closed (client requested)");
457            transport.close().await?;
458            Ok(false)
459        },
460    }
461}
462
463/// An active realtime session.
464///
465/// Cheap to share indirectly via the channels it owns; call
466/// [`RealtimeSession::close`] to terminate the background task.
467pub struct RealtimeSession {
468    cmd_tx: mpsc::Sender<String>,
469    shutdown_tx: watch::Sender<bool>,
470    events_tx: broadcast::Sender<ServerEvent>,
471    audio_tx: broadcast::Sender<RealtimeAudioChunk>,
472    initial_events_rx: Mutex<Option<broadcast::Receiver<ServerEvent>>>,
473    initial_audio_rx: Mutex<Option<broadcast::Receiver<RealtimeAudioChunk>>>,
474    completion_rx: watch::Receiver<Option<ZaiResult<()>>>,
475    model_name: String,
476    input_audio_format: InputAudioFormat,
477    join: JoinHandle<ZaiResult<()>>,
478}
479
480impl RealtimeSession {
481    /// Send raw 16-bit little-endian mono PCM.
482    ///
483    /// With [`InputAudioFormat::Wav`] (the default), the bytes are wrapped in a
484    /// 16 kHz WAV container. With `Pcm16` or `Pcm24`, they are sent as raw PCM
485    /// and the selected format declares the corresponding sample rate.
486    pub async fn send_audio(&self, pcm: Bytes) -> ZaiResult<()> {
487        if pcm.is_empty() {
488            return Err(protocol_error("realtime audio frame must not be empty"));
489        }
490        if pcm.len() as u64 > REALTIME_AUDIO_FRAME_MAX {
491            return Err(protocol_error(format!(
492                "realtime audio frame exceeds {REALTIME_AUDIO_FRAME_MAX} bytes"
493            )));
494        }
495        if pcm.len() % 2 != 0 {
496            return Err(protocol_error(
497                "16-bit PCM input must contain an even number of bytes",
498            ));
499        }
500        let audio = match self.input_audio_format {
501            InputAudioFormat::Wav => encode_wav_pcm_base64(&pcm, 16_000)?,
502            InputAudioFormat::Pcm16 | InputAudioFormat::Pcm24 => encode_base64(&pcm),
503        };
504        self.dispatch(ClientEvent::InputAudioBufferAppend {
505            audio,
506            client_timestamp: Some(now_ms()),
507        })
508        .await
509    }
510
511    /// Upload a JPEG frame for passive-video mode.
512    pub async fn send_video_frame(&self, jpeg: Bytes) -> ZaiResult<()> {
513        if jpeg.len() as u64 > REALTIME_AUDIO_FRAME_MAX {
514            return Err(protocol_error(format!(
515                "realtime video frame exceeds {REALTIME_AUDIO_FRAME_MAX} bytes"
516            )));
517        }
518        if !jpeg.starts_with(&[0xff, 0xd8]) || !jpeg.ends_with(&[0xff, 0xd9]) {
519            return Err(protocol_error(
520                "realtime video frame must be a complete JPEG image",
521            ));
522        }
523        self.dispatch(ClientEvent::InputAudioBufferAppendVideoFrame {
524            video_frame: encode_jpeg_frame_base64(&jpeg),
525            client_timestamp: Some(now_ms()),
526        })
527        .await
528    }
529
530    /// Commit buffered audio for inference in client-VAD mode. Server-VAD
531    /// commits automatically and normally does not need this command.
532    pub async fn commit_audio(&self) -> ZaiResult<()> {
533        self.dispatch(ClientEvent::InputAudioBufferCommit {
534            client_timestamp: Some(now_ms()),
535        })
536        .await
537    }
538
539    /// Clear audio buffered by the server without triggering inference.
540    pub async fn clear_audio(&self) -> ZaiResult<()> {
541        self.dispatch(ClientEvent::InputAudioBufferClear).await
542    }
543
544    /// Inject a user text message into the conversation history.
545    pub async fn send_text(&self, text: impl Into<String>) -> ZaiResult<()> {
546        let text = text.into();
547        if text.trim().is_empty() {
548            return Err(protocol_error("realtime text must not be blank"));
549        }
550        self.dispatch(ClientEvent::ConversationItemCreate {
551            event_id: Some(new_event_id()),
552            item: super::protocol::RealtimeConversationItem::user_text(text),
553        })
554        .await
555    }
556
557    /// Feed back a function-call result.
558    pub async fn send_function_output(
559        &self,
560        call_name: impl Into<String>,
561        output: impl Into<String>,
562    ) -> ZaiResult<()> {
563        let call_name = call_name.into();
564        if call_name.trim().is_empty() {
565            return Err(protocol_error("realtime function name must not be blank"));
566        }
567        let output = output.into();
568        if output.trim().is_empty() {
569            return Err(protocol_error("realtime function output must not be blank"));
570        }
571        self.dispatch(ClientEvent::ConversationItemCreate {
572            event_id: Some(new_event_id()),
573            item: super::protocol::RealtimeConversationItem::function_output(call_name, output),
574        })
575        .await
576    }
577
578    /// Delete an item from the server-side conversation history.
579    pub async fn delete_item(&self, item_id: impl Into<String>) -> ZaiResult<()> {
580        let item_id = item_id.into();
581        if item_id.trim().is_empty() {
582            return Err(protocol_error("realtime item id must not be blank"));
583        }
584        self.dispatch(ClientEvent::ConversationItemDelete {
585            event_id: Some(new_event_id()),
586            client_timestamp: Some(now_ms()),
587            item_id,
588        })
589        .await
590    }
591
592    /// Ask the server to emit the current representation of one conversation
593    /// item via [`ServerEvent::ConversationItemRetrieved`].
594    pub async fn retrieve_item(&self, item_id: impl Into<String>) -> ZaiResult<()> {
595        let item_id = item_id.into();
596        if item_id.trim().is_empty() {
597            return Err(protocol_error("realtime item id must not be blank"));
598        }
599        self.dispatch(ClientEvent::ConversationItemRetrieve {
600            event_id: Some(new_event_id()),
601            client_timestamp: Some(now_ms()),
602            item_id,
603        })
604        .await
605    }
606
607    /// Trigger model inference (`response.create`).
608    pub async fn create_response(&self) -> ZaiResult<()> {
609        self.dispatch(ClientEvent::ResponseCreate {
610            client_timestamp: Some(now_ms()),
611        })
612        .await
613    }
614
615    /// Cancel the in-flight response (`response.cancel`), e.g. on interruption.
616    pub async fn cancel(&self) -> ZaiResult<()> {
617        self.dispatch(ClientEvent::ResponseCancel {
618            client_timestamp: Some(now_ms()),
619        })
620        .await
621    }
622
623    /// Stream of server metadata events (transcripts, response lifecycle,
624    /// errors, heartbeats). Audio deltas are decoded only onto
625    /// [`Self::audio_stream`] to avoid retaining a second base64 copy. The
626    /// first subscriber receives events buffered since session creation;
627    /// later subscribers start at the live tail. A lagged consumer or
628    /// background session failure is surfaced as an error instead of silently
629    /// losing protocol events.
630    pub fn events(&self) -> Pin<Box<dyn Stream<Item = ZaiResult<ServerEvent>> + Send + '_>> {
631        observable_broadcast_stream(
632            subscribe_with_initial_backlog(&self.events_tx, &self.initial_events_rx),
633            self.completion_rx.clone(),
634            "realtime event",
635        )
636    }
637
638    /// Stream of decoded 24 kHz, mono, 16-bit PCM output chunks.
639    ///
640    /// Lag is an error because dropping a PCM chunk would silently corrupt the
641    /// resulting audio stream.
642    pub fn audio_stream(
643        &self,
644    ) -> Pin<Box<dyn Stream<Item = ZaiResult<RealtimeAudioChunk>> + Send + '_>> {
645        observable_broadcast_stream(
646            subscribe_with_initial_backlog(&self.audio_tx, &self.initial_audio_rx),
647            self.completion_rx.clone(),
648            "realtime audio",
649        )
650    }
651
652    /// The model id sent in the initial `session.update` event.
653    pub fn model_name(&self) -> &str {
654        &self.model_name
655    }
656
657    #[tracing::instrument(name = "realtime.dispatch", skip(self, event))]
658    async fn dispatch(&self, event: ClientEvent) -> ZaiResult<()> {
659        let message = serialize_event(&event)?;
660        self.cmd_tx
661            .send(message)
662            .await
663            .map_err(|_| RealtimeErrorKind::Closed.into())
664    }
665
666    /// Signal the background task to close without awaiting it.
667    ///
668    /// Use when the session is shared (e.g. behind an `Arc`), driven from a
669    /// `tokio::select!`, or closed reactively on a shutdown signal. This only
670    /// signals a dedicated shutdown channel; the background loop observes it
671    /// independently of the bounded outbound queue and exits.
672    /// For deterministic, awaited teardown use [`RealtimeSession::close`].
673    pub async fn request_close(&self) -> ZaiResult<()> {
674        self.shutdown_tx
675            .send(true)
676            .map_err(|_| RealtimeErrorKind::Closed.into())
677    }
678
679    /// Close the session and wait for the event loop to finish.
680    pub async fn close(self) -> ZaiResult<()> {
681        // Best-effort: the loop may already have ended due to a peer close or
682        // protocol error.
683        let _ = self.shutdown_tx.send(true);
684        // Preserve both task failures and transport-close failures for the
685        // caller instead of converting abnormal teardown into success.
686        match self.join.await {
687            Ok(result) => result,
688            Err(join_error) => Err(protocol_error(format!(
689                "realtime event loop join failed: {join_error}"
690            ))),
691        }
692    }
693}
694
695struct BroadcastState<T> {
696    receiver: broadcast::Receiver<T>,
697    completion: watch::Receiver<Option<ZaiResult<()>>>,
698    channel_name: &'static str,
699    terminal_reported: bool,
700    completion_lost: bool,
701}
702
703fn subscribe_with_initial_backlog<T: Clone>(
704    sender: &broadcast::Sender<T>,
705    initial: &Mutex<Option<broadcast::Receiver<T>>>,
706) -> broadcast::Receiver<T> {
707    initial
708        .lock()
709        .unwrap_or_else(|poisoned| poisoned.into_inner())
710        .take()
711        .unwrap_or_else(|| sender.subscribe())
712}
713
714fn observable_broadcast_stream<T>(
715    receiver: broadcast::Receiver<T>,
716    completion: watch::Receiver<Option<ZaiResult<()>>>,
717    channel_name: &'static str,
718) -> Pin<Box<dyn Stream<Item = ZaiResult<T>> + Send>>
719where
720    T: Clone + Send + 'static,
721{
722    let state = BroadcastState {
723        receiver,
724        completion,
725        channel_name,
726        terminal_reported: false,
727        completion_lost: false,
728    };
729    Box::pin(stream::unfold(state, |mut state| async move {
730        loop {
731            if state.terminal_reported {
732                return None;
733            }
734
735            if state.completion_lost {
736                match state.receiver.try_recv() {
737                    Ok(value) => return Some((Ok(value), state)),
738                    Err(broadcast::error::TryRecvError::Lagged(skipped)) => {
739                        let error = lagged_stream_error(state.channel_name, skipped);
740                        state.terminal_reported = true;
741                        return Some((Err(error), state));
742                    },
743                    Err(
744                        broadcast::error::TryRecvError::Empty
745                        | broadcast::error::TryRecvError::Closed,
746                    ) => {
747                        state.terminal_reported = true;
748                        return Some((
749                            Err(protocol_error(
750                                "realtime background task ended without a completion status",
751                            )),
752                            state,
753                        ));
754                    },
755                }
756            }
757
758            let completion = state.completion.borrow().clone();
759            if let Some(result) = completion {
760                match state.receiver.try_recv() {
761                    Ok(value) => return Some((Ok(value), state)),
762                    Err(broadcast::error::TryRecvError::Lagged(skipped)) => {
763                        let error = lagged_stream_error(state.channel_name, skipped);
764                        state.terminal_reported = true;
765                        return Some((Err(error), state));
766                    },
767                    Err(
768                        broadcast::error::TryRecvError::Empty
769                        | broadcast::error::TryRecvError::Closed,
770                    ) => match result {
771                        Ok(()) => return None,
772                        Err(error) => {
773                            state.terminal_reported = true;
774                            return Some((Err(error), state));
775                        },
776                    },
777                }
778            }
779
780            tokio::select! {
781                value = state.receiver.recv() => match value {
782                    Ok(value) => return Some((Ok(value), state)),
783                    Err(broadcast::error::RecvError::Lagged(skipped)) => {
784                        let error = lagged_stream_error(state.channel_name, skipped);
785                        state.terminal_reported = true;
786                        return Some((Err(error), state));
787                    },
788                    Err(broadcast::error::RecvError::Closed) => return None,
789                },
790                changed = state.completion.changed() => {
791                    if changed.is_err() {
792                        state.completion_lost = true;
793                    }
794                },
795            }
796        }
797    }))
798}
799
800fn lagged_stream_error(channel_name: &str, skipped: u64) -> crate::ZaiError {
801    protocol_error(format!(
802        "{channel_name} consumer lagged and lost {skipped} message(s)"
803    ))
804}
805
806fn validate_session_config(config: &SessionConfig) -> ZaiResult<()> {
807    if let Some(temperature) = config.temperature
808        && (!temperature.is_finite() || !(0.0..=1.0).contains(&temperature))
809    {
810        return Err(protocol_error(
811            "realtime temperature must be a finite value between 0 and 1",
812        ));
813    }
814    if let Some(tokens) = config.max_response_output_tokens
815        && !(1..=1024).contains(&tokens)
816    {
817        return Err(protocol_error(
818            "realtime max_response_output_tokens must be between 1 and 1024",
819        ));
820    }
821    if config.modalities.is_empty() {
822        return Err(protocol_error(
823            "realtime modalities must contain text, audio, or both",
824        ));
825    }
826    if config.modalities.len() > 2
827        || (config.modalities.len() == 2 && config.modalities[0] == config.modalities[1])
828    {
829        return Err(protocol_error(
830            "realtime modalities must not contain duplicate values",
831        ));
832    }
833    let turn_detection = &config.turn_detection;
834    let has_server_vad_options = turn_detection.create_response.is_some()
835        || turn_detection.interrupt_response.is_some()
836        || turn_detection.prefix_padding_ms.is_some()
837        || turn_detection.silence_duration_ms.is_some()
838        || turn_detection.threshold.is_some();
839    if turn_detection.type_ == TurnDetectionType::ClientVad && has_server_vad_options {
840        return Err(protocol_error(
841            "realtime server-VAD options require turn_detection type server_vad",
842        ));
843    }
844    if let Some(threshold) = turn_detection.threshold
845        && (!threshold.is_finite() || !(0.0..=1.0).contains(&threshold))
846    {
847        return Err(protocol_error(
848            "realtime VAD threshold must be a finite value between 0 and 1",
849        ));
850    }
851    if config.beta_fields.chat_mode.is_none() {
852        return Err(protocol_error(
853            "realtime beta_fields.chat_mode is required when beta_fields is present",
854        ));
855    }
856    if config
857        .beta_fields
858        .tts_source
859        .as_deref()
860        .is_some_and(|source| source != "e2e")
861    {
862        return Err(protocol_error(
863            "unsupported realtime beta_fields.tts_source; the current protocol supports only \"e2e\"",
864        ));
865    }
866    if config.tools.iter().any(|tool| tool.type_ != "function") {
867        return Err(protocol_error("realtime tools must use type \"function\""));
868    }
869    if config
870        .tools
871        .iter()
872        .any(|tool| tool.name.trim().is_empty() || tool.description.trim().is_empty())
873    {
874        return Err(protocol_error(
875            "realtime tools require non-blank names and descriptions",
876        ));
877    }
878    if config.tools.iter().any(|tool| !tool.parameters.is_object()) {
879        return Err(protocol_error(
880            "realtime tool parameters must be a JSON Schema object",
881        ));
882    }
883    let mut tool_names = HashSet::with_capacity(config.tools.len());
884    if config
885        .tools
886        .iter()
887        .any(|tool| !tool_names.insert(tool.name.as_str()))
888    {
889        return Err(protocol_error("realtime tool names must be unique"));
890    }
891    if !config.tools.is_empty() && config.beta_fields.chat_mode != Some(ChatMode::Audio) {
892        return Err(protocol_error(
893            "realtime function tools are supported only in audio chat mode",
894        ));
895    }
896    if let Some(content) = config
897        .greeting_config
898        .as_ref()
899        .and_then(|greeting| greeting.content.as_deref())
900        && content.chars().count() > 1024
901    {
902        return Err(protocol_error(
903            "realtime greeting content must not exceed 1024 characters",
904        ));
905    }
906    Ok(())
907}
908
909fn serialize_event(event: &ClientEvent) -> ZaiResult<String> {
910    let json = serde_json::to_string(event)?;
911    if json.len() as u64 > WS_MESSAGE_MAX {
912        return Err(protocol_error(format!(
913            "realtime message exceeds {WS_MESSAGE_MAX} bytes"
914        )));
915    }
916    Ok(json)
917}
918
919fn protocol_error(message: impl Into<String>) -> crate::ZaiError {
920    RealtimeErrorKind::Protocol(message.into()).into()
921}
922
923fn now_ms() -> i64 {
924    chrono::Utc::now().timestamp_millis()
925}
926
927fn new_event_id() -> String {
928    format!("evt_{}", uuid::Uuid::new_v4().simple())
929}
930
931#[cfg(test)]
932mod teardown_tests {
933    use super::*;
934    use async_trait::async_trait;
935    use std::sync::Mutex;
936
937    /// A mock transport whose `recv` never resolves, so the only way the event
938    /// loop can exit is via the command branch — pinning the drop→teardown
939    /// invariant `RealtimeSession::close` depends on.
940    struct HangingTransport {
941        closed: Arc<Mutex<bool>>,
942    }
943
944    #[async_trait]
945    impl RealtimeTransport for HangingTransport {
946        async fn send(&mut self, _msg: String) -> crate::ZaiResult<()> {
947            Ok(())
948        }
949        async fn recv(&mut self) -> crate::ZaiResult<Option<WsMessage>> {
950            // Never resolves: forces the loop to exit via the command channel.
951            std::future::pending().await
952        }
953        async fn close(&mut self) -> crate::ZaiResult<()> {
954            *self.closed.lock().unwrap() = true;
955            Ok(())
956        }
957    }
958
959    /// Regression guard: dropping the last command `Sender` must terminate the
960    /// background loop AND close the transport. The consuming `close()` and the
961    /// implicit-drop teardown both rely on this; a future refactor that drops
962    /// the `None => close` arm would leak the task. Uses a 2s timeout so a
963    /// regression fails fast instead of hanging the test binary.
964    #[tokio::test]
965    async fn dropping_command_sender_terminates_loop_and_closes_transport() {
966        let closed = Arc::new(Mutex::new(false));
967        let transport = HangingTransport {
968            closed: Arc::clone(&closed),
969        };
970        let (cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
971        let (_shutdown_tx, shutdown_rx) = watch::channel(false);
972        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
973        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
974        let join = tokio::spawn(run_loop(
975            transport,
976            cmd_rx,
977            shutdown_rx,
978            events_tx,
979            audio_tx,
980        ));
981
982        // Drop the last sender → `cmd_rx.recv()` returns `None` → the loop
983        // calls `transport.close()` and exits.
984        drop(cmd_tx);
985
986        let joined = tokio::time::timeout(std::time::Duration::from_secs(2), join)
987            .await
988            .expect("run_loop did not terminate after the command sender dropped");
989        joined
990            .expect("run_loop task panicked")
991            .expect("run_loop returned an error");
992        assert!(
993            *closed.lock().unwrap(),
994            "transport.close() was not invoked on teardown"
995        );
996    }
997}
998
999#[cfg(test)]
1000mod run_loop_tests {
1001    use super::*;
1002    use async_trait::async_trait;
1003    use base64::Engine as _;
1004    use futures_util::StreamExt as _;
1005    use std::collections::VecDeque;
1006    use std::sync::{
1007        Arc, Mutex,
1008        atomic::{AtomicBool, AtomicUsize, Ordering},
1009    };
1010    use std::time::Duration;
1011
1012    /// Mock transport with scripted responses.
1013    struct ScriptedTransport {
1014        messages: VecDeque<String>,
1015        disconnect_when_empty: bool,
1016        sent: Arc<Mutex<Vec<String>>>,
1017        closed: Arc<Mutex<bool>>,
1018    }
1019
1020    impl ScriptedTransport {
1021        fn new(msgs: Vec<&str>) -> Self {
1022            Self {
1023                messages: msgs.into_iter().map(String::from).collect(),
1024                disconnect_when_empty: false,
1025                sent: Arc::new(Mutex::new(Vec::new())),
1026                closed: Arc::new(Mutex::new(false)),
1027            }
1028        }
1029
1030        fn disconnecting() -> Self {
1031            Self {
1032                disconnect_when_empty: true,
1033                ..Self::new(Vec::new())
1034            }
1035        }
1036    }
1037
1038    #[async_trait]
1039    impl RealtimeTransport for ScriptedTransport {
1040        async fn send(&mut self, msg: String) -> crate::ZaiResult<()> {
1041            self.sent.lock().unwrap().push(msg);
1042            Ok(())
1043        }
1044        async fn recv(&mut self) -> crate::ZaiResult<Option<WsMessage>> {
1045            match self.messages.pop_front() {
1046                Some(message) => Ok(Some(WsMessage::Text(message))),
1047                None if self.disconnect_when_empty => Ok(None),
1048                None => std::future::pending().await,
1049            }
1050        }
1051        async fn close(&mut self) -> crate::ZaiResult<()> {
1052            *self.closed.lock().unwrap() = true;
1053            Ok(())
1054        }
1055    }
1056
1057    struct FloodTransport {
1058        received: Arc<AtomicUsize>,
1059        sent: Arc<AtomicUsize>,
1060        closed: Arc<AtomicBool>,
1061    }
1062
1063    #[async_trait]
1064    impl RealtimeTransport for FloodTransport {
1065        async fn send(&mut self, _msg: String) -> crate::ZaiResult<()> {
1066            self.sent.fetch_add(1, Ordering::Relaxed);
1067            Ok(())
1068        }
1069
1070        async fn recv(&mut self) -> crate::ZaiResult<Option<WsMessage>> {
1071            self.received.fetch_add(1, Ordering::Relaxed);
1072            Ok(Some(WsMessage::Text(r#"{"type":"heartbeat"}"#.into())))
1073        }
1074
1075        async fn close(&mut self) -> crate::ZaiResult<()> {
1076            self.closed.store(true, Ordering::Relaxed);
1077            Ok(())
1078        }
1079    }
1080
1081    struct BoundaryHeartbeatTransport {
1082        delivered: bool,
1083        closed: Arc<AtomicBool>,
1084    }
1085
1086    #[async_trait]
1087    impl RealtimeTransport for BoundaryHeartbeatTransport {
1088        async fn send(&mut self, _msg: String) -> crate::ZaiResult<()> {
1089            Ok(())
1090        }
1091
1092        async fn recv(&mut self) -> crate::ZaiResult<Option<WsMessage>> {
1093            if self.delivered {
1094                return std::future::pending().await;
1095            }
1096            self.delivered = true;
1097            tokio::time::sleep(INBOUND_IDLE_TIMEOUT).await;
1098            Ok(Some(WsMessage::Text(r#"{"type":"heartbeat"}"#.into())))
1099        }
1100
1101        async fn close(&mut self) -> crate::ZaiResult<()> {
1102            self.closed.store(true, Ordering::Relaxed);
1103            Ok(())
1104        }
1105    }
1106
1107    async fn assert_loop_ok(join: JoinHandle<ZaiResult<()>>) {
1108        tokio::time::timeout(Duration::from_secs(2), join)
1109            .await
1110            .expect("run_loop timed out")
1111            .expect("run_loop task panicked")
1112            .expect("run_loop returned an error");
1113    }
1114
1115    fn spawn_test_loop<T>(
1116        transport: T,
1117        cmd_rx: mpsc::Receiver<String>,
1118        events_tx: broadcast::Sender<ServerEvent>,
1119        audio_tx: broadcast::Sender<RealtimeAudioChunk>,
1120    ) -> (watch::Sender<bool>, JoinHandle<ZaiResult<()>>)
1121    where
1122        T: RealtimeTransport + 'static,
1123    {
1124        let (shutdown_tx, shutdown_rx) = watch::channel(false);
1125        let join = tokio::spawn(run_loop(
1126            transport,
1127            cmd_rx,
1128            shutdown_rx,
1129            events_tx,
1130            audio_tx,
1131        ));
1132        (shutdown_tx, join)
1133    }
1134
1135    #[tokio::test]
1136    async fn run_loop_processes_server_events() {
1137        let transport = ScriptedTransport::new(vec![
1138            r#"{"type":"session.created","session":{"id":"s1"}}"#,
1139            r#"{"type":"session.updated","session":{"input_audio_format":"wav","output_audio_format":"pcm","turn_detection":{"type":"client_vad"}}}"#,
1140        ]);
1141        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1142        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1143        let mut events_rx = events_tx.subscribe();
1144        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1145        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1146        assert!(matches!(
1147            events_rx.recv().await,
1148            Ok(ServerEvent::SessionCreated { session })
1149                if session.id.as_deref() == Some("s1")
1150        ));
1151        assert!(matches!(
1152            events_rx.recv().await,
1153            Ok(ServerEvent::SessionUpdated { session })
1154                if session.input_audio_format == InputAudioFormat::Wav
1155        ));
1156        shutdown_tx.send(true).unwrap();
1157        assert_loop_ok(join).await;
1158    }
1159
1160    #[tokio::test]
1161    async fn run_loop_handles_audio_delta() {
1162        // Base64-encoded "hello"
1163        let audio_b64 = base64::engine::general_purpose::STANDARD.encode(b"hello");
1164        let json = format!(
1165            r#"{{"type":"response.audio.delta","response_id":"r1","item_id":"i1","delta":"{audio_b64}","event_id":"e1"}}"#
1166        );
1167        let transport = ScriptedTransport::new(vec![&json]);
1168        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1169        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1170        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1171        let mut audio_rx = audio_tx.subscribe();
1172        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1173        let chunk = audio_rx.recv().await.unwrap();
1174        assert_eq!(chunk.response_id, "r1");
1175        assert_eq!(chunk.item_id, "i1");
1176        assert_eq!(chunk.data, Bytes::from_static(b"hello"));
1177        shutdown_tx.send(true).unwrap();
1178        assert_loop_ok(join).await;
1179    }
1180
1181    #[tokio::test]
1182    async fn run_loop_handles_error_event() {
1183        let transport = ScriptedTransport::new(vec![
1184            r#"{"type":"error","error":{"type":"server_error","code":"server_error","message":"oops"}}"#,
1185        ]);
1186        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1187        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1188        let mut events_rx = events_tx.subscribe();
1189        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1190        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1191        assert!(matches!(
1192            events_rx.recv().await,
1193            Ok(ServerEvent::Error { .. })
1194        ));
1195        shutdown_tx.send(true).unwrap();
1196        assert_loop_ok(join).await;
1197    }
1198
1199    #[tokio::test]
1200    async fn run_loop_forwards_text_delta_and_done() {
1201        let transport = ScriptedTransport::new(vec![
1202            r#"{"type":"response.text.delta","response_id":"r1","item_id":"i1","delta":"hello "}"#,
1203            r#"{"type":"response.text.done","response_id":"r1","item_id":"i1","text":"hello world"}"#,
1204        ]);
1205        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1206        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1207        let mut events_rx = events_tx.subscribe();
1208        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1209        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1210
1211        assert!(matches!(
1212            events_rx.recv().await,
1213            Ok(ServerEvent::ResponseTextDelta { delta, .. }) if delta == "hello "
1214        ));
1215        assert!(matches!(
1216            events_rx.recv().await,
1217            Ok(ServerEvent::ResponseTextDone {
1218                text: Some(text),
1219                ..
1220            }) if text == "hello world"
1221        ));
1222
1223        shutdown_tx.send(true).unwrap();
1224        assert_loop_ok(join).await;
1225    }
1226
1227    #[tokio::test]
1228    async fn run_loop_keeps_both_directions_fair_under_flooding() {
1229        let received = Arc::new(AtomicUsize::new(0));
1230        let sent = Arc::new(AtomicUsize::new(0));
1231        let closed = Arc::new(AtomicBool::new(false));
1232        let transport = FloodTransport {
1233            received: Arc::clone(&received),
1234            sent: Arc::clone(&sent),
1235            closed: Arc::clone(&closed),
1236        };
1237        let (cmd_tx, cmd_rx) = mpsc::channel::<String>(64);
1238        for _ in 0..32 {
1239            cmd_tx
1240                .try_send("{}".into())
1241                .expect("command queue has capacity");
1242        }
1243        drop(cmd_tx);
1244        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1245        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1246
1247        let (_shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1248        assert_loop_ok(join).await;
1249
1250        assert_eq!(sent.load(Ordering::Relaxed), 32);
1251        assert!(received.load(Ordering::Relaxed) >= 32);
1252        assert!(closed.load(Ordering::Relaxed));
1253    }
1254
1255    #[tokio::test(start_paused = true)]
1256    async fn heartbeat_at_idle_boundary_wins_timeout_race() {
1257        let closed = Arc::new(AtomicBool::new(false));
1258        let transport = BoundaryHeartbeatTransport {
1259            delivered: false,
1260            closed: Arc::clone(&closed),
1261        };
1262        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1263        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1264        let mut events_rx = events_tx.subscribe();
1265        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1266        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1267
1268        tokio::task::yield_now().await;
1269        tokio::time::advance(INBOUND_IDLE_TIMEOUT).await;
1270        assert!(matches!(events_rx.recv().await, Ok(ServerEvent::Heartbeat)));
1271        assert!(
1272            !join.is_finished(),
1273            "boundary heartbeat caused a false timeout"
1274        );
1275
1276        shutdown_tx.send(true).unwrap();
1277        assert_loop_ok(join).await;
1278        assert!(closed.load(Ordering::Relaxed));
1279    }
1280
1281    #[tokio::test]
1282    async fn run_loop_sends_client_event() {
1283        let transport = ScriptedTransport::new(vec![]);
1284        let sent = Arc::clone(&transport.sent);
1285        let (cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1286        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1287        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1288        let (_shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1289        cmd_tx
1290            .send(
1291                serialize_event(&ClientEvent::ResponseCreate {
1292                    client_timestamp: None,
1293                })
1294                .unwrap(),
1295            )
1296            .await
1297            .unwrap();
1298        drop(cmd_tx);
1299        assert_loop_ok(join).await;
1300        assert_eq!(sent.lock().unwrap().len(), 1);
1301    }
1302
1303    #[tokio::test]
1304    async fn run_loop_shutdown_signal_terminates() {
1305        let transport = ScriptedTransport::new(vec![]);
1306        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1307        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1308        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1309        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1310        shutdown_tx.send(true).unwrap();
1311        assert_loop_ok(join).await;
1312    }
1313
1314    #[tokio::test]
1315    async fn run_loop_peer_disconnect_terminates() {
1316        // Empty message queue → recv returns None → peer disconnected
1317        let transport = ScriptedTransport::disconnecting();
1318        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1319        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1320        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1321        let (_shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1322        assert_loop_ok(join).await;
1323    }
1324
1325    #[tokio::test]
1326    async fn unknown_event_is_ignored_without_hiding_following_known_event() {
1327        let transport = ScriptedTransport::new(vec![
1328            r#"{"type":"future.event","payload":true}"#,
1329            r#"{"type":"heartbeat"}"#,
1330        ]);
1331        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1332        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1333        let mut events_rx = events_tx.subscribe();
1334        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1335        let (shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1336
1337        assert!(matches!(events_rx.recv().await, Ok(ServerEvent::Heartbeat)));
1338        shutdown_tx.send(true).unwrap();
1339        assert_loop_ok(join).await;
1340    }
1341
1342    #[tokio::test]
1343    async fn malformed_known_event_closes_session() {
1344        let transport = ScriptedTransport::new(vec![
1345            r#"{"type":"response.text.delta","response_id":"r1","delta":"missing item"}"#,
1346        ]);
1347        let closed = Arc::clone(&transport.closed);
1348        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1349        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1350        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1351        let (_shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1352        let error = join
1353            .await
1354            .expect("run_loop task panicked")
1355            .expect_err("malformed known event was silently ignored");
1356
1357        assert!(error.message().contains("malformed realtime server event"));
1358        assert!(*closed.lock().unwrap());
1359    }
1360
1361    #[tokio::test(start_paused = true)]
1362    async fn run_loop_closes_half_open_session_after_missed_heartbeats() {
1363        let transport = ScriptedTransport::new(vec![]);
1364        let closed = Arc::clone(&transport.closed);
1365        // Keep the sender alive so only the inbound idle deadline can stop the
1366        // loop; this models a half-open connection with a healthy local task.
1367        let (_cmd_tx, cmd_rx) = mpsc::channel::<String>(8);
1368        let (events_tx, _) = broadcast::channel::<ServerEvent>(16);
1369        let (audio_tx, _) = broadcast::channel::<RealtimeAudioChunk>(16);
1370        let (_shutdown_tx, join) = spawn_test_loop(transport, cmd_rx, events_tx, audio_tx);
1371
1372        tokio::task::yield_now().await;
1373        tokio::time::advance(INBOUND_IDLE_TIMEOUT).await;
1374        let error = join
1375            .await
1376            .expect("run_loop task panicked")
1377            .expect_err("half-open session did not time out");
1378
1379        assert!(error.message().contains("inbound heartbeat timed out"));
1380        assert!(*closed.lock().unwrap());
1381    }
1382
1383    #[test]
1384    fn new_event_id_format() {
1385        let first = new_event_id();
1386        let second = new_event_id();
1387        assert!(first.starts_with("evt_"));
1388        assert_eq!(first.len(), 36);
1389        assert_ne!(first, second);
1390    }
1391
1392    #[test]
1393    fn session_queues_match_frozen_contract_capacity() {
1394        assert_eq!(SESSION_CHANNEL_CAPACITY, 8);
1395    }
1396
1397    #[test]
1398    fn oversized_event_is_rejected_before_enqueue() {
1399        let event = ClientEvent::ConversationItemCreate {
1400            event_id: None,
1401            item: super::super::protocol::RealtimeConversationItem::user_text(
1402                "x".repeat(WS_MESSAGE_MAX as usize),
1403            ),
1404        };
1405        assert!(serialize_event(&event).is_err());
1406    }
1407
1408    #[test]
1409    fn session_config_validates_numeric_limits() {
1410        let mut config = SessionConfig {
1411            temperature: Some(f64::NAN),
1412            ..SessionConfig::default()
1413        };
1414        assert!(validate_session_config(&config).is_err());
1415
1416        config.temperature = Some(0.5);
1417        config.max_response_output_tokens = Some(0);
1418        assert!(validate_session_config(&config).is_err());
1419
1420        config.max_response_output_tokens = Some(1025);
1421        assert!(validate_session_config(&config).is_err());
1422
1423        config.max_response_output_tokens = Some(1024);
1424        assert!(validate_session_config(&config).is_ok());
1425    }
1426
1427    #[test]
1428    fn session_config_validates_modalities_vad_and_tools() {
1429        let mut config = SessionConfig {
1430            modalities: Vec::new(),
1431            ..SessionConfig::default()
1432        };
1433        assert!(validate_session_config(&config).is_err());
1434
1435        config.modalities = vec![RealtimeModality::Text, RealtimeModality::Text];
1436        assert!(validate_session_config(&config).is_err());
1437
1438        config.modalities = vec![RealtimeModality::Text, RealtimeModality::Audio];
1439        config.turn_detection.create_response = Some(true);
1440        assert!(validate_session_config(&config).is_err());
1441
1442        config.turn_detection.type_ = TurnDetectionType::ServerVad;
1443        config.turn_detection.threshold = Some(f64::NAN);
1444        assert!(validate_session_config(&config).is_err());
1445
1446        config.turn_detection.threshold = Some(0.5);
1447        config.tools = vec![RealtimeTool::function(
1448            "weather",
1449            "Get weather",
1450            serde_json::json!({"type": "object"}),
1451        )];
1452        config.beta_fields.chat_mode = Some(ChatMode::VideoPassive);
1453        assert!(validate_session_config(&config).is_err());
1454
1455        config.beta_fields.chat_mode = Some(ChatMode::Audio);
1456        config.tools.push(config.tools[0].clone());
1457        assert!(validate_session_config(&config).is_err());
1458
1459        config.tools.pop();
1460        assert!(validate_session_config(&config).is_ok());
1461    }
1462
1463    #[tokio::test]
1464    async fn observable_stream_reports_lag_and_background_failure() {
1465        let (events_tx, events_rx) = broadcast::channel(2);
1466        let (_completion_tx, completion_rx) = watch::channel(None);
1467        let mut events = observable_broadcast_stream(events_rx, completion_rx, "test events");
1468        for value in 0..3 {
1469            events_tx.send(value).unwrap();
1470        }
1471        let lag = events
1472            .next()
1473            .await
1474            .expect("stream ended")
1475            .expect_err("lag was silently discarded");
1476        assert!(lag.message().contains("lost 1 message"));
1477        assert!(
1478            events.next().await.is_none(),
1479            "a corrupted stream continued after reporting lag"
1480        );
1481
1482        let (_events_tx, events_rx) = broadcast::channel::<u8>(2);
1483        let (completion_tx, completion_rx) = watch::channel(None);
1484        let mut events = observable_broadcast_stream(events_rx, completion_rx, "test events");
1485        completion_tx.send_replace(Some(Err(protocol_error("background failed"))));
1486        let failure = events
1487            .next()
1488            .await
1489            .expect("stream ended before reporting failure")
1490            .expect_err("background failure was hidden");
1491        assert!(failure.message().contains("background failed"));
1492        assert!(events.next().await.is_none());
1493
1494        let (events_tx, events_rx) = broadcast::channel::<u8>(2);
1495        let (completion_tx, completion_rx) = watch::channel(None);
1496        let mut events = observable_broadcast_stream(events_rx, completion_rx, "test events");
1497        events_tx.send(9).unwrap();
1498        drop(completion_tx);
1499        assert_eq!(events.next().await.unwrap().unwrap(), 9);
1500        let failure = events
1501            .next()
1502            .await
1503            .expect("stream ended before reporting task loss")
1504            .expect_err("missing completion status was hidden");
1505        assert!(failure.message().contains("without a completion status"));
1506    }
1507
1508    #[test]
1509    fn first_subscription_keeps_pre_subscription_backlog() {
1510        let (events_tx, initial_rx) = broadcast::channel(SESSION_CHANNEL_CAPACITY);
1511        let initial = Mutex::new(Some(initial_rx));
1512        events_tx.send(7_u8).unwrap();
1513
1514        let mut first = subscribe_with_initial_backlog(&events_tx, &initial);
1515        assert_eq!(first.try_recv().unwrap(), 7);
1516
1517        let mut second = subscribe_with_initial_backlog(&events_tx, &initial);
1518        assert!(matches!(
1519            second.try_recv(),
1520            Err(broadcast::error::TryRecvError::Empty)
1521        ));
1522    }
1523}