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