1use crate::error::{Result, VoxRtcError};
2use crate::socket::RawSocketChannel;
3use crate::types::*;
4use serde_json::Value;
5use std::ops::ControlFlow;
6use std::sync::{Arc, Mutex};
7use tokio::sync::broadcast::error::RecvError;
8use uuid::Uuid;
9use tokio::task::JoinHandle;
10use tokio::time::{Duration, timeout};
11
12#[derive(Clone)]
13pub struct VoxRtcControlSession {
14 channel: RawSocketChannel,
15 session_id: String,
16 channel_name: String,
17 join_timeout: Duration,
18 response_generation: Arc<Mutex<ResponseGeneration>>,
19}
20
21#[derive(Default)]
22struct ResponseGeneration {
23 counter: u64,
24 id: Option<String>,
25}
26
27pub struct Listener {
28 handle: JoinHandle<()>,
29}
30
31impl Drop for Listener {
32 fn drop(&mut self) {
33 self.handle.abort();
34 }
35}
36
37impl VoxRtcControlSession {
38 pub(crate) fn new(
39 channel: RawSocketChannel,
40 session_id: String,
41 join_timeout: Duration,
42 ) -> Self {
43 let channel_name = format!("/rtc/{session_id}");
44 Self {
45 channel,
46 session_id,
47 channel_name,
48 join_timeout,
49 response_generation: Arc::new(Mutex::new(ResponseGeneration::default())),
50 }
51 }
52
53 pub fn session_id(&self) -> &str {
54 &self.session_id
55 }
56
57 pub fn channel_name(&self) -> &str {
58 &self.channel_name
59 }
60
61 pub async fn join(&self) -> Result<()> {
62 let mut states = self.channel.subscribe_state();
63 self.channel.join().await?;
64 let channel_name = self.channel.name().to_owned();
65 let channel = self.channel.clone();
66 timeout(self.join_timeout, async move {
67 loop {
68 let state = *states.borrow_and_update();
69 match state {
70 ChannelState::Joined => return Ok(()),
71 ChannelState::Closed | ChannelState::Declined => {
72 let reason = join_decline_reason(channel.decline_reason().await);
73 return Err(VoxRtcError::JoinFailed {
74 channel: channel_name,
75 state: format!("{state:?}"),
76 reason,
77 });
78 }
79 _ => {}
80 }
81 if states.changed().await.is_err() {
82 return Err(VoxRtcError::Disconnected);
83 }
84 }
85 })
86 .await
87 .map_err(|_| VoxRtcError::JoinTimeout(self.channel_name.clone()))?
88 }
89
90 pub async fn close(&self) -> Result<()> {
91 self.channel.leave().await
92 }
93
94 pub fn on_event<F>(&self, handler: F) -> Listener
95 where
96 F: Fn(WireEvent) + Send + Sync + 'static,
97 {
98 let mut messages = self.channel.subscribe_messages();
99 let session_id = self.session_id.clone();
100 let channel_name = self.channel_name.clone();
101 Listener {
102 handle: tokio::spawn(async move {
103 loop {
104 match next_message(messages.recv().await) {
105 ControlFlow::Break(()) => break,
106 ControlFlow::Continue(None) => continue,
107 ControlFlow::Continue(Some((event, payload))) => handler(WireEvent {
108 r#type: event,
109 data: payload,
110 session_id: session_id.clone(),
111 channel_name: channel_name.clone(),
112 }),
113 }
114 }
115 }),
116 }
117 }
118
119 pub fn on<F>(&self, event_name: impl Into<String>, handler: F) -> Listener
120 where
121 F: Fn(EventData) + Send + Sync + 'static,
122 {
123 let event_name = event_name.into();
124 let mut messages = self.channel.subscribe_messages();
125 Listener {
126 handle: tokio::spawn(async move {
127 loop {
128 match next_message(messages.recv().await) {
129 ControlFlow::Break(()) => break,
130 ControlFlow::Continue(None) => continue,
131 ControlFlow::Continue(Some((event, payload))) => {
132 if event == event_name {
133 handler(payload);
134 }
135 }
136 }
137 }
138 }),
139 }
140 }
141
142 pub fn on_session_attached<F>(&self, handler: F) -> Listener
143 where
144 F: Fn(SessionAttachedEvent) + Send + Sync + 'static,
145 {
146 let session_id = self.session_id.clone();
147 let channel_name = self.channel_name.clone();
148 self.on(EVENT_RTC_SESSION_ATTACHED, move |payload| {
149 handler(SessionAttachedEvent {
150 session_id: base_session_id(&payload, &session_id),
151 channel_name: channel_name.clone(),
152 data: payload,
153 })
154 })
155 }
156
157 pub fn on_session_created<F>(&self, handler: F) -> Listener
158 where
159 F: Fn(SessionCreatedEvent) + Send + Sync + 'static,
160 {
161 let session_id = self.session_id.clone();
162 let channel_name = self.channel_name.clone();
163 self.on(EVENT_SESSION_CREATED, move |payload| {
164 let session = payload.get("session").and_then(Value::as_object).cloned();
165 handler(SessionCreatedEvent {
166 session_id: base_session_id(&payload, &session_id),
167 channel_name: channel_name.clone(),
168 data: payload,
169 session,
170 });
171 })
172 }
173
174 pub fn on_transcript<F>(&self, handler: F) -> Listener
175 where
176 F: Fn(TranscriptEvent) + Send + Sync + 'static,
177 {
178 let session_id = self.session_id.clone();
179 let channel_name = self.channel_name.clone();
180 self.on(EVENT_TRANSCRIPT_COMPLETED, move |payload| {
181 handler(TranscriptEvent {
182 session_id: base_session_id(&payload, &session_id),
183 channel_name: channel_name.clone(),
184 transcript: required_string(&payload, "transcript", ""),
185 language: optional_string(&payload, "language"),
186 start_ms: optional_number(&payload, "start_ms"),
187 end_ms: optional_number(&payload, "end_ms"),
188 eou_probability: optional_number(&payload, "eou_probability"),
189 topics: optional_string_vec(&payload, "topics"),
190 data: payload,
191 });
192 })
193 }
194
195 pub fn on_turn_state_changed<F>(&self, handler: F) -> Listener
196 where
197 F: Fn(TurnStateEvent) + Send + Sync + 'static,
198 {
199 let session_id = self.session_id.clone();
200 let channel_name = self.channel_name.clone();
201 self.on(EVENT_TURN_STATE_CHANGED, move |payload| {
202 handler(TurnStateEvent {
203 session_id: base_session_id(&payload, &session_id),
204 channel_name: channel_name.clone(),
205 state: required_string(&payload, "state", "unknown"),
206 previous_state: optional_string(&payload, "previous_state"),
207 data: payload,
208 });
209 })
210 }
211
212 pub fn on_speech_started<F>(&self, handler: F) -> Listener
213 where
214 F: Fn(SpeechStartedEvent) + Send + Sync + 'static,
215 {
216 let session_id = self.session_id.clone();
217 let channel_name = self.channel_name.clone();
218 self.on(EVENT_SPEECH_STARTED, move |payload| {
219 handler(SpeechStartedEvent {
220 session_id: base_session_id(&payload, &session_id),
221 channel_name: channel_name.clone(),
222 timestamp_ms: optional_number(&payload, "timestamp_ms"),
223 data: payload,
224 });
225 })
226 }
227
228 pub fn on_speech_stopped<F>(&self, handler: F) -> Listener
229 where
230 F: Fn(SpeechStoppedEvent) + Send + Sync + 'static,
231 {
232 let session_id = self.session_id.clone();
233 let channel_name = self.channel_name.clone();
234 self.on(EVENT_SPEECH_STOPPED, move |payload| {
235 handler(SpeechStoppedEvent {
236 session_id: base_session_id(&payload, &session_id),
237 channel_name: channel_name.clone(),
238 timestamp_ms: optional_number(&payload, "timestamp_ms"),
239 data: payload,
240 });
241 })
242 }
243
244 pub fn on_transcript_delta<F>(&self, handler: F) -> Listener
245 where
246 F: Fn(TranscriptDeltaEvent) + Send + Sync + 'static,
247 {
248 let session_id = self.session_id.clone();
249 let channel_name = self.channel_name.clone();
250 self.on(EVENT_TRANSCRIPT_DELTA, move |payload| {
251 handler(TranscriptDeltaEvent {
252 session_id: base_session_id(&payload, &session_id),
253 channel_name: channel_name.clone(),
254 delta: required_string(&payload, "delta", ""),
255 start_ms: optional_number(&payload, "start_ms"),
256 end_ms: optional_number(&payload, "end_ms"),
257 data: payload,
258 });
259 })
260 }
261
262 pub fn on_turn_eou_predicted<F>(&self, handler: F) -> Listener
263 where
264 F: Fn(TurnEouPredictedEvent) + Send + Sync + 'static,
265 {
266 let session_id = self.session_id.clone();
267 let channel_name = self.channel_name.clone();
268 self.on(EVENT_TURN_EOU_PREDICTED, move |payload| {
269 handler(TurnEouPredictedEvent {
270 session_id: base_session_id(&payload, &session_id),
271 channel_name: channel_name.clone(),
272 probability: optional_number(&payload, "probability"),
273 threshold: optional_number(&payload, "threshold"),
274 delay_ms: optional_number(&payload, "delay_ms"),
275 start_ms: optional_number(&payload, "start_ms"),
276 end_ms: optional_number(&payload, "end_ms"),
277 decision: optional_string(&payload, "decision"),
278 action: optional_string(&payload, "action"),
279 turn_detector: optional_string(&payload, "turn_detector"),
280 data: payload,
281 });
282 })
283 }
284
285 pub fn on_response_created<F>(&self, handler: F) -> Listener
286 where
287 F: Fn(ResponseEvent) + Send + Sync + 'static,
288 {
289 self.on_response_event(EVENT_RESPONSE_CREATED, handler)
290 }
291
292 pub fn on_response_committed<F>(&self, handler: F) -> Listener
293 where
294 F: Fn(ResponseEvent) + Send + Sync + 'static,
295 {
296 self.on_response_event(EVENT_RESPONSE_COMMITTED, handler)
297 }
298
299 pub fn on_response_done<F>(&self, handler: F) -> Listener
300 where
301 F: Fn(ResponseEvent) + Send + Sync + 'static,
302 {
303 self.on_response_event(EVENT_RESPONSE_DONE, handler)
304 }
305
306 pub fn on_response_cancelled<F>(&self, handler: F) -> Listener
307 where
308 F: Fn(ResponseEvent) + Send + Sync + 'static,
309 {
310 self.on_response_event(EVENT_RESPONSE_CANCELLED, handler)
311 }
312
313 pub fn on_response_audio_clear<F>(&self, handler: F) -> Listener
314 where
315 F: Fn(ResponseEvent) + Send + Sync + 'static,
316 {
317 self.on_response_event(EVENT_RESPONSE_AUDIO_CLEAR, handler)
318 }
319
320 fn on_response_event<F>(&self, event_name: &'static str, handler: F) -> Listener
321 where
322 F: Fn(ResponseEvent) + Send + Sync + 'static,
323 {
324 let session_id = self.session_id.clone();
325 let channel_name = self.channel_name.clone();
326 self.on(event_name, move |payload| {
327 handler(response_event(payload, &session_id, &channel_name));
328 })
329 }
330
331 pub fn on_interruption_detected<F>(&self, handler: F) -> Listener
332 where
333 F: Fn(InterruptionEvent) + Send + Sync + 'static,
334 {
335 self.on_interruption_event(EVENT_INTERRUPTION_DETECTED, handler)
336 }
337
338 pub fn on_interruption_false_positive<F>(&self, handler: F) -> Listener
339 where
340 F: Fn(InterruptionEvent) + Send + Sync + 'static,
341 {
342 self.on_interruption_event(EVENT_INTERRUPTION_FALSE_POSITIVE, handler)
343 }
344
345 fn on_interruption_event<F>(&self, event_name: &'static str, handler: F) -> Listener
346 where
347 F: Fn(InterruptionEvent) + Send + Sync + 'static,
348 {
349 let session_id = self.session_id.clone();
350 let channel_name = self.channel_name.clone();
351 self.on(event_name, move |payload| {
352 handler(InterruptionEvent {
353 response: response_event(payload.clone(), &session_id, &channel_name),
354 vad_active_ms: optional_number(&payload, "vad_active_ms"),
355 partial_transcript: optional_string(&payload, "partial_transcript"),
356 });
357 })
358 }
359
360 pub fn on_browser_event<F>(&self, handler: F) -> Listener
361 where
362 F: Fn(BrowserEvent) + Send + Sync + 'static,
363 {
364 let session_id = self.session_id.clone();
365 let channel_name = self.channel_name.clone();
366 self.on(EVENT_BROWSER_EVENT, move |payload| {
367 handler(BrowserEvent {
368 session_id: base_session_id(&payload, &session_id),
369 channel_name: channel_name.clone(),
370 event: required_string(&payload, "event", ""),
371 payload: payload.get("payload").cloned().unwrap_or(Value::Null),
372 data: payload,
373 });
374 })
375 }
376
377 pub fn on_close<F>(&self, handler: F) -> Listener
378 where
379 F: Fn(CloseEvent) + Send + Sync + 'static,
380 {
381 let session_id = self.session_id.clone();
382 let channel_name = self.channel_name.clone();
383 self.on(EVENT_RTC_CLIENT_DISCONNECTED, move |payload| {
384 handler(CloseEvent {
385 session_id: base_session_id(&payload, &session_id),
386 channel_name: channel_name.clone(),
387 reason: required_string(&payload, "reason", "unknown"),
388 connection_state: optional_string(&payload, "connection_state"),
389 ice_connection_state: optional_string(&payload, "ice_connection_state"),
390 data_channel_state: optional_string(&payload, "data_channel_state"),
391 data: payload,
392 });
393 })
394 }
395
396 pub fn on_error<F>(&self, handler: F) -> Listener
397 where
398 F: Fn(ErrorEvent) + Send + Sync + 'static,
399 {
400 let session_id = self.session_id.clone();
401 let channel_name = self.channel_name.clone();
402 self.on(EVENT_ERROR, move |payload| {
403 handler(ErrorEvent {
404 session_id: base_session_id(&payload, &session_id),
405 channel_name: channel_name.clone(),
406 message: optional_string(&payload, "message"),
407 code: optional_nonempty_string(&payload, "code"),
408 recoverable: recoverable_flag(&payload),
409 generation_id: optional_nonempty_string(&payload, "generation_id"),
410 data: payload,
411 });
412 })
413 }
414
415 pub async fn send_control(&self, event: &str, payload: EventData) -> Result<()> {
416 self.channel.send_message(event, payload).await
417 }
418
419 pub async fn configure(&self, config: SessionConfig) -> Result<()> {
420 let mut session = config.extra;
421 insert_opt(&mut session, "stt_model", config.stt_model);
422 insert_opt(&mut session, "tts_model", config.tts_model);
423 insert_opt(&mut session, "voice", config.voice);
424 insert_opt(&mut session, "turn_profile", config.turn_profile);
425 insert_opt(&mut session, "vad_backend", config.vad_backend);
426 insert_opt(&mut session, "turn_detector", config.turn_detector);
427
428 let mut payload = EventData::new();
429 payload.insert("session".to_owned(), Value::Object(session));
430 self.send_control("session.update", payload).await
431 }
432
433 pub async fn start_response(&self, options: Option<ResponseOptions>) -> Result<()> {
434 let (_, payload) = self.start_payload(options);
435 self.send_control("response.start", payload).await
436 }
437
438 pub async fn start_response_and_wait(
439 &self,
440 options: Option<ResponseOptions>,
441 wait_timeout: Duration,
442 ) -> Result<StartAck> {
443 let (generation_id, payload) = self.start_payload(options);
444 let mut messages = self.channel.subscribe_messages();
445 self.send_control("response.start", payload).await?;
446 timeout(wait_timeout, async move {
447 loop {
448 match next_message(messages.recv().await) {
449 ControlFlow::Break(()) => return Err(VoxRtcError::ChannelClosed),
450 ControlFlow::Continue(None) => continue,
451 ControlFlow::Continue(Some((event, data))) => {
452 if optional_nonempty_string(&data, "generation_id").as_deref()
453 != Some(generation_id.as_str())
454 {
455 continue;
456 }
457 if event == EVENT_RESPONSE_CREATED {
458 return Ok(StartAck {
459 accepted: true,
460 generation_id: generation_id.clone(),
461 response_id: optional_string(&data, "response_id"),
462 error_code: None,
463 error_message: None,
464 recoverable: true,
465 });
466 }
467 if event == EVENT_ERROR {
468 return Ok(StartAck {
469 accepted: false,
470 generation_id: generation_id.clone(),
471 response_id: optional_string(&data, "response_id"),
472 error_code: optional_nonempty_string(&data, "code"),
473 error_message: optional_string(&data, "message"),
474 recoverable: recoverable_flag(&data),
475 });
476 }
477 }
478 }
479 }
480 })
481 .await
482 .map_err(|_| VoxRtcError::Timeout("response.start acknowledgement"))?
483 }
484
485 pub async fn append_response_text(
486 &self,
487 delta: impl Into<String>,
488 options: Option<ResponseOptions>,
489 ) -> Result<()> {
490 let explicit = explicit_generation(&options);
491 let mut payload = response_options_payload(options);
492 payload.insert("delta".to_owned(), Value::String(delta.into()));
493 self.thread_generation(&mut payload, explicit);
494 self.send_control("response.delta", payload).await
495 }
496
497 pub async fn commit_response(&self, options: Option<ResponseOptions>) -> Result<()> {
498 let explicit = explicit_generation(&options);
499 let mut payload = EventData::new();
500 self.thread_generation(&mut payload, explicit);
501 self.send_control("response.commit", payload).await
502 }
503
504 pub async fn cancel_response(&self, options: Option<ResponseOptions>) -> Result<()> {
505 let explicit = explicit_generation(&options);
506 let mut payload = EventData::new();
507 self.thread_generation(&mut payload, explicit);
508 self.clear_response_generation();
509 self.send_control("response.cancel", payload).await
510 }
511
512 pub async fn replace_response_text(
513 &self,
514 text: impl Into<String>,
515 options: Option<ResponseOptions>,
516 ) -> Result<()> {
517 self.clear_response_generation();
518 let explicit = explicit_generation(&options);
519 let mut payload = response_options_payload(options);
520 payload.insert("text".to_owned(), Value::String(text.into()));
521 if let Some(generation_id) = explicit {
522 payload.insert("generation_id".to_owned(), Value::String(generation_id));
523 }
524 self.send_control("response.replace_text", payload).await
525 }
526
527 pub async fn send_text_response(
528 &self,
529 text: impl Into<String>,
530 options: Option<ResponseOptions>,
531 cancel_first: bool,
532 ) -> Result<()> {
533 let text = text.into();
534 if cancel_first {
535 return self.replace_response_text(text, options).await;
536 }
537 self.start_response(options.clone()).await?;
538 self.append_response_text(text, options.clone()).await?;
539 self.commit_response(options).await
540 }
541
542 pub async fn send_client_event(&self, envelope: ClientEventEnvelope) -> Result<()> {
543 let mut payload = EventData::new();
544 payload.insert("event".to_owned(), Value::String(envelope.event));
545 payload.insert("payload".to_owned(), envelope.payload);
546 self.send_control(EVENT_CLIENT_EVENT, payload).await
547 }
548
549 fn start_payload(&self, options: Option<ResponseOptions>) -> (String, EventData) {
550 let explicit = explicit_generation(&options);
551 let mut payload = response_options_payload(options);
552 let generation_id = match explicit {
553 Some(id) => self.set_response_generation(id),
554 None => self.next_response_generation(),
555 };
556 payload.insert(
557 "generation_id".to_owned(),
558 Value::String(generation_id.clone()),
559 );
560 (generation_id, payload)
561 }
562
563 fn next_response_generation(&self) -> String {
564 let mut state = self
565 .response_generation
566 .lock()
567 .expect("response generation mutex poisoned");
568 state.counter += 1;
569 let generation_id = format!("generation_{}_{}", state.counter, Uuid::new_v4());
570 state.id = Some(generation_id.clone());
571 generation_id
572 }
573
574 fn set_response_generation(&self, generation_id: String) -> String {
575 let mut state = self
576 .response_generation
577 .lock()
578 .expect("response generation mutex poisoned");
579 state.counter += 1;
580 state.id = Some(generation_id.clone());
581 generation_id
582 }
583
584 fn thread_generation(&self, payload: &mut EventData, explicit: Option<String>) {
585 match explicit {
586 Some(generation_id) => {
587 payload.insert("generation_id".to_owned(), Value::String(generation_id));
588 }
589 None => self.add_response_generation(payload),
590 }
591 }
592
593 fn add_response_generation(&self, payload: &mut EventData) {
594 let state = self
595 .response_generation
596 .lock()
597 .expect("response generation mutex poisoned");
598 if let Some(generation_id) = &state.id {
599 payload.insert(
600 "generation_id".to_owned(),
601 Value::String(generation_id.clone()),
602 );
603 }
604 }
605
606 fn clear_response_generation(&self) {
607 self.response_generation
608 .lock()
609 .expect("response generation mutex poisoned")
610 .id = None;
611 }
612}
613
614fn next_message(
615 result: std::result::Result<(String, EventData), RecvError>,
616) -> ControlFlow<(), Option<(String, EventData)>> {
617 match result {
618 Ok(message) => ControlFlow::Continue(Some(message)),
619 Err(RecvError::Lagged(_)) => ControlFlow::Continue(None),
620 Err(RecvError::Closed) => ControlFlow::Break(()),
621 }
622}
623
624fn join_decline_reason(reason: Option<EventData>) -> Option<String> {
625 let reason = reason?;
626 for key in ["message", "reason", "error"] {
627 if let Some(value) = reason.get(key).and_then(Value::as_str)
628 && !value.is_empty()
629 {
630 return Some(value.to_owned());
631 }
632 }
633 if reason.is_empty() {
634 None
635 } else {
636 Some(Value::Object(reason).to_string())
637 }
638}
639
640fn insert_opt(session: &mut EventData, key: &str, value: Option<String>) {
641 if let Some(value) = value {
642 session.insert(key.to_owned(), Value::String(value));
643 }
644}
645
646fn response_options_payload(options: Option<ResponseOptions>) -> EventData {
647 let mut payload = EventData::new();
648 if let Some(options) = options
649 && let Some(allow) = options.allow_interruptions
650 {
651 payload.insert("allow_interruptions".to_owned(), Value::Bool(allow));
652 }
653 payload
654}
655
656fn explicit_generation(options: &Option<ResponseOptions>) -> Option<String> {
657 options
658 .as_ref()
659 .and_then(|options| options.generation_id.clone())
660 .filter(|id| !id.is_empty())
661}
662
663fn base_session_id(payload: &EventData, fallback: &str) -> String {
664 required_string(payload, "session_id", fallback)
665}
666
667fn response_event(payload: EventData, session_id: &str, channel_name: &str) -> ResponseEvent {
668 ResponseEvent {
669 session_id: base_session_id(&payload, session_id),
670 channel_name: channel_name.to_owned(),
671 response_id: optional_string(&payload, "response_id"),
672 generation_id: optional_nonempty_string(&payload, "generation_id"),
673 data: payload,
674 }
675}
676
677#[cfg(test)]
678mod tests {
679 use super::*;
680 use crate::socket::test_channel;
681 use serde_json::json;
682 use tokio::sync::broadcast;
683 use tokio::sync::mpsc;
684
685 async fn session() -> (VoxRtcControlSession, broadcast::Sender<(String, EventData)>) {
686 let (channel, sender) = test_channel().await;
687 let session =
688 VoxRtcControlSession::new(channel, "sess-1".to_owned(), Duration::from_secs(1));
689 (session, sender)
690 }
691
692 fn payload(value: Value) -> EventData {
693 value.as_object().cloned().expect("object payload")
694 }
695
696 #[test]
697 fn join_decline_reason_prefers_structured_message_fields() {
698 assert_eq!(
699 join_decline_reason(Some(payload(json!({ "message": "expired" })))),
700 Some("expired".to_owned())
701 );
702 assert_eq!(
703 join_decline_reason(Some(payload(json!({ "reason": "missing" })))),
704 Some("missing".to_owned())
705 );
706 assert_eq!(
707 join_decline_reason(Some(payload(json!({ "channel": "/rtc/abc" })))),
708 Some(r#"{"channel":"/rtc/abc"}"#.to_owned())
709 );
710 assert_eq!(join_decline_reason(Some(EventData::new())), None);
711 assert_eq!(join_decline_reason(None), None);
712 }
713
714 async fn recv<T>(rx: &mut mpsc::UnboundedReceiver<T>) -> T {
715 timeout(Duration::from_secs(1), rx.recv())
716 .await
717 .expect("handler fired within timeout")
718 .expect("handler produced an event")
719 }
720
721 #[test]
722 fn next_message_classifies_lag_close_and_ok() {
723 assert!(matches!(
724 next_message(Ok(("e".to_owned(), EventData::new()))),
725 ControlFlow::Continue(Some(_))
726 ));
727 assert!(matches!(
728 next_message(Err(RecvError::Lagged(7))),
729 ControlFlow::Continue(None)
730 ));
731 assert!(matches!(
732 next_message(Err(RecvError::Closed)),
733 ControlFlow::Break(())
734 ));
735 }
736
737 #[tokio::test]
738 async fn response_commands_share_one_generation_id() {
739 let (session, _) = session().await;
740 let generation_id = session.next_response_generation();
741 let mut delta = payload(json!({ "delta": "hello" }));
742 session.add_response_generation(&mut delta);
743 let mut commit = EventData::new();
744 session.add_response_generation(&mut commit);
745
746 assert_eq!(
747 delta.get("generation_id"),
748 Some(&Value::String(generation_id.clone()))
749 );
750 assert_eq!(
751 commit.get("generation_id"),
752 Some(&Value::String(generation_id))
753 );
754 }
755
756 #[tokio::test]
757 async fn on_error_parses_typed_fields() {
758 let (session, sender) = session().await;
759 let (tx, mut rx) = mpsc::unbounded_channel();
760 let _listener = session.on_error(move |event| {
761 tx.send(event).unwrap();
762 });
763 sender
764 .send((
765 EVENT_ERROR.to_owned(),
766 payload(json!({
767 "message": "cannot start now",
768 "code": ERROR_CODE_SESSION_FAILED,
769 "recoverable": false,
770 "generation_id": "gen-9"
771 })),
772 ))
773 .unwrap();
774 let event = recv(&mut rx).await;
775 assert_eq!(event.message.as_deref(), Some("cannot start now"));
776 assert_eq!(event.code.as_deref(), Some(ERROR_CODE_SESSION_FAILED));
777 assert!(!event.recoverable);
778 assert_eq!(event.generation_id.as_deref(), Some("gen-9"));
779 }
780
781 #[tokio::test]
782 async fn on_error_defaults_missing_recoverable_to_true() {
783 let (session, sender) = session().await;
784 let (tx, mut rx) = mpsc::unbounded_channel();
785 let _listener = session.on_error(move |event| {
786 tx.send(event).unwrap();
787 });
788 sender
789 .send((
790 EVENT_ERROR.to_owned(),
791 payload(json!({ "message": "legacy server", "code": "" })),
792 ))
793 .unwrap();
794 let event = recv(&mut rx).await;
795 assert!(event.recoverable);
796 assert_eq!(event.code, None);
797 assert_eq!(event.generation_id, None);
798 }
799
800 #[tokio::test]
801 async fn start_payload_uses_explicit_generation_id() {
802 let (session, _) = session().await;
803 let options = ResponseOptions {
804 allow_interruptions: Some(false),
805 generation_id: Some("gen-7".to_owned()),
806 };
807 let (generation_id, start) = session.start_payload(Some(options));
808 assert_eq!(generation_id, "gen-7");
809 assert_eq!(
810 start.get("generation_id"),
811 Some(&Value::String("gen-7".to_owned()))
812 );
813 assert_eq!(start.get("allow_interruptions"), Some(&Value::Bool(false)));
814
815 let mut commit = EventData::new();
816 session.thread_generation(&mut commit, None);
817 assert_eq!(
818 commit.get("generation_id"),
819 Some(&Value::String("gen-7".to_owned()))
820 );
821 }
822
823 #[tokio::test]
824 async fn start_payload_generates_generation_id_when_absent() {
825 let (session, _) = session().await;
826 let (generation_id, start) = session.start_payload(None);
827 assert!(generation_id.starts_with("generation_1_"));
828 assert!(generation_id.len() > "generation_1_".len());
829 assert_eq!(
830 start.get("generation_id"),
831 Some(&Value::String(generation_id))
832 );
833 }
834
835 #[tokio::test]
836 async fn explicit_generation_id_overrides_tracked_one() {
837 let (session, _) = session().await;
838 let tracked = session.next_response_generation();
839 let mut delta = payload(json!({ "delta": "hi" }));
840 session.thread_generation(&mut delta, Some("gen-42".to_owned()));
841 assert_eq!(
842 delta.get("generation_id"),
843 Some(&Value::String("gen-42".to_owned()))
844 );
845 assert_ne!(tracked, "gen-42");
846 }
847
848 #[tokio::test]
849 async fn response_events_expose_generation_id() {
850 let (session, sender) = session().await;
851 let (tx, mut rx) = mpsc::unbounded_channel();
852 let _listener = session.on_response_created(move |event| {
853 tx.send(event).unwrap();
854 });
855 sender
856 .send((
857 EVENT_RESPONSE_CREATED.to_owned(),
858 payload(json!({ "response_id": "resp-1", "generation_id": "gen-1" })),
859 ))
860 .unwrap();
861 let event = recv(&mut rx).await;
862 assert_eq!(event.response_id.as_deref(), Some("resp-1"));
863 assert_eq!(event.generation_id.as_deref(), Some("gen-1"));
864 }
865
866 #[tokio::test]
867 async fn audio_clear_and_interruption_expose_generation_id() {
868 let (session, sender) = session().await;
869 let (clear_tx, mut clear_rx) = mpsc::unbounded_channel();
870 let _clear = session.on_response_audio_clear(move |event| {
871 clear_tx.send(event).unwrap();
872 });
873 let (int_tx, mut int_rx) = mpsc::unbounded_channel();
874 let _interruption = session.on_interruption_detected(move |event| {
875 int_tx.send(event).unwrap();
876 });
877 sender
878 .send((
879 EVENT_RESPONSE_AUDIO_CLEAR.to_owned(),
880 payload(json!({ "response_id": "resp-2", "generation_id": "gen-2" })),
881 ))
882 .unwrap();
883 sender
884 .send((
885 EVENT_INTERRUPTION_DETECTED.to_owned(),
886 payload(json!({
887 "response_id": "resp-2",
888 "generation_id": "gen-2",
889 "vad_active_ms": 250
890 })),
891 ))
892 .unwrap();
893 let clear = recv(&mut clear_rx).await;
894 assert_eq!(clear.generation_id.as_deref(), Some("gen-2"));
895 let interruption = recv(&mut int_rx).await;
896 assert_eq!(interruption.response.generation_id.as_deref(), Some("gen-2"));
897 assert_eq!(interruption.vad_active_ms, Some(250.0));
898 }
899
900 #[tokio::test]
901 async fn start_response_and_wait_resolves_on_matching_created() {
902 let (session, sender) = session().await;
903 let options = ResponseOptions {
904 generation_id: Some("gen-ack".to_owned()),
905 ..Default::default()
906 };
907 tokio::spawn(async move {
908 tokio::time::sleep(Duration::from_millis(100)).await;
909 sender
910 .send((
911 EVENT_RESPONSE_CREATED.to_owned(),
912 payload(json!({ "response_id": "resp-other", "generation_id": "gen-other" })),
913 ))
914 .unwrap();
915 sender
916 .send((
917 EVENT_RESPONSE_CREATED.to_owned(),
918 payload(json!({ "response_id": "resp-9", "generation_id": "gen-ack" })),
919 ))
920 .unwrap();
921 });
922 let ack = session
923 .start_response_and_wait(Some(options), Duration::from_secs(2))
924 .await
925 .expect("ack within timeout");
926 assert!(ack.accepted);
927 assert_eq!(ack.generation_id, "gen-ack");
928 assert_eq!(ack.response_id.as_deref(), Some("resp-9"));
929 assert!(ack.recoverable);
930 assert_eq!(ack.error_code, None);
931 }
932
933 #[tokio::test]
934 async fn start_response_and_wait_surfaces_typed_rejection() {
935 let (session, sender) = session().await;
936 let options = ResponseOptions {
937 generation_id: Some("gen-rejected".to_owned()),
938 ..Default::default()
939 };
940 tokio::spawn(async move {
941 tokio::time::sleep(Duration::from_millis(100)).await;
942 sender
943 .send((
944 EVENT_ERROR.to_owned(),
945 payload(json!({
946 "message": "busy",
947 "code": ERROR_CODE_RESPONSE_ALREADY_ACTIVE,
948 "recoverable": true,
949 "generation_id": "gen-rejected"
950 })),
951 ))
952 .unwrap();
953 });
954 let ack = session
955 .start_response_and_wait(Some(options), Duration::from_secs(2))
956 .await
957 .expect("rejection within timeout");
958 assert!(!ack.accepted);
959 assert_eq!(ack.generation_id, "gen-rejected");
960 assert_eq!(
961 ack.error_code.as_deref(),
962 Some(ERROR_CODE_RESPONSE_ALREADY_ACTIVE)
963 );
964 assert_eq!(ack.error_message.as_deref(), Some("busy"));
965 assert!(ack.recoverable);
966 }
967
968 #[tokio::test]
969 async fn start_response_and_wait_times_out_without_ack() {
970 let (session, _sender) = session().await;
971 let error = session
972 .start_response_and_wait(None, Duration::from_millis(100))
973 .await
974 .expect_err("no ack must time out");
975 assert!(matches!(error, VoxRtcError::Timeout(_)));
976 }
977
978 #[tokio::test]
979 async fn on_speech_started_fires_with_timestamp() {
980 let (session, sender) = session().await;
981 let (tx, mut rx) = mpsc::unbounded_channel();
982 let _listener = session.on_speech_started(move |event| {
983 tx.send(event).unwrap();
984 });
985 sender
986 .send((
987 EVENT_SPEECH_STARTED.to_owned(),
988 payload(json!({ "session_id": "sess-1", "timestamp_ms": 1234 })),
989 ))
990 .unwrap();
991 let event = recv(&mut rx).await;
992 assert_eq!(event.session_id, "sess-1");
993 assert_eq!(event.channel_name, "/rtc/sess-1");
994 assert_eq!(event.timestamp_ms, Some(1234.0));
995 }
996
997 #[tokio::test]
998 async fn on_speech_stopped_fires_with_timestamp() {
999 let (session, sender) = session().await;
1000 let (tx, mut rx) = mpsc::unbounded_channel();
1001 let _listener = session.on_speech_stopped(move |event| {
1002 tx.send(event).unwrap();
1003 });
1004 sender
1005 .send((
1006 EVENT_SPEECH_STOPPED.to_owned(),
1007 payload(json!({ "timestamp_ms": 5678 })),
1008 ))
1009 .unwrap();
1010 let event = recv(&mut rx).await;
1011 assert_eq!(event.timestamp_ms, Some(5678.0));
1012 }
1013
1014 #[tokio::test]
1015 async fn on_transcript_delta_fires_with_fields() {
1016 let (session, sender) = session().await;
1017 let (tx, mut rx) = mpsc::unbounded_channel();
1018 let _listener = session.on_transcript_delta(move |event| {
1019 tx.send(event).unwrap();
1020 });
1021 sender
1022 .send((
1023 EVENT_TRANSCRIPT_DELTA.to_owned(),
1024 payload(json!({ "delta": "hel", "start_ms": 10, "end_ms": 20 })),
1025 ))
1026 .unwrap();
1027 let event = recv(&mut rx).await;
1028 assert_eq!(event.delta, "hel");
1029 assert_eq!(event.start_ms, Some(10.0));
1030 assert_eq!(event.end_ms, Some(20.0));
1031 }
1032
1033 #[tokio::test]
1034 async fn on_turn_eou_predicted_fires_with_fields() {
1035 let (session, sender) = session().await;
1036 let (tx, mut rx) = mpsc::unbounded_channel();
1037 let _listener = session.on_turn_eou_predicted(move |event| {
1038 tx.send(event).unwrap();
1039 });
1040 sender
1041 .send((
1042 EVENT_TURN_EOU_PREDICTED.to_owned(),
1043 payload(json!({
1044 "probability": 0.82,
1045 "threshold": 0.5,
1046 "delay_ms": 120,
1047 "start_ms": 0,
1048 "end_ms": 300,
1049 "decision": "end",
1050 "action": "commit",
1051 "turn_detector": "smart"
1052 })),
1053 ))
1054 .unwrap();
1055 let event = recv(&mut rx).await;
1056 assert_eq!(event.probability, Some(0.82));
1057 assert_eq!(event.threshold, Some(0.5));
1058 assert_eq!(event.delay_ms, Some(120.0));
1059 assert_eq!(event.start_ms, Some(0.0));
1060 assert_eq!(event.end_ms, Some(300.0));
1061 assert_eq!(event.decision.as_deref(), Some("end"));
1062 assert_eq!(event.action.as_deref(), Some("commit"));
1063 assert_eq!(event.turn_detector.as_deref(), Some("smart"));
1064 }
1065
1066 #[tokio::test]
1067 async fn handler_survives_a_lagged_broadcast() {
1068 let (session, sender) = session().await;
1069 let (tx, mut rx) = mpsc::unbounded_channel();
1070 let _listener = session.on_speech_started(move |event| {
1071 tx.send(event.timestamp_ms).unwrap();
1072 });
1073
1074 for index in 0..2100u32 {
1075 let _ = sender.send((
1076 EVENT_SPEECH_STARTED.to_owned(),
1077 payload(json!({ "timestamp_ms": index })),
1078 ));
1079 }
1080 let _ = sender.send((
1081 EVENT_SPEECH_STARTED.to_owned(),
1082 payload(json!({ "timestamp_ms": 9999 })),
1083 ));
1084
1085 let mut saw_final = false;
1086 while let Ok(Some(value)) = timeout(Duration::from_secs(1), rx.recv()).await {
1087 if value == Some(9999.0) {
1088 saw_final = true;
1089 break;
1090 }
1091 }
1092 assert!(
1093 saw_final,
1094 "loop must keep delivering events after a broadcast lag"
1095 );
1096 }
1097}