Skip to main content

vox_rtc_server/
session.rs

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}