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