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