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                    .cloned()
195                    .and_then(|value| serde_json::from_value::<SpeechContext>(value).ok())
196                    .filter(SpeechContext::is_valid),
197                data: payload,
198            });
199        })
200    }
201
202    pub fn on_turn_state_changed<F>(&self, handler: F) -> Listener
203    where
204        F: Fn(TurnStateEvent) + Send + Sync + 'static,
205    {
206        let session_id = self.session_id.clone();
207        let channel_name = self.channel_name.clone();
208        self.on(EVENT_TURN_STATE_CHANGED, move |payload| {
209            handler(TurnStateEvent {
210                session_id: session_id.clone(),
211                channel_name: channel_name.clone(),
212                state: required_string(&payload, "state", "unknown"),
213                previous_state: optional_string(&payload, "previous_state"),
214                data: payload,
215            });
216        })
217    }
218
219    pub fn on_speech_started<F>(&self, handler: F) -> Listener
220    where
221        F: Fn(SpeechStartedEvent) + Send + Sync + 'static,
222    {
223        let session_id = self.session_id.clone();
224        let channel_name = self.channel_name.clone();
225        self.on(EVENT_SPEECH_STARTED, move |payload| {
226            handler(SpeechStartedEvent {
227                session_id: session_id.clone(),
228                channel_name: channel_name.clone(),
229                timestamp_ms: optional_number(&payload, "timestamp_ms"),
230                data: payload,
231            });
232        })
233    }
234
235    pub fn on_speech_stopped<F>(&self, handler: F) -> Listener
236    where
237        F: Fn(SpeechStoppedEvent) + Send + Sync + 'static,
238    {
239        let session_id = self.session_id.clone();
240        let channel_name = self.channel_name.clone();
241        self.on(EVENT_SPEECH_STOPPED, move |payload| {
242            handler(SpeechStoppedEvent {
243                session_id: session_id.clone(),
244                channel_name: channel_name.clone(),
245                timestamp_ms: optional_number(&payload, "timestamp_ms"),
246                data: payload,
247            });
248        })
249    }
250
251    pub fn on_transcript_delta<F>(&self, handler: F) -> Listener
252    where
253        F: Fn(TranscriptDeltaEvent) + Send + Sync + 'static,
254    {
255        let session_id = self.session_id.clone();
256        let channel_name = self.channel_name.clone();
257        self.on(EVENT_TRANSCRIPT_DELTA, move |payload| {
258            handler(TranscriptDeltaEvent {
259                session_id: session_id.clone(),
260                channel_name: channel_name.clone(),
261                delta: required_string(&payload, "delta", ""),
262                start_ms: optional_number(&payload, "start_ms"),
263                end_ms: optional_number(&payload, "end_ms"),
264                data: payload,
265            });
266        })
267    }
268
269    pub fn on_turn_eou_predicted<F>(&self, handler: F) -> Listener
270    where
271        F: Fn(TurnEouPredictedEvent) + Send + Sync + 'static,
272    {
273        let session_id = self.session_id.clone();
274        let channel_name = self.channel_name.clone();
275        self.on(EVENT_TURN_EOU_PREDICTED, move |payload| {
276            handler(TurnEouPredictedEvent {
277                session_id: session_id.clone(),
278                channel_name: channel_name.clone(),
279                probability: optional_number(&payload, "probability"),
280                threshold: optional_number(&payload, "threshold"),
281                delay_ms: optional_number(&payload, "delay_ms"),
282                start_ms: optional_number(&payload, "start_ms"),
283                end_ms: optional_number(&payload, "end_ms"),
284                decision: optional_string(&payload, "decision"),
285                action: optional_string(&payload, "action"),
286                turn_detector: optional_string(&payload, "turn_detector"),
287                data: payload,
288            });
289        })
290    }
291
292    pub fn on_response_created<F>(&self, handler: F) -> Listener
293    where
294        F: Fn(ResponseEvent) + Send + Sync + 'static,
295    {
296        self.on_response_event(EVENT_RESPONSE_CREATED, handler)
297    }
298
299    pub fn on_response_committed<F>(&self, handler: F) -> Listener
300    where
301        F: Fn(ResponseEvent) + Send + Sync + 'static,
302    {
303        self.on_response_event(EVENT_RESPONSE_COMMITTED, handler)
304    }
305
306    pub fn on_response_done<F>(&self, handler: F) -> Listener
307    where
308        F: Fn(ResponseEvent) + Send + Sync + 'static,
309    {
310        self.on_response_event(EVENT_RESPONSE_DONE, handler)
311    }
312
313    pub fn on_response_cancelled<F>(&self, handler: F) -> Listener
314    where
315        F: Fn(ResponseEvent) + Send + Sync + 'static,
316    {
317        self.on_response_event(EVENT_RESPONSE_CANCELLED, handler)
318    }
319
320    pub fn on_response_audio_clear<F>(&self, handler: F) -> Listener
321    where
322        F: Fn(ResponseEvent) + Send + Sync + 'static,
323    {
324        self.on_response_event(EVENT_RESPONSE_AUDIO_CLEAR, handler)
325    }
326
327    fn on_response_event<F>(&self, event_name: &'static str, handler: F) -> Listener
328    where
329        F: Fn(ResponseEvent) + Send + Sync + 'static,
330    {
331        let session_id = self.session_id.clone();
332        let channel_name = self.channel_name.clone();
333        self.on(event_name, move |payload| {
334            handler(response_event(payload, &session_id, &channel_name));
335        })
336    }
337
338    pub fn on_interruption_detected<F>(&self, handler: F) -> Listener
339    where
340        F: Fn(InterruptionEvent) + Send + Sync + 'static,
341    {
342        self.on_interruption_event(EVENT_INTERRUPTION_DETECTED, handler)
343    }
344
345    pub fn on_interruption_false_positive<F>(&self, handler: F) -> Listener
346    where
347        F: Fn(InterruptionEvent) + Send + Sync + 'static,
348    {
349        self.on_interruption_event(EVENT_INTERRUPTION_FALSE_POSITIVE, handler)
350    }
351
352    fn on_interruption_event<F>(&self, event_name: &'static str, handler: F) -> Listener
353    where
354        F: Fn(InterruptionEvent) + Send + Sync + 'static,
355    {
356        let session_id = self.session_id.clone();
357        let channel_name = self.channel_name.clone();
358        self.on(event_name, move |payload| {
359            handler(InterruptionEvent {
360                response: response_event(payload.clone(), &session_id, &channel_name),
361                vad_active_ms: optional_number(&payload, "vad_active_ms"),
362                partial_transcript: optional_string(&payload, "partial_transcript"),
363                reason: optional_nonempty_string(&payload, "reason"),
364            });
365        })
366    }
367
368    pub fn on_browser_event<F>(&self, handler: F) -> Listener
369    where
370        F: Fn(BrowserEvent) + Send + Sync + 'static,
371    {
372        let session_id = self.session_id.clone();
373        let channel_name = self.channel_name.clone();
374        self.on(EVENT_BROWSER_EVENT, move |payload| {
375            handler(BrowserEvent {
376                session_id: session_id.clone(),
377                channel_name: channel_name.clone(),
378                event: required_string(&payload, "event", ""),
379                payload: payload.get("payload").cloned().unwrap_or(Value::Null),
380                data: payload,
381            });
382        })
383    }
384
385    pub fn on_close<F>(&self, handler: F) -> Listener
386    where
387        F: Fn(CloseEvent) + Send + Sync + 'static,
388    {
389        let session_id = self.session_id.clone();
390        let channel_name = self.channel_name.clone();
391        self.on(EVENT_RTC_CLIENT_DISCONNECTED, move |payload| {
392            handler(CloseEvent {
393                session_id: session_id.clone(),
394                channel_name: channel_name.clone(),
395                reason: required_string(&payload, "reason", "unknown"),
396                connection_state: optional_string(&payload, "connection_state"),
397                ice_connection_state: optional_string(&payload, "ice_connection_state"),
398                data_channel_state: optional_string(&payload, "data_channel_state"),
399                data: payload,
400            });
401        })
402    }
403
404    pub fn on_error<F>(&self, handler: F) -> Listener
405    where
406        F: Fn(ErrorEvent) + Send + Sync + 'static,
407    {
408        let session_id = self.session_id.clone();
409        let channel_name = self.channel_name.clone();
410        self.on(EVENT_ERROR, move |payload| {
411            handler(ErrorEvent {
412                session_id: session_id.clone(),
413                channel_name: channel_name.clone(),
414                message: optional_string(&payload, "message"),
415                code: optional_nonempty_string(&payload, "code"),
416                recoverable: recoverable_flag(&payload),
417                generation_id: optional_nonempty_string(&payload, "generation_id"),
418                data: payload,
419            });
420        })
421    }
422
423    pub fn on_signaling_error<F>(&self, handler: F) -> Listener
424    where
425        F: Fn(SignalingErrorEvent) + Send + Sync + 'static,
426    {
427        let session_id = self.session_id.clone();
428        let channel_name = self.channel_name.clone();
429        self.on(EVENT_RTC_SIGNALING_ERROR, move |payload| {
430            handler(SignalingErrorEvent {
431                session_id: session_id.clone(),
432                channel_name: channel_name.clone(),
433                message: optional_string(&payload, "message"),
434                generation: optional_i64(&payload, "generation"),
435                data: payload,
436            });
437        })
438    }
439
440    pub async fn send_control(&self, event: &str, payload: EventData) -> Result<()> {
441        self.channel.send_message(event, payload).await
442    }
443
444    pub async fn configure(&self, config: SessionConfig) -> Result<()> {
445        let mut payload = EventData::new();
446        payload.insert(
447            "session".to_owned(),
448            Value::Object(session_config_payload(config)),
449        );
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 session_config_payload(config: SessionConfig) -> EventData {
667    let mut session = config.extra;
668    insert_opt(&mut session, "stt_model", config.stt_model);
669    insert_opt(&mut session, "tts_model", config.tts_model);
670    insert_opt(&mut session, "voice", config.voice);
671    insert_opt(&mut session, "turn_profile", config.turn_profile);
672    insert_opt(&mut session, "vad_backend", config.vad_backend);
673    insert_opt(&mut session, "turn_detector", config.turn_detector);
674    if let Some(enabled) = config.speech_context {
675        session.insert("speech_context".to_owned(), Value::Bool(enabled));
676    }
677    session
678}
679
680fn response_options_payload(options: Option<ResponseOptions>) -> EventData {
681    let mut payload = EventData::new();
682    if let Some(options) = options
683        && let Some(allow) = options.allow_interruptions
684    {
685        payload.insert("allow_interruptions".to_owned(), Value::Bool(allow));
686    }
687    payload
688}
689
690fn explicit_generation(options: &Option<ResponseOptions>) -> Option<String> {
691    options
692        .as_ref()
693        .and_then(|options| options.generation_id.clone())
694        .filter(|id| !id.is_empty())
695}
696
697fn response_event(payload: EventData, session_id: &str, channel_name: &str) -> ResponseEvent {
698    ResponseEvent {
699        session_id: session_id.to_owned(),
700        channel_name: channel_name.to_owned(),
701        response_id: optional_string(&payload, "response_id"),
702        generation_id: optional_nonempty_string(&payload, "generation_id"),
703        data: payload,
704    }
705}
706
707#[cfg(test)]
708mod tests {
709    use super::*;
710    use crate::socket::test_channel;
711    use serde_json::json;
712    use tokio::sync::broadcast;
713    use tokio::sync::mpsc;
714
715    async fn session() -> (VoxRtcControlSession, broadcast::Sender<(String, EventData)>) {
716        let (channel, sender) = test_channel().await;
717        let session =
718            VoxRtcControlSession::new(channel, "sess-1".to_owned(), Duration::from_secs(1));
719        (session, sender)
720    }
721
722    fn payload(value: Value) -> EventData {
723        value.as_object().cloned().expect("object payload")
724    }
725
726    #[test]
727    fn join_decline_reason_prefers_structured_message_fields() {
728        assert_eq!(
729            join_decline_reason(Some(payload(json!({ "message": "expired" })))),
730            Some("expired".to_owned())
731        );
732        assert_eq!(
733            join_decline_reason(Some(payload(json!({ "reason": "missing" })))),
734            Some("missing".to_owned())
735        );
736        assert_eq!(
737            join_decline_reason(Some(payload(json!({ "channel": "/rtc/abc" })))),
738            Some(r#"{"channel":"/rtc/abc"}"#.to_owned())
739        );
740        assert_eq!(join_decline_reason(Some(EventData::new())), None);
741        assert_eq!(join_decline_reason(None), None);
742    }
743
744    async fn recv<T>(rx: &mut mpsc::UnboundedReceiver<T>) -> T {
745        timeout(Duration::from_secs(1), rx.recv())
746            .await
747            .expect("handler fired within timeout")
748            .expect("handler produced an event")
749    }
750
751    #[test]
752    fn next_message_classifies_lag_close_and_ok() {
753        assert!(matches!(
754            next_message(Ok(("e".to_owned(), EventData::new()))),
755            ControlFlow::Continue(Some(_))
756        ));
757        assert!(matches!(
758            next_message(Err(RecvError::Lagged(7))),
759            ControlFlow::Continue(None)
760        ));
761        assert!(matches!(
762            next_message(Err(RecvError::Closed)),
763            ControlFlow::Break(())
764        ));
765    }
766
767    #[test]
768    fn session_config_serializes_explicit_false_speech_context() {
769        let payload = session_config_payload(SessionConfig {
770            speech_context: Some(false),
771            ..Default::default()
772        });
773        assert_eq!(payload.get("speech_context"), Some(&Value::Bool(false)));
774    }
775
776    #[tokio::test]
777    async fn response_commands_share_one_generation_id() {
778        let (session, _) = session().await;
779        let generation_id = session.next_response_generation();
780        let mut delta = payload(json!({ "delta": "hello" }));
781        session.add_response_generation(&mut delta);
782        let mut commit = EventData::new();
783        session.add_response_generation(&mut commit);
784
785        assert_eq!(
786            delta.get("generation_id"),
787            Some(&Value::String(generation_id.clone()))
788        );
789        assert_eq!(
790            commit.get("generation_id"),
791            Some(&Value::String(generation_id))
792        );
793    }
794
795    #[tokio::test]
796    async fn on_error_parses_typed_fields() {
797        let (session, sender) = session().await;
798        let (tx, mut rx) = mpsc::unbounded_channel();
799        let _listener = session.on_error(move |event| {
800            tx.send(event).unwrap();
801        });
802        sender
803            .send((
804                EVENT_ERROR.to_owned(),
805                payload(json!({
806                    "message": "cannot start now",
807                    "code": ERROR_CODE_SESSION_FAILED,
808                    "recoverable": false,
809                    "generation_id": "gen-9"
810                })),
811            ))
812            .unwrap();
813        let event = recv(&mut rx).await;
814        assert_eq!(event.message.as_deref(), Some("cannot start now"));
815        assert_eq!(event.code.as_deref(), Some(ERROR_CODE_SESSION_FAILED));
816        assert!(!event.recoverable);
817        assert_eq!(event.generation_id.as_deref(), Some("gen-9"));
818    }
819
820    #[tokio::test]
821    async fn on_error_defaults_missing_recoverable_to_true() {
822        let (session, sender) = session().await;
823        let (tx, mut rx) = mpsc::unbounded_channel();
824        let _listener = session.on_error(move |event| {
825            tx.send(event).unwrap();
826        });
827        sender
828            .send((
829                EVENT_ERROR.to_owned(),
830                payload(json!({ "message": "legacy server", "code": "" })),
831            ))
832            .unwrap();
833        let event = recv(&mut rx).await;
834        assert!(event.recoverable);
835        assert_eq!(event.code, None);
836        assert_eq!(event.generation_id, None);
837    }
838
839    #[tokio::test]
840    async fn start_payload_uses_explicit_generation_id() {
841        let (session, _) = session().await;
842        let options = ResponseOptions {
843            allow_interruptions: Some(false),
844            generation_id: Some("gen-7".to_owned()),
845        };
846        let (generation_id, start) = session.start_payload(Some(options));
847        assert_eq!(generation_id, "gen-7");
848        assert_eq!(
849            start.get("generation_id"),
850            Some(&Value::String("gen-7".to_owned()))
851        );
852        assert_eq!(start.get("allow_interruptions"), Some(&Value::Bool(false)));
853
854        let mut commit = EventData::new();
855        session.thread_generation(&mut commit, None);
856        assert_eq!(
857            commit.get("generation_id"),
858            Some(&Value::String("gen-7".to_owned()))
859        );
860    }
861
862    #[tokio::test]
863    async fn start_payload_generates_generation_id_when_absent() {
864        let (session, _) = session().await;
865        let (generation_id, start) = session.start_payload(None);
866        assert!(generation_id.starts_with("generation_1_"));
867        assert!(generation_id.len() > "generation_1_".len());
868        assert_eq!(
869            start.get("generation_id"),
870            Some(&Value::String(generation_id))
871        );
872    }
873
874    #[tokio::test]
875    async fn explicit_generation_id_overrides_tracked_one() {
876        let (session, _) = session().await;
877        let tracked = session.next_response_generation();
878        let mut delta = payload(json!({ "delta": "hi" }));
879        session.thread_generation(&mut delta, Some("gen-42".to_owned()));
880        assert_eq!(
881            delta.get("generation_id"),
882            Some(&Value::String("gen-42".to_owned()))
883        );
884        assert_ne!(tracked, "gen-42");
885    }
886
887    #[tokio::test]
888    async fn response_events_expose_generation_id() {
889        let (session, sender) = session().await;
890        let (tx, mut rx) = mpsc::unbounded_channel();
891        let _listener = session.on_response_created(move |event| {
892            tx.send(event).unwrap();
893        });
894        sender
895            .send((
896                EVENT_RESPONSE_CREATED.to_owned(),
897                payload(json!({ "response_id": "resp-1", "generation_id": "gen-1" })),
898            ))
899            .unwrap();
900        let event = recv(&mut rx).await;
901        assert_eq!(event.response_id.as_deref(), Some("resp-1"));
902        assert_eq!(event.generation_id.as_deref(), Some("gen-1"));
903    }
904
905    #[tokio::test]
906    async fn audio_clear_and_interruption_expose_generation_id() {
907        let (session, sender) = session().await;
908        let (clear_tx, mut clear_rx) = mpsc::unbounded_channel();
909        let _clear = session.on_response_audio_clear(move |event| {
910            clear_tx.send(event).unwrap();
911        });
912        let (int_tx, mut int_rx) = mpsc::unbounded_channel();
913        let _interruption = session.on_interruption_detected(move |event| {
914            int_tx.send(event).unwrap();
915        });
916        sender
917            .send((
918                EVENT_RESPONSE_AUDIO_CLEAR.to_owned(),
919                payload(json!({ "response_id": "resp-2", "generation_id": "gen-2" })),
920            ))
921            .unwrap();
922        sender
923            .send((
924                EVENT_INTERRUPTION_DETECTED.to_owned(),
925                payload(json!({
926                    "response_id": "resp-2",
927                    "generation_id": "gen-2",
928                    "vad_active_ms": 250
929                })),
930            ))
931            .unwrap();
932        let clear = recv(&mut clear_rx).await;
933        assert_eq!(clear.generation_id.as_deref(), Some("gen-2"));
934        let interruption = recv(&mut int_rx).await;
935        assert_eq!(interruption.response.generation_id.as_deref(), Some("gen-2"));
936        assert_eq!(interruption.vad_active_ms, Some(250.0));
937    }
938
939    #[tokio::test]
940    async fn on_signaling_error_parses_message_and_generation() {
941        let (session, sender) = session().await;
942        let (tx, mut rx) = mpsc::unbounded_channel();
943        let _listener = session.on_signaling_error(move |event| {
944            tx.send(event).unwrap();
945        });
946        sender
947            .send((
948                EVENT_RTC_SIGNALING_ERROR.to_owned(),
949                payload(json!({
950                    "message": "setLocalDescription failed",
951                    "generation": 3
952                })),
953            ))
954            .unwrap();
955        let event = recv(&mut rx).await;
956        assert_eq!(event.message.as_deref(), Some("setLocalDescription failed"));
957        assert_eq!(event.generation, Some(3));
958    }
959
960    #[tokio::test]
961    async fn on_signaling_error_leaves_generation_none_when_absent() {
962        let (session, sender) = session().await;
963        let (tx, mut rx) = mpsc::unbounded_channel();
964        let _listener = session.on_signaling_error(move |event| {
965            tx.send(event).unwrap();
966        });
967        sender
968            .send((
969                EVENT_RTC_SIGNALING_ERROR.to_owned(),
970                payload(json!({ "message": "RTC signaling failed" })),
971            ))
972            .unwrap();
973        let event = recv(&mut rx).await;
974        assert_eq!(event.message.as_deref(), Some("RTC signaling failed"));
975        assert_eq!(event.generation, None);
976    }
977
978    #[tokio::test]
979    async fn on_transcript_exposes_entities_and_words() {
980        let (session, sender) = session().await;
981        let (tx, mut rx) = mpsc::unbounded_channel();
982        let _listener = session.on_transcript(move |event| {
983            tx.send(event).unwrap();
984        });
985        sender
986            .send((
987                EVENT_TRANSCRIPT_COMPLETED.to_owned(),
988                payload(json!({
989                    "transcript": "call Ada",
990                    "entities": [
991                        { "type": "PRODUCT", "text": "Ada", "start_char": 5, "end_char": 8 }
992                    ],
993                    "words": [
994                        { "word": "call", "start_ms": 0, "end_ms": 300 },
995                        { "word": "Ada", "start_ms": 300, "end_ms": 600, "confidence": 0.91 }
996                    ],
997                    "speech_context": serde_json::from_str::<Value>(include_str!(
998                        "../../../fixtures/speech-context-v2.json"
999                    )).unwrap()
1000                })),
1001            ))
1002            .unwrap();
1003        let event = recv(&mut rx).await;
1004        assert_eq!(
1005            event.entities,
1006            vec![TranscriptEntity {
1007                r#type: "PRODUCT".to_owned(),
1008                text: "Ada".to_owned(),
1009                start_char: 5,
1010                end_char: 8,
1011            }]
1012        );
1013        assert_eq!(event.words.len(), 2);
1014        assert_eq!(event.words[0].word, "call");
1015        assert_eq!(event.words[0].start_ms, 0.0);
1016        assert_eq!(event.words[0].confidence, None);
1017        assert_eq!(event.words[1].confidence, Some(0.91));
1018        let context = event.speech_context.expect("speech context");
1019        assert_eq!(context.schema_version, 2);
1020        assert_eq!(context.status, SpeechContextStatus::Complete);
1021        assert_eq!(
1022            context.emotions.as_deref(),
1023            Some(
1024                &[SpeechContextSpan {
1025                    label: "surprised".to_owned(),
1026                    start_ms: 0,
1027                    end_ms: 2500,
1028                }][..]
1029            )
1030        );
1031        let sounds = context.sounds.expect("sound spans");
1032        assert_eq!(sounds.len(), 2);
1033        assert_eq!(sounds[0].span.label, "fireworks");
1034        assert_eq!(sounds[0].score, 0.42);
1035    }
1036
1037    #[tokio::test]
1038    async fn on_transcript_defaults_entities_and_words_to_empty() {
1039        let (session, sender) = session().await;
1040        let (tx, mut rx) = mpsc::unbounded_channel();
1041        let _listener = session.on_transcript(move |event| {
1042            tx.send(event).unwrap();
1043        });
1044        sender
1045            .send((
1046                EVENT_TRANSCRIPT_COMPLETED.to_owned(),
1047                payload(json!({ "transcript": "hello" })),
1048            ))
1049            .unwrap();
1050        let event = recv(&mut rx).await;
1051        assert!(event.entities.is_empty());
1052        assert!(event.words.is_empty());
1053        assert!(event.speech_context.is_none());
1054    }
1055
1056    #[tokio::test]
1057    async fn on_transcript_preserves_text_but_rejects_malformed_speech_context() {
1058        let (session, sender) = session().await;
1059        let (tx, mut rx) = mpsc::unbounded_channel();
1060        let _listener = session.on_transcript(move |event| {
1061            tx.send(event).unwrap();
1062        });
1063        sender
1064            .send((
1065                EVENT_TRANSCRIPT_COMPLETED.to_owned(),
1066                payload(json!({
1067                    "transcript": "still delivered",
1068                    "speech_context": {
1069                        "schema_version": 2,
1070                        "status": "complete",
1071                        "emotions": [],
1072                        "vocal": [],
1073                        "sounds": [
1074                            {
1075                                "label": "fireworks",
1076                                "start_ms": 0,
1077                                "end_ms": 960,
1078                                "score": 1.1
1079                            }
1080                        ]
1081                    }
1082                })),
1083            ))
1084            .unwrap();
1085        let event = recv(&mut rx).await;
1086        assert_eq!(event.transcript, "still delivered");
1087        assert!(event.speech_context.is_none());
1088    }
1089
1090    #[tokio::test]
1091    async fn interruption_events_expose_reason() {
1092        let (session, sender) = session().await;
1093        let (det_tx, mut det_rx) = mpsc::unbounded_channel();
1094        let _detected = session.on_interruption_detected(move |event| {
1095            det_tx.send(event).unwrap();
1096        });
1097        let (fp_tx, mut fp_rx) = mpsc::unbounded_channel();
1098        let _false_positive = session.on_interruption_false_positive(move |event| {
1099            fp_tx.send(event).unwrap();
1100        });
1101        sender
1102            .send((
1103                EVENT_INTERRUPTION_DETECTED.to_owned(),
1104                payload(json!({
1105                    "response_id": "resp-3",
1106                    "generation_id": "gen-3",
1107                    "reason": "speech_overlap"
1108                })),
1109            ))
1110            .unwrap();
1111        sender
1112            .send((
1113                EVENT_INTERRUPTION_FALSE_POSITIVE.to_owned(),
1114                payload(json!({ "response_id": "resp-3", "reason": "backchannel" })),
1115            ))
1116            .unwrap();
1117        let detected = recv(&mut det_rx).await;
1118        assert_eq!(detected.reason.as_deref(), Some("speech_overlap"));
1119        let false_positive = recv(&mut fp_rx).await;
1120        assert_eq!(false_positive.reason.as_deref(), Some("backchannel"));
1121    }
1122
1123    #[tokio::test]
1124    async fn start_response_and_wait_resolves_on_matching_created() {
1125        let (session, sender) = session().await;
1126        let options = ResponseOptions {
1127            generation_id: Some("gen-ack".to_owned()),
1128            ..Default::default()
1129        };
1130        tokio::spawn(async move {
1131            tokio::time::sleep(Duration::from_millis(100)).await;
1132            sender
1133                .send((
1134                    EVENT_RESPONSE_CREATED.to_owned(),
1135                    payload(json!({ "response_id": "resp-other", "generation_id": "gen-other" })),
1136                ))
1137                .unwrap();
1138            sender
1139                .send((
1140                    EVENT_RESPONSE_CREATED.to_owned(),
1141                    payload(json!({ "response_id": "resp-9", "generation_id": "gen-ack" })),
1142                ))
1143                .unwrap();
1144        });
1145        let ack = session
1146            .start_response_and_wait(Some(options), Duration::from_secs(2))
1147            .await
1148            .expect("ack within timeout");
1149        assert!(ack.accepted);
1150        assert_eq!(ack.generation_id, "gen-ack");
1151        assert_eq!(ack.response_id.as_deref(), Some("resp-9"));
1152        assert!(ack.recoverable);
1153        assert_eq!(ack.error_code, None);
1154    }
1155
1156    #[tokio::test]
1157    async fn start_response_and_wait_surfaces_typed_rejection() {
1158        let (session, sender) = session().await;
1159        let options = ResponseOptions {
1160            generation_id: Some("gen-rejected".to_owned()),
1161            ..Default::default()
1162        };
1163        tokio::spawn(async move {
1164            tokio::time::sleep(Duration::from_millis(100)).await;
1165            sender
1166                .send((
1167                    EVENT_ERROR.to_owned(),
1168                    payload(json!({
1169                        "message": "busy",
1170                        "code": ERROR_CODE_RESPONSE_ALREADY_ACTIVE,
1171                        "recoverable": true,
1172                        "generation_id": "gen-rejected"
1173                    })),
1174                ))
1175                .unwrap();
1176        });
1177        let ack = session
1178            .start_response_and_wait(Some(options), Duration::from_secs(2))
1179            .await
1180            .expect("rejection within timeout");
1181        assert!(!ack.accepted);
1182        assert_eq!(ack.generation_id, "gen-rejected");
1183        assert_eq!(
1184            ack.error_code.as_deref(),
1185            Some(ERROR_CODE_RESPONSE_ALREADY_ACTIVE)
1186        );
1187        assert_eq!(ack.error_message.as_deref(), Some("busy"));
1188        assert!(ack.recoverable);
1189    }
1190
1191    #[tokio::test]
1192    async fn start_response_and_wait_times_out_without_ack() {
1193        let (session, _sender) = session().await;
1194        let error = session
1195            .start_response_and_wait(None, Duration::from_millis(100))
1196            .await
1197            .expect_err("no ack must time out");
1198        assert!(matches!(error, VoxRtcError::Timeout(_)));
1199    }
1200
1201    #[tokio::test]
1202    async fn on_speech_started_fires_with_timestamp() {
1203        let (session, sender) = session().await;
1204        let (tx, mut rx) = mpsc::unbounded_channel();
1205        let _listener = session.on_speech_started(move |event| {
1206            tx.send(event).unwrap();
1207        });
1208        sender
1209            .send((
1210                EVENT_SPEECH_STARTED.to_owned(),
1211                payload(json!({ "session_id": "sess-1", "timestamp_ms": 1234 })),
1212            ))
1213            .unwrap();
1214        let event = recv(&mut rx).await;
1215        assert_eq!(event.session_id, "sess-1");
1216        assert_eq!(event.channel_name, "/rtc/sess-1");
1217        assert_eq!(event.timestamp_ms, Some(1234.0));
1218    }
1219
1220    #[tokio::test]
1221    async fn on_speech_stopped_fires_with_timestamp() {
1222        let (session, sender) = session().await;
1223        let (tx, mut rx) = mpsc::unbounded_channel();
1224        let _listener = session.on_speech_stopped(move |event| {
1225            tx.send(event).unwrap();
1226        });
1227        sender
1228            .send((
1229                EVENT_SPEECH_STOPPED.to_owned(),
1230                payload(json!({ "timestamp_ms": 5678 })),
1231            ))
1232            .unwrap();
1233        let event = recv(&mut rx).await;
1234        assert_eq!(event.timestamp_ms, Some(5678.0));
1235    }
1236
1237    #[tokio::test]
1238    async fn on_transcript_delta_fires_with_fields() {
1239        let (session, sender) = session().await;
1240        let (tx, mut rx) = mpsc::unbounded_channel();
1241        let _listener = session.on_transcript_delta(move |event| {
1242            tx.send(event).unwrap();
1243        });
1244        sender
1245            .send((
1246                EVENT_TRANSCRIPT_DELTA.to_owned(),
1247                payload(json!({ "delta": "hel", "start_ms": 10, "end_ms": 20 })),
1248            ))
1249            .unwrap();
1250        let event = recv(&mut rx).await;
1251        assert_eq!(event.delta, "hel");
1252        assert_eq!(event.start_ms, Some(10.0));
1253        assert_eq!(event.end_ms, Some(20.0));
1254    }
1255
1256    #[tokio::test]
1257    async fn on_turn_eou_predicted_fires_with_fields() {
1258        let (session, sender) = session().await;
1259        let (tx, mut rx) = mpsc::unbounded_channel();
1260        let _listener = session.on_turn_eou_predicted(move |event| {
1261            tx.send(event).unwrap();
1262        });
1263        sender
1264            .send((
1265                EVENT_TURN_EOU_PREDICTED.to_owned(),
1266                payload(json!({
1267                    "probability": 0.82,
1268                    "threshold": 0.5,
1269                    "delay_ms": 120,
1270                    "start_ms": 0,
1271                    "end_ms": 300,
1272                    "decision": "end",
1273                    "action": "commit",
1274                    "turn_detector": "smart"
1275                })),
1276            ))
1277            .unwrap();
1278        let event = recv(&mut rx).await;
1279        assert_eq!(event.probability, Some(0.82));
1280        assert_eq!(event.threshold, Some(0.5));
1281        assert_eq!(event.delay_ms, Some(120.0));
1282        assert_eq!(event.start_ms, Some(0.0));
1283        assert_eq!(event.end_ms, Some(300.0));
1284        assert_eq!(event.decision.as_deref(), Some("end"));
1285        assert_eq!(event.action.as_deref(), Some("commit"));
1286        assert_eq!(event.turn_detector.as_deref(), Some("smart"));
1287    }
1288
1289    #[tokio::test]
1290    async fn handler_survives_a_lagged_broadcast() {
1291        let (session, sender) = session().await;
1292        let (tx, mut rx) = mpsc::unbounded_channel();
1293        let _listener = session.on_speech_started(move |event| {
1294            tx.send(event.timestamp_ms).unwrap();
1295        });
1296
1297        for index in 0..2100u32 {
1298            let _ = sender.send((
1299                EVENT_SPEECH_STARTED.to_owned(),
1300                payload(json!({ "timestamp_ms": index })),
1301            ));
1302        }
1303        let _ = sender.send((
1304            EVENT_SPEECH_STARTED.to_owned(),
1305            payload(json!({ "timestamp_ms": 9999 })),
1306        ));
1307
1308        let mut saw_final = false;
1309        while let Ok(Some(value)) = timeout(Duration::from_secs(1), rx.recv()).await {
1310            if value == Some(9999.0) {
1311                saw_final = true;
1312                break;
1313            }
1314        }
1315        assert!(
1316            saw_final,
1317            "loop must keep delivering events after a broadcast lag"
1318        );
1319    }
1320}