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: session_id.clone(),
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: session_id.clone(),
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: session_id.clone(),
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 entities: transcript_entities(&payload),
191 words: transcript_words(&payload),
192 data: payload,
193 });
194 })
195 }
196
197 pub fn on_turn_state_changed<F>(&self, handler: F) -> Listener
198 where
199 F: Fn(TurnStateEvent) + Send + Sync + 'static,
200 {
201 let session_id = self.session_id.clone();
202 let channel_name = self.channel_name.clone();
203 self.on(EVENT_TURN_STATE_CHANGED, move |payload| {
204 handler(TurnStateEvent {
205 session_id: session_id.clone(),
206 channel_name: channel_name.clone(),
207 state: required_string(&payload, "state", "unknown"),
208 previous_state: optional_string(&payload, "previous_state"),
209 data: payload,
210 });
211 })
212 }
213
214 pub fn on_speech_started<F>(&self, handler: F) -> Listener
215 where
216 F: Fn(SpeechStartedEvent) + Send + Sync + 'static,
217 {
218 let session_id = self.session_id.clone();
219 let channel_name = self.channel_name.clone();
220 self.on(EVENT_SPEECH_STARTED, move |payload| {
221 handler(SpeechStartedEvent {
222 session_id: session_id.clone(),
223 channel_name: channel_name.clone(),
224 timestamp_ms: optional_number(&payload, "timestamp_ms"),
225 data: payload,
226 });
227 })
228 }
229
230 pub fn on_speech_stopped<F>(&self, handler: F) -> Listener
231 where
232 F: Fn(SpeechStoppedEvent) + Send + Sync + 'static,
233 {
234 let session_id = self.session_id.clone();
235 let channel_name = self.channel_name.clone();
236 self.on(EVENT_SPEECH_STOPPED, move |payload| {
237 handler(SpeechStoppedEvent {
238 session_id: session_id.clone(),
239 channel_name: channel_name.clone(),
240 timestamp_ms: optional_number(&payload, "timestamp_ms"),
241 data: payload,
242 });
243 })
244 }
245
246 pub fn on_transcript_delta<F>(&self, handler: F) -> Listener
247 where
248 F: Fn(TranscriptDeltaEvent) + Send + Sync + 'static,
249 {
250 let session_id = self.session_id.clone();
251 let channel_name = self.channel_name.clone();
252 self.on(EVENT_TRANSCRIPT_DELTA, move |payload| {
253 handler(TranscriptDeltaEvent {
254 session_id: session_id.clone(),
255 channel_name: channel_name.clone(),
256 delta: required_string(&payload, "delta", ""),
257 start_ms: optional_number(&payload, "start_ms"),
258 end_ms: optional_number(&payload, "end_ms"),
259 data: payload,
260 });
261 })
262 }
263
264 pub fn on_turn_eou_predicted<F>(&self, handler: F) -> Listener
265 where
266 F: Fn(TurnEouPredictedEvent) + Send + Sync + 'static,
267 {
268 let session_id = self.session_id.clone();
269 let channel_name = self.channel_name.clone();
270 self.on(EVENT_TURN_EOU_PREDICTED, move |payload| {
271 handler(TurnEouPredictedEvent {
272 session_id: session_id.clone(),
273 channel_name: channel_name.clone(),
274 probability: optional_number(&payload, "probability"),
275 threshold: optional_number(&payload, "threshold"),
276 delay_ms: optional_number(&payload, "delay_ms"),
277 start_ms: optional_number(&payload, "start_ms"),
278 end_ms: optional_number(&payload, "end_ms"),
279 decision: optional_string(&payload, "decision"),
280 action: optional_string(&payload, "action"),
281 turn_detector: optional_string(&payload, "turn_detector"),
282 data: payload,
283 });
284 })
285 }
286
287 pub fn on_response_created<F>(&self, handler: F) -> Listener
288 where
289 F: Fn(ResponseEvent) + Send + Sync + 'static,
290 {
291 self.on_response_event(EVENT_RESPONSE_CREATED, handler)
292 }
293
294 pub fn on_response_committed<F>(&self, handler: F) -> Listener
295 where
296 F: Fn(ResponseEvent) + Send + Sync + 'static,
297 {
298 self.on_response_event(EVENT_RESPONSE_COMMITTED, handler)
299 }
300
301 pub fn on_response_done<F>(&self, handler: F) -> Listener
302 where
303 F: Fn(ResponseEvent) + Send + Sync + 'static,
304 {
305 self.on_response_event(EVENT_RESPONSE_DONE, handler)
306 }
307
308 pub fn on_response_cancelled<F>(&self, handler: F) -> Listener
309 where
310 F: Fn(ResponseEvent) + Send + Sync + 'static,
311 {
312 self.on_response_event(EVENT_RESPONSE_CANCELLED, handler)
313 }
314
315 pub fn on_response_audio_clear<F>(&self, handler: F) -> Listener
316 where
317 F: Fn(ResponseEvent) + Send + Sync + 'static,
318 {
319 self.on_response_event(EVENT_RESPONSE_AUDIO_CLEAR, handler)
320 }
321
322 fn on_response_event<F>(&self, event_name: &'static str, handler: F) -> Listener
323 where
324 F: Fn(ResponseEvent) + Send + Sync + 'static,
325 {
326 let session_id = self.session_id.clone();
327 let channel_name = self.channel_name.clone();
328 self.on(event_name, move |payload| {
329 handler(response_event(payload, &session_id, &channel_name));
330 })
331 }
332
333 pub fn on_interruption_detected<F>(&self, handler: F) -> Listener
334 where
335 F: Fn(InterruptionEvent) + Send + Sync + 'static,
336 {
337 self.on_interruption_event(EVENT_INTERRUPTION_DETECTED, handler)
338 }
339
340 pub fn on_interruption_false_positive<F>(&self, handler: F) -> Listener
341 where
342 F: Fn(InterruptionEvent) + Send + Sync + 'static,
343 {
344 self.on_interruption_event(EVENT_INTERRUPTION_FALSE_POSITIVE, handler)
345 }
346
347 fn on_interruption_event<F>(&self, event_name: &'static str, handler: F) -> Listener
348 where
349 F: Fn(InterruptionEvent) + Send + Sync + 'static,
350 {
351 let session_id = self.session_id.clone();
352 let channel_name = self.channel_name.clone();
353 self.on(event_name, move |payload| {
354 handler(InterruptionEvent {
355 response: response_event(payload.clone(), &session_id, &channel_name),
356 vad_active_ms: optional_number(&payload, "vad_active_ms"),
357 partial_transcript: optional_string(&payload, "partial_transcript"),
358 reason: optional_nonempty_string(&payload, "reason"),
359 });
360 })
361 }
362
363 pub fn on_browser_event<F>(&self, handler: F) -> Listener
364 where
365 F: Fn(BrowserEvent) + Send + Sync + 'static,
366 {
367 let session_id = self.session_id.clone();
368 let channel_name = self.channel_name.clone();
369 self.on(EVENT_BROWSER_EVENT, move |payload| {
370 handler(BrowserEvent {
371 session_id: session_id.clone(),
372 channel_name: channel_name.clone(),
373 event: required_string(&payload, "event", ""),
374 payload: payload.get("payload").cloned().unwrap_or(Value::Null),
375 data: payload,
376 });
377 })
378 }
379
380 pub fn on_close<F>(&self, handler: F) -> Listener
381 where
382 F: Fn(CloseEvent) + Send + Sync + 'static,
383 {
384 let session_id = self.session_id.clone();
385 let channel_name = self.channel_name.clone();
386 self.on(EVENT_RTC_CLIENT_DISCONNECTED, move |payload| {
387 handler(CloseEvent {
388 session_id: session_id.clone(),
389 channel_name: channel_name.clone(),
390 reason: required_string(&payload, "reason", "unknown"),
391 connection_state: optional_string(&payload, "connection_state"),
392 ice_connection_state: optional_string(&payload, "ice_connection_state"),
393 data_channel_state: optional_string(&payload, "data_channel_state"),
394 data: payload,
395 });
396 })
397 }
398
399 pub fn on_error<F>(&self, handler: F) -> Listener
400 where
401 F: Fn(ErrorEvent) + Send + Sync + 'static,
402 {
403 let session_id = self.session_id.clone();
404 let channel_name = self.channel_name.clone();
405 self.on(EVENT_ERROR, move |payload| {
406 handler(ErrorEvent {
407 session_id: session_id.clone(),
408 channel_name: channel_name.clone(),
409 message: optional_string(&payload, "message"),
410 code: optional_nonempty_string(&payload, "code"),
411 recoverable: recoverable_flag(&payload),
412 generation_id: optional_nonempty_string(&payload, "generation_id"),
413 data: payload,
414 });
415 })
416 }
417
418 pub fn on_signaling_error<F>(&self, handler: F) -> Listener
419 where
420 F: Fn(SignalingErrorEvent) + Send + Sync + 'static,
421 {
422 let session_id = self.session_id.clone();
423 let channel_name = self.channel_name.clone();
424 self.on(EVENT_RTC_SIGNALING_ERROR, move |payload| {
425 handler(SignalingErrorEvent {
426 session_id: session_id.clone(),
427 channel_name: channel_name.clone(),
428 message: optional_string(&payload, "message"),
429 generation: optional_i64(&payload, "generation"),
430 data: payload,
431 });
432 })
433 }
434
435 pub async fn send_control(&self, event: &str, payload: EventData) -> Result<()> {
436 self.channel.send_message(event, payload).await
437 }
438
439 pub async fn configure(&self, config: SessionConfig) -> Result<()> {
440 let mut session = config.extra;
441 insert_opt(&mut session, "stt_model", config.stt_model);
442 insert_opt(&mut session, "tts_model", config.tts_model);
443 insert_opt(&mut session, "voice", config.voice);
444 insert_opt(&mut session, "turn_profile", config.turn_profile);
445 insert_opt(&mut session, "vad_backend", config.vad_backend);
446 insert_opt(&mut session, "turn_detector", config.turn_detector);
447
448 let mut payload = EventData::new();
449 payload.insert("session".to_owned(), Value::Object(session));
450 self.send_control("session.update", payload).await
451 }
452
453 pub async fn start_response(&self, options: Option<ResponseOptions>) -> Result<()> {
454 let (_, payload) = self.start_payload(options);
455 self.send_control("response.start", payload).await
456 }
457
458 pub async fn start_response_and_wait(
459 &self,
460 options: Option<ResponseOptions>,
461 wait_timeout: Duration,
462 ) -> Result<StartAck> {
463 let (generation_id, payload) = self.start_payload(options);
464 let mut messages = self.channel.subscribe_messages();
465 self.send_control("response.start", payload).await?;
466 timeout(wait_timeout, async move {
467 loop {
468 match next_message(messages.recv().await) {
469 ControlFlow::Break(()) => return Err(VoxRtcError::ChannelClosed),
470 ControlFlow::Continue(None) => continue,
471 ControlFlow::Continue(Some((event, data))) => {
472 if optional_nonempty_string(&data, "generation_id").as_deref()
473 != Some(generation_id.as_str())
474 {
475 continue;
476 }
477 if event == EVENT_RESPONSE_CREATED {
478 return Ok(StartAck {
479 accepted: true,
480 generation_id: generation_id.clone(),
481 response_id: optional_string(&data, "response_id"),
482 error_code: None,
483 error_message: None,
484 recoverable: true,
485 });
486 }
487 if event == EVENT_ERROR {
488 return Ok(StartAck {
489 accepted: false,
490 generation_id: generation_id.clone(),
491 response_id: optional_string(&data, "response_id"),
492 error_code: optional_nonempty_string(&data, "code"),
493 error_message: optional_string(&data, "message"),
494 recoverable: recoverable_flag(&data),
495 });
496 }
497 }
498 }
499 }
500 })
501 .await
502 .map_err(|_| VoxRtcError::Timeout("response.start acknowledgement"))?
503 }
504
505 pub async fn append_response_text(
506 &self,
507 delta: impl Into<String>,
508 options: Option<ResponseOptions>,
509 ) -> Result<()> {
510 let explicit = explicit_generation(&options);
511 let mut payload = response_options_payload(options);
512 payload.insert("delta".to_owned(), Value::String(delta.into()));
513 self.thread_generation(&mut payload, explicit);
514 self.send_control("response.delta", payload).await
515 }
516
517 pub async fn commit_response(&self, options: Option<ResponseOptions>) -> Result<()> {
518 let explicit = explicit_generation(&options);
519 let mut payload = EventData::new();
520 self.thread_generation(&mut payload, explicit);
521 self.send_control("response.commit", payload).await
522 }
523
524 pub async fn cancel_response(&self, options: Option<ResponseOptions>) -> Result<()> {
525 let explicit = explicit_generation(&options);
526 let mut payload = EventData::new();
527 self.thread_generation(&mut payload, explicit);
528 self.clear_response_generation();
529 self.send_control("response.cancel", payload).await
530 }
531
532 pub async fn replace_response_text(
533 &self,
534 text: impl Into<String>,
535 options: Option<ResponseOptions>,
536 ) -> Result<()> {
537 self.clear_response_generation();
538 let explicit = explicit_generation(&options);
539 let mut payload = response_options_payload(options);
540 payload.insert("text".to_owned(), Value::String(text.into()));
541 if let Some(generation_id) = explicit {
542 payload.insert("generation_id".to_owned(), Value::String(generation_id));
543 }
544 self.send_control("response.replace_text", payload).await
545 }
546
547 pub async fn send_text_response(
548 &self,
549 text: impl Into<String>,
550 options: Option<ResponseOptions>,
551 cancel_first: bool,
552 ) -> Result<()> {
553 let text = text.into();
554 if cancel_first {
555 return self.replace_response_text(text, options).await;
556 }
557 self.start_response(options.clone()).await?;
558 self.append_response_text(text, options.clone()).await?;
559 self.commit_response(options).await
560 }
561
562 pub async fn send_client_event(&self, envelope: ClientEventEnvelope) -> Result<()> {
563 let mut payload = EventData::new();
564 payload.insert("event".to_owned(), Value::String(envelope.event));
565 payload.insert("payload".to_owned(), envelope.payload);
566 self.send_control(EVENT_CLIENT_EVENT, payload).await
567 }
568
569 fn start_payload(&self, options: Option<ResponseOptions>) -> (String, EventData) {
570 let explicit = explicit_generation(&options);
571 let mut payload = response_options_payload(options);
572 let generation_id = match explicit {
573 Some(id) => self.set_response_generation(id),
574 None => self.next_response_generation(),
575 };
576 payload.insert(
577 "generation_id".to_owned(),
578 Value::String(generation_id.clone()),
579 );
580 (generation_id, payload)
581 }
582
583 fn next_response_generation(&self) -> String {
584 let mut state = self
585 .response_generation
586 .lock()
587 .expect("response generation mutex poisoned");
588 state.counter += 1;
589 let generation_id = format!("generation_{}_{}", state.counter, Uuid::new_v4());
590 state.id = Some(generation_id.clone());
591 generation_id
592 }
593
594 fn set_response_generation(&self, generation_id: String) -> String {
595 let mut state = self
596 .response_generation
597 .lock()
598 .expect("response generation mutex poisoned");
599 state.counter += 1;
600 state.id = Some(generation_id.clone());
601 generation_id
602 }
603
604 fn thread_generation(&self, payload: &mut EventData, explicit: Option<String>) {
605 match explicit {
606 Some(generation_id) => {
607 payload.insert("generation_id".to_owned(), Value::String(generation_id));
608 }
609 None => self.add_response_generation(payload),
610 }
611 }
612
613 fn add_response_generation(&self, payload: &mut EventData) {
614 let state = self
615 .response_generation
616 .lock()
617 .expect("response generation mutex poisoned");
618 if let Some(generation_id) = &state.id {
619 payload.insert(
620 "generation_id".to_owned(),
621 Value::String(generation_id.clone()),
622 );
623 }
624 }
625
626 fn clear_response_generation(&self) {
627 self.response_generation
628 .lock()
629 .expect("response generation mutex poisoned")
630 .id = None;
631 }
632}
633
634fn next_message(
635 result: std::result::Result<(String, EventData), RecvError>,
636) -> ControlFlow<(), Option<(String, EventData)>> {
637 match result {
638 Ok(message) => ControlFlow::Continue(Some(message)),
639 Err(RecvError::Lagged(_)) => ControlFlow::Continue(None),
640 Err(RecvError::Closed) => ControlFlow::Break(()),
641 }
642}
643
644fn join_decline_reason(reason: Option<EventData>) -> Option<String> {
645 let reason = reason?;
646 for key in ["message", "reason", "error"] {
647 if let Some(value) = reason.get(key).and_then(Value::as_str)
648 && !value.is_empty()
649 {
650 return Some(value.to_owned());
651 }
652 }
653 if reason.is_empty() {
654 None
655 } else {
656 Some(Value::Object(reason).to_string())
657 }
658}
659
660fn insert_opt(session: &mut EventData, key: &str, value: Option<String>) {
661 if let Some(value) = value {
662 session.insert(key.to_owned(), Value::String(value));
663 }
664}
665
666fn response_options_payload(options: Option<ResponseOptions>) -> EventData {
667 let mut payload = EventData::new();
668 if let Some(options) = options
669 && let Some(allow) = options.allow_interruptions
670 {
671 payload.insert("allow_interruptions".to_owned(), Value::Bool(allow));
672 }
673 payload
674}
675
676fn explicit_generation(options: &Option<ResponseOptions>) -> Option<String> {
677 options
678 .as_ref()
679 .and_then(|options| options.generation_id.clone())
680 .filter(|id| !id.is_empty())
681}
682
683fn response_event(payload: EventData, session_id: &str, channel_name: &str) -> ResponseEvent {
684 ResponseEvent {
685 session_id: session_id.to_owned(),
686 channel_name: channel_name.to_owned(),
687 response_id: optional_string(&payload, "response_id"),
688 generation_id: optional_nonempty_string(&payload, "generation_id"),
689 data: payload,
690 }
691}
692
693#[cfg(test)]
694mod tests {
695 use super::*;
696 use crate::socket::test_channel;
697 use serde_json::json;
698 use tokio::sync::broadcast;
699 use tokio::sync::mpsc;
700
701 async fn session() -> (VoxRtcControlSession, broadcast::Sender<(String, EventData)>) {
702 let (channel, sender) = test_channel().await;
703 let session =
704 VoxRtcControlSession::new(channel, "sess-1".to_owned(), Duration::from_secs(1));
705 (session, sender)
706 }
707
708 fn payload(value: Value) -> EventData {
709 value.as_object().cloned().expect("object payload")
710 }
711
712 #[test]
713 fn join_decline_reason_prefers_structured_message_fields() {
714 assert_eq!(
715 join_decline_reason(Some(payload(json!({ "message": "expired" })))),
716 Some("expired".to_owned())
717 );
718 assert_eq!(
719 join_decline_reason(Some(payload(json!({ "reason": "missing" })))),
720 Some("missing".to_owned())
721 );
722 assert_eq!(
723 join_decline_reason(Some(payload(json!({ "channel": "/rtc/abc" })))),
724 Some(r#"{"channel":"/rtc/abc"}"#.to_owned())
725 );
726 assert_eq!(join_decline_reason(Some(EventData::new())), None);
727 assert_eq!(join_decline_reason(None), None);
728 }
729
730 async fn recv<T>(rx: &mut mpsc::UnboundedReceiver<T>) -> T {
731 timeout(Duration::from_secs(1), rx.recv())
732 .await
733 .expect("handler fired within timeout")
734 .expect("handler produced an event")
735 }
736
737 #[test]
738 fn next_message_classifies_lag_close_and_ok() {
739 assert!(matches!(
740 next_message(Ok(("e".to_owned(), EventData::new()))),
741 ControlFlow::Continue(Some(_))
742 ));
743 assert!(matches!(
744 next_message(Err(RecvError::Lagged(7))),
745 ControlFlow::Continue(None)
746 ));
747 assert!(matches!(
748 next_message(Err(RecvError::Closed)),
749 ControlFlow::Break(())
750 ));
751 }
752
753 #[tokio::test]
754 async fn response_commands_share_one_generation_id() {
755 let (session, _) = session().await;
756 let generation_id = session.next_response_generation();
757 let mut delta = payload(json!({ "delta": "hello" }));
758 session.add_response_generation(&mut delta);
759 let mut commit = EventData::new();
760 session.add_response_generation(&mut commit);
761
762 assert_eq!(
763 delta.get("generation_id"),
764 Some(&Value::String(generation_id.clone()))
765 );
766 assert_eq!(
767 commit.get("generation_id"),
768 Some(&Value::String(generation_id))
769 );
770 }
771
772 #[tokio::test]
773 async fn on_error_parses_typed_fields() {
774 let (session, sender) = session().await;
775 let (tx, mut rx) = mpsc::unbounded_channel();
776 let _listener = session.on_error(move |event| {
777 tx.send(event).unwrap();
778 });
779 sender
780 .send((
781 EVENT_ERROR.to_owned(),
782 payload(json!({
783 "message": "cannot start now",
784 "code": ERROR_CODE_SESSION_FAILED,
785 "recoverable": false,
786 "generation_id": "gen-9"
787 })),
788 ))
789 .unwrap();
790 let event = recv(&mut rx).await;
791 assert_eq!(event.message.as_deref(), Some("cannot start now"));
792 assert_eq!(event.code.as_deref(), Some(ERROR_CODE_SESSION_FAILED));
793 assert!(!event.recoverable);
794 assert_eq!(event.generation_id.as_deref(), Some("gen-9"));
795 }
796
797 #[tokio::test]
798 async fn on_error_defaults_missing_recoverable_to_true() {
799 let (session, sender) = session().await;
800 let (tx, mut rx) = mpsc::unbounded_channel();
801 let _listener = session.on_error(move |event| {
802 tx.send(event).unwrap();
803 });
804 sender
805 .send((
806 EVENT_ERROR.to_owned(),
807 payload(json!({ "message": "legacy server", "code": "" })),
808 ))
809 .unwrap();
810 let event = recv(&mut rx).await;
811 assert!(event.recoverable);
812 assert_eq!(event.code, None);
813 assert_eq!(event.generation_id, None);
814 }
815
816 #[tokio::test]
817 async fn start_payload_uses_explicit_generation_id() {
818 let (session, _) = session().await;
819 let options = ResponseOptions {
820 allow_interruptions: Some(false),
821 generation_id: Some("gen-7".to_owned()),
822 };
823 let (generation_id, start) = session.start_payload(Some(options));
824 assert_eq!(generation_id, "gen-7");
825 assert_eq!(
826 start.get("generation_id"),
827 Some(&Value::String("gen-7".to_owned()))
828 );
829 assert_eq!(start.get("allow_interruptions"), Some(&Value::Bool(false)));
830
831 let mut commit = EventData::new();
832 session.thread_generation(&mut commit, None);
833 assert_eq!(
834 commit.get("generation_id"),
835 Some(&Value::String("gen-7".to_owned()))
836 );
837 }
838
839 #[tokio::test]
840 async fn start_payload_generates_generation_id_when_absent() {
841 let (session, _) = session().await;
842 let (generation_id, start) = session.start_payload(None);
843 assert!(generation_id.starts_with("generation_1_"));
844 assert!(generation_id.len() > "generation_1_".len());
845 assert_eq!(
846 start.get("generation_id"),
847 Some(&Value::String(generation_id))
848 );
849 }
850
851 #[tokio::test]
852 async fn explicit_generation_id_overrides_tracked_one() {
853 let (session, _) = session().await;
854 let tracked = session.next_response_generation();
855 let mut delta = payload(json!({ "delta": "hi" }));
856 session.thread_generation(&mut delta, Some("gen-42".to_owned()));
857 assert_eq!(
858 delta.get("generation_id"),
859 Some(&Value::String("gen-42".to_owned()))
860 );
861 assert_ne!(tracked, "gen-42");
862 }
863
864 #[tokio::test]
865 async fn response_events_expose_generation_id() {
866 let (session, sender) = session().await;
867 let (tx, mut rx) = mpsc::unbounded_channel();
868 let _listener = session.on_response_created(move |event| {
869 tx.send(event).unwrap();
870 });
871 sender
872 .send((
873 EVENT_RESPONSE_CREATED.to_owned(),
874 payload(json!({ "response_id": "resp-1", "generation_id": "gen-1" })),
875 ))
876 .unwrap();
877 let event = recv(&mut rx).await;
878 assert_eq!(event.response_id.as_deref(), Some("resp-1"));
879 assert_eq!(event.generation_id.as_deref(), Some("gen-1"));
880 }
881
882 #[tokio::test]
883 async fn audio_clear_and_interruption_expose_generation_id() {
884 let (session, sender) = session().await;
885 let (clear_tx, mut clear_rx) = mpsc::unbounded_channel();
886 let _clear = session.on_response_audio_clear(move |event| {
887 clear_tx.send(event).unwrap();
888 });
889 let (int_tx, mut int_rx) = mpsc::unbounded_channel();
890 let _interruption = session.on_interruption_detected(move |event| {
891 int_tx.send(event).unwrap();
892 });
893 sender
894 .send((
895 EVENT_RESPONSE_AUDIO_CLEAR.to_owned(),
896 payload(json!({ "response_id": "resp-2", "generation_id": "gen-2" })),
897 ))
898 .unwrap();
899 sender
900 .send((
901 EVENT_INTERRUPTION_DETECTED.to_owned(),
902 payload(json!({
903 "response_id": "resp-2",
904 "generation_id": "gen-2",
905 "vad_active_ms": 250
906 })),
907 ))
908 .unwrap();
909 let clear = recv(&mut clear_rx).await;
910 assert_eq!(clear.generation_id.as_deref(), Some("gen-2"));
911 let interruption = recv(&mut int_rx).await;
912 assert_eq!(interruption.response.generation_id.as_deref(), Some("gen-2"));
913 assert_eq!(interruption.vad_active_ms, Some(250.0));
914 }
915
916 #[tokio::test]
917 async fn on_signaling_error_parses_message_and_generation() {
918 let (session, sender) = session().await;
919 let (tx, mut rx) = mpsc::unbounded_channel();
920 let _listener = session.on_signaling_error(move |event| {
921 tx.send(event).unwrap();
922 });
923 sender
924 .send((
925 EVENT_RTC_SIGNALING_ERROR.to_owned(),
926 payload(json!({
927 "message": "setLocalDescription failed",
928 "generation": 3
929 })),
930 ))
931 .unwrap();
932 let event = recv(&mut rx).await;
933 assert_eq!(event.message.as_deref(), Some("setLocalDescription failed"));
934 assert_eq!(event.generation, Some(3));
935 }
936
937 #[tokio::test]
938 async fn on_signaling_error_leaves_generation_none_when_absent() {
939 let (session, sender) = session().await;
940 let (tx, mut rx) = mpsc::unbounded_channel();
941 let _listener = session.on_signaling_error(move |event| {
942 tx.send(event).unwrap();
943 });
944 sender
945 .send((
946 EVENT_RTC_SIGNALING_ERROR.to_owned(),
947 payload(json!({ "message": "RTC signaling failed" })),
948 ))
949 .unwrap();
950 let event = recv(&mut rx).await;
951 assert_eq!(event.message.as_deref(), Some("RTC signaling failed"));
952 assert_eq!(event.generation, None);
953 }
954
955 #[tokio::test]
956 async fn on_transcript_exposes_entities_and_words() {
957 let (session, sender) = session().await;
958 let (tx, mut rx) = mpsc::unbounded_channel();
959 let _listener = session.on_transcript(move |event| {
960 tx.send(event).unwrap();
961 });
962 sender
963 .send((
964 EVENT_TRANSCRIPT_COMPLETED.to_owned(),
965 payload(json!({
966 "transcript": "call Ada",
967 "entities": [
968 { "type": "PRODUCT", "text": "Ada", "start_char": 5, "end_char": 8 }
969 ],
970 "words": [
971 { "word": "call", "start_ms": 0, "end_ms": 300 },
972 { "word": "Ada", "start_ms": 300, "end_ms": 600, "confidence": 0.91 }
973 ]
974 })),
975 ))
976 .unwrap();
977 let event = recv(&mut rx).await;
978 assert_eq!(
979 event.entities,
980 vec![TranscriptEntity {
981 r#type: "PRODUCT".to_owned(),
982 text: "Ada".to_owned(),
983 start_char: 5,
984 end_char: 8,
985 }]
986 );
987 assert_eq!(event.words.len(), 2);
988 assert_eq!(event.words[0].word, "call");
989 assert_eq!(event.words[0].start_ms, 0.0);
990 assert_eq!(event.words[0].confidence, None);
991 assert_eq!(event.words[1].confidence, Some(0.91));
992 }
993
994 #[tokio::test]
995 async fn on_transcript_defaults_entities_and_words_to_empty() {
996 let (session, sender) = session().await;
997 let (tx, mut rx) = mpsc::unbounded_channel();
998 let _listener = session.on_transcript(move |event| {
999 tx.send(event).unwrap();
1000 });
1001 sender
1002 .send((
1003 EVENT_TRANSCRIPT_COMPLETED.to_owned(),
1004 payload(json!({ "transcript": "hello" })),
1005 ))
1006 .unwrap();
1007 let event = recv(&mut rx).await;
1008 assert!(event.entities.is_empty());
1009 assert!(event.words.is_empty());
1010 }
1011
1012 #[tokio::test]
1013 async fn interruption_events_expose_reason() {
1014 let (session, sender) = session().await;
1015 let (det_tx, mut det_rx) = mpsc::unbounded_channel();
1016 let _detected = session.on_interruption_detected(move |event| {
1017 det_tx.send(event).unwrap();
1018 });
1019 let (fp_tx, mut fp_rx) = mpsc::unbounded_channel();
1020 let _false_positive = session.on_interruption_false_positive(move |event| {
1021 fp_tx.send(event).unwrap();
1022 });
1023 sender
1024 .send((
1025 EVENT_INTERRUPTION_DETECTED.to_owned(),
1026 payload(json!({
1027 "response_id": "resp-3",
1028 "generation_id": "gen-3",
1029 "reason": "speech_overlap"
1030 })),
1031 ))
1032 .unwrap();
1033 sender
1034 .send((
1035 EVENT_INTERRUPTION_FALSE_POSITIVE.to_owned(),
1036 payload(json!({ "response_id": "resp-3", "reason": "backchannel" })),
1037 ))
1038 .unwrap();
1039 let detected = recv(&mut det_rx).await;
1040 assert_eq!(detected.reason.as_deref(), Some("speech_overlap"));
1041 let false_positive = recv(&mut fp_rx).await;
1042 assert_eq!(false_positive.reason.as_deref(), Some("backchannel"));
1043 }
1044
1045 #[tokio::test]
1046 async fn start_response_and_wait_resolves_on_matching_created() {
1047 let (session, sender) = session().await;
1048 let options = ResponseOptions {
1049 generation_id: Some("gen-ack".to_owned()),
1050 ..Default::default()
1051 };
1052 tokio::spawn(async move {
1053 tokio::time::sleep(Duration::from_millis(100)).await;
1054 sender
1055 .send((
1056 EVENT_RESPONSE_CREATED.to_owned(),
1057 payload(json!({ "response_id": "resp-other", "generation_id": "gen-other" })),
1058 ))
1059 .unwrap();
1060 sender
1061 .send((
1062 EVENT_RESPONSE_CREATED.to_owned(),
1063 payload(json!({ "response_id": "resp-9", "generation_id": "gen-ack" })),
1064 ))
1065 .unwrap();
1066 });
1067 let ack = session
1068 .start_response_and_wait(Some(options), Duration::from_secs(2))
1069 .await
1070 .expect("ack within timeout");
1071 assert!(ack.accepted);
1072 assert_eq!(ack.generation_id, "gen-ack");
1073 assert_eq!(ack.response_id.as_deref(), Some("resp-9"));
1074 assert!(ack.recoverable);
1075 assert_eq!(ack.error_code, None);
1076 }
1077
1078 #[tokio::test]
1079 async fn start_response_and_wait_surfaces_typed_rejection() {
1080 let (session, sender) = session().await;
1081 let options = ResponseOptions {
1082 generation_id: Some("gen-rejected".to_owned()),
1083 ..Default::default()
1084 };
1085 tokio::spawn(async move {
1086 tokio::time::sleep(Duration::from_millis(100)).await;
1087 sender
1088 .send((
1089 EVENT_ERROR.to_owned(),
1090 payload(json!({
1091 "message": "busy",
1092 "code": ERROR_CODE_RESPONSE_ALREADY_ACTIVE,
1093 "recoverable": true,
1094 "generation_id": "gen-rejected"
1095 })),
1096 ))
1097 .unwrap();
1098 });
1099 let ack = session
1100 .start_response_and_wait(Some(options), Duration::from_secs(2))
1101 .await
1102 .expect("rejection within timeout");
1103 assert!(!ack.accepted);
1104 assert_eq!(ack.generation_id, "gen-rejected");
1105 assert_eq!(
1106 ack.error_code.as_deref(),
1107 Some(ERROR_CODE_RESPONSE_ALREADY_ACTIVE)
1108 );
1109 assert_eq!(ack.error_message.as_deref(), Some("busy"));
1110 assert!(ack.recoverable);
1111 }
1112
1113 #[tokio::test]
1114 async fn start_response_and_wait_times_out_without_ack() {
1115 let (session, _sender) = session().await;
1116 let error = session
1117 .start_response_and_wait(None, Duration::from_millis(100))
1118 .await
1119 .expect_err("no ack must time out");
1120 assert!(matches!(error, VoxRtcError::Timeout(_)));
1121 }
1122
1123 #[tokio::test]
1124 async fn on_speech_started_fires_with_timestamp() {
1125 let (session, sender) = session().await;
1126 let (tx, mut rx) = mpsc::unbounded_channel();
1127 let _listener = session.on_speech_started(move |event| {
1128 tx.send(event).unwrap();
1129 });
1130 sender
1131 .send((
1132 EVENT_SPEECH_STARTED.to_owned(),
1133 payload(json!({ "session_id": "sess-1", "timestamp_ms": 1234 })),
1134 ))
1135 .unwrap();
1136 let event = recv(&mut rx).await;
1137 assert_eq!(event.session_id, "sess-1");
1138 assert_eq!(event.channel_name, "/rtc/sess-1");
1139 assert_eq!(event.timestamp_ms, Some(1234.0));
1140 }
1141
1142 #[tokio::test]
1143 async fn on_speech_stopped_fires_with_timestamp() {
1144 let (session, sender) = session().await;
1145 let (tx, mut rx) = mpsc::unbounded_channel();
1146 let _listener = session.on_speech_stopped(move |event| {
1147 tx.send(event).unwrap();
1148 });
1149 sender
1150 .send((
1151 EVENT_SPEECH_STOPPED.to_owned(),
1152 payload(json!({ "timestamp_ms": 5678 })),
1153 ))
1154 .unwrap();
1155 let event = recv(&mut rx).await;
1156 assert_eq!(event.timestamp_ms, Some(5678.0));
1157 }
1158
1159 #[tokio::test]
1160 async fn on_transcript_delta_fires_with_fields() {
1161 let (session, sender) = session().await;
1162 let (tx, mut rx) = mpsc::unbounded_channel();
1163 let _listener = session.on_transcript_delta(move |event| {
1164 tx.send(event).unwrap();
1165 });
1166 sender
1167 .send((
1168 EVENT_TRANSCRIPT_DELTA.to_owned(),
1169 payload(json!({ "delta": "hel", "start_ms": 10, "end_ms": 20 })),
1170 ))
1171 .unwrap();
1172 let event = recv(&mut rx).await;
1173 assert_eq!(event.delta, "hel");
1174 assert_eq!(event.start_ms, Some(10.0));
1175 assert_eq!(event.end_ms, Some(20.0));
1176 }
1177
1178 #[tokio::test]
1179 async fn on_turn_eou_predicted_fires_with_fields() {
1180 let (session, sender) = session().await;
1181 let (tx, mut rx) = mpsc::unbounded_channel();
1182 let _listener = session.on_turn_eou_predicted(move |event| {
1183 tx.send(event).unwrap();
1184 });
1185 sender
1186 .send((
1187 EVENT_TURN_EOU_PREDICTED.to_owned(),
1188 payload(json!({
1189 "probability": 0.82,
1190 "threshold": 0.5,
1191 "delay_ms": 120,
1192 "start_ms": 0,
1193 "end_ms": 300,
1194 "decision": "end",
1195 "action": "commit",
1196 "turn_detector": "smart"
1197 })),
1198 ))
1199 .unwrap();
1200 let event = recv(&mut rx).await;
1201 assert_eq!(event.probability, Some(0.82));
1202 assert_eq!(event.threshold, Some(0.5));
1203 assert_eq!(event.delay_ms, Some(120.0));
1204 assert_eq!(event.start_ms, Some(0.0));
1205 assert_eq!(event.end_ms, Some(300.0));
1206 assert_eq!(event.decision.as_deref(), Some("end"));
1207 assert_eq!(event.action.as_deref(), Some("commit"));
1208 assert_eq!(event.turn_detector.as_deref(), Some("smart"));
1209 }
1210
1211 #[tokio::test]
1212 async fn handler_survives_a_lagged_broadcast() {
1213 let (session, sender) = session().await;
1214 let (tx, mut rx) = mpsc::unbounded_channel();
1215 let _listener = session.on_speech_started(move |event| {
1216 tx.send(event.timestamp_ms).unwrap();
1217 });
1218
1219 for index in 0..2100u32 {
1220 let _ = sender.send((
1221 EVENT_SPEECH_STARTED.to_owned(),
1222 payload(json!({ "timestamp_ms": index })),
1223 ));
1224 }
1225 let _ = sender.send((
1226 EVENT_SPEECH_STARTED.to_owned(),
1227 payload(json!({ "timestamp_ms": 9999 })),
1228 ));
1229
1230 let mut saw_final = false;
1231 while let Ok(Some(value)) = timeout(Duration::from_secs(1), rx.recv()).await {
1232 if value == Some(9999.0) {
1233 saw_final = true;
1234 break;
1235 }
1236 }
1237 assert!(
1238 saw_final,
1239 "loop must keep delivering events after a broadcast lag"
1240 );
1241 }
1242}