1use 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
47const INBOUND_IDLE_TIMEOUT: Duration = Duration::from_secs(90);
50const SESSION_CHANNEL_CAPACITY: usize = 8;
54
55#[derive(Debug, Clone)]
57pub struct RealtimeAudioChunk {
58 pub response_id: String,
60 pub item_id: String,
62 pub output_index: Option<u64>,
64 pub content_index: Option<u64>,
66 pub data: Bytes,
68}
69
70pub 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 pub fn instructions(mut self, instructions: impl Into<String>) -> Self {
101 self.session_config.instructions = Some(instructions.into());
102 self
103 }
104
105 pub fn turn_detection(mut self, vad: TurnDetectionType) -> Self {
107 self.session_config.turn_detection.type_ = vad;
108 self
109 }
110
111 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 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 pub fn vad_threshold(mut self, threshold: f64) -> Self {
127 self.session_config.turn_detection.threshold = Some(threshold);
128 self
129 }
130
131 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 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 pub fn input_audio_format(mut self, format: InputAudioFormat) -> Self {
145 self.session_config.input_audio_format = format;
146 self
147 }
148
149 pub fn output_audio_format(mut self, format: OutputAudioFormat) -> Self {
151 self.session_config.output_audio_format = format;
152 self
153 }
154
155 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 pub fn voice(mut self, voice: RealtimeVoice) -> Self {
163 self.session_config.voice = Some(voice);
164 self
165 }
166
167 pub fn temperature(mut self, temperature: f64) -> Self {
170 self.session_config.temperature = Some(temperature);
171 self
172 }
173
174 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 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 pub fn chat_mode(mut self, mode: ChatMode) -> Self {
190 self.session_config.beta_fields.chat_mode = Some(mode);
191 self
192 }
193
194 pub fn auto_search(mut self, enabled: bool) -> Self {
196 self.session_config.beta_fields.auto_search = Some(enabled);
197 self
198 }
199
200 pub fn greeting_config(mut self, greeting: GreetingConfig) -> Self {
202 self.session_config.greeting_config = Some(greeting);
203 self
204 }
205
206 pub fn tools(mut self, tools: Vec<RealtimeTool>) -> Self {
208 self.session_config.tools = tools;
209 self
210 }
211
212 pub fn session_config(mut self, config: SessionConfig) -> Self {
217 self.session_config = config;
218 self
219 }
220
221 #[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 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 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 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
301async 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 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 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 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
463pub 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 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 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 pub async fn commit_audio(&self) -> ZaiResult<()> {
533 self.dispatch(ClientEvent::InputAudioBufferCommit {
534 client_timestamp: Some(now_ms()),
535 })
536 .await
537 }
538
539 pub async fn clear_audio(&self) -> ZaiResult<()> {
541 self.dispatch(ClientEvent::InputAudioBufferClear).await
542 }
543
544 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 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 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 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 pub async fn create_response(&self) -> ZaiResult<()> {
609 self.dispatch(ClientEvent::ResponseCreate {
610 client_timestamp: Some(now_ms()),
611 })
612 .await
613 }
614
615 pub async fn cancel(&self) -> ZaiResult<()> {
617 self.dispatch(ClientEvent::ResponseCancel {
618 client_timestamp: Some(now_ms()),
619 })
620 .await
621 }
622
623 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 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 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 pub async fn request_close(&self) -> ZaiResult<()> {
674 self.shutdown_tx
675 .send(true)
676 .map_err(|_| RealtimeErrorKind::Closed.into())
677 }
678
679 pub async fn close(self) -> ZaiResult<()> {
681 let _ = self.shutdown_tx.send(true);
684 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 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 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 #[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(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 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 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 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 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}