Skip to main content

gpui_component/speech/
state.rs

1use std::{cell::OnceCell, rc::Rc, time::Duration};
2
3use gpui::{App, Context, EventEmitter, SharedString, Subscription, Task, WeakEntity};
4
5use super::{
6    AudioInput, AudioSink, RecognitionSession, SpeechError, SpeechRecognizer, SpeechSink,
7    level::LevelMeter,
8};
9
10/// Default time [`SpeechState::stop`] waits for the final result.
11const DEFAULT_STOP_TIMEOUT: Duration = Duration::from_secs(3);
12
13/// Where a [`SpeechState`] is in its session.
14#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
15pub enum SpeechStatus {
16    /// No session is running.
17    #[default]
18    Idle,
19    /// Audio is being captured while the recognizer connects.
20    Connecting,
21    /// Audio is being captured and recognized.
22    Recording,
23    /// Capture stopped; waiting for the recognizer's final result.
24    Stopping,
25}
26
27impl SpeechStatus {
28    /// Whether a session is running.
29    pub fn is_active(self) -> bool {
30        self != Self::Idle
31    }
32
33    /// Whether the microphone is capturing.
34    pub fn is_capturing(self) -> bool {
35        matches!(self, Self::Connecting | Self::Recording)
36    }
37}
38
39/// Events emitted by [`SpeechState`].
40#[derive(Debug, Clone)]
41pub enum SpeechEvent {
42    /// A session started and audio is being captured.
43    Started,
44    /// The transcript changed: every committed phrase followed by the current
45    /// hypothesis. Later events supersede earlier ones.
46    Partial(SharedString),
47    /// The session ended normally with this transcript, possibly empty.
48    Final(SharedString),
49    /// The session was cancelled and its transcript discarded.
50    Cancelled,
51    /// The session failed and ended.
52    Error(SpeechError),
53}
54
55struct Session {
56    id: usize,
57    /// Dropping this stops capture.
58    capture: Option<Subscription>,
59    recognition: Box<dyn RecognitionSession>,
60    _stop_timeout: Option<Task<()>>,
61}
62
63/// The state of a speech input: captures audio from an [`AudioInput`], feeds it
64/// to a [`SpeechRecognizer`] and tracks the transcript.
65///
66/// Render it with [`SpeechButton`](super::SpeechButton) and
67/// [`SpeechWaveform`](super::SpeechWaveform), and subscribe to [`SpeechEvent`]
68/// to receive the text.
69///
70/// The recognizer is, in order: the one passed to [`Self::recognizer`]; else
71/// the platform's `SystemRecognizer`, unless
72/// [`Self::system_fallback`] turned it off; else none, and the state is not
73/// available. The input defaults to the `Microphone`.
74/// Both defaults need the `speech` feature.
75pub struct SpeechState {
76    recognizer: Option<Rc<dyn SpeechRecognizer>>,
77    input: Option<Rc<dyn AudioInput>>,
78    system_fallback: bool,
79    system_recognizer: OnceCell<Option<Rc<dyn SpeechRecognizer>>>,
80    stop_timeout: Duration,
81    status: SpeechStatus,
82    session: Option<Session>,
83    next_session: usize,
84    committed: String,
85    hypothesis: SharedString,
86    meter: LevelMeter,
87}
88
89impl EventEmitter<SpeechEvent> for SpeechState {}
90
91impl SpeechState {
92    /// Create a speech state with the default recognizer and input.
93    pub fn new(_: &mut Context<Self>) -> Self {
94        Self {
95            recognizer: None,
96            input: super::default_input(),
97            system_fallback: true,
98            system_recognizer: OnceCell::new(),
99            stop_timeout: DEFAULT_STOP_TIMEOUT,
100            status: SpeechStatus::Idle,
101            session: None,
102            next_session: 0,
103            committed: String::new(),
104            hypothesis: SharedString::default(),
105            meter: LevelMeter::new(),
106        }
107    }
108
109    /// Recognize speech with `recognizer` instead of the system's.
110    pub fn recognizer(mut self, recognizer: impl SpeechRecognizer) -> Self {
111        self.recognizer = Some(Rc::new(recognizer));
112        self
113    }
114
115    /// Capture audio from `input` instead of the microphone.
116    pub fn input(mut self, input: impl AudioInput) -> Self {
117        self.input = Some(Rc::new(input));
118        self
119    }
120
121    /// Whether to fall back to the platform's recognizer when no
122    /// [`Self::recognizer`] is set, default `true`.
123    ///
124    /// On Windows the system recognizer dictates through Microsoft's online
125    /// service; turn this off when audio must not leave the application.
126    pub fn system_fallback(mut self, system_fallback: bool) -> Self {
127        self.system_fallback = system_fallback;
128        self
129    }
130
131    /// Set how long [`Self::stop`] waits for the recognizer's final result
132    /// before ending the session with the transcript so far, default 3 seconds.
133    pub fn stop_timeout(mut self, timeout: Duration) -> Self {
134        self.stop_timeout = timeout;
135        self
136    }
137
138    /// Whether a recognizer and an input are configured, so that speech input
139    /// can work on this platform at all.
140    pub fn has_recognizer(&self) -> bool {
141        self.input.is_some() && self.active_recognizer().is_some()
142    }
143
144    /// Whether a session can start now: [`Self::has_recognizer`] and the
145    /// recognizer reports itself available.
146    pub fn is_available(&self, cx: &App) -> bool {
147        self.input.is_some()
148            && self
149                .active_recognizer()
150                .is_some_and(|recognizer| recognizer.is_available(cx))
151    }
152
153    /// Where the state is in its session.
154    pub fn status(&self) -> SpeechStatus {
155        self.status
156    }
157
158    /// The transcript of the current or last session: every committed phrase
159    /// followed by the current hypothesis.
160    pub fn transcript(&self) -> SharedString {
161        if self.hypothesis.is_empty() {
162            self.committed.clone().into()
163        } else {
164            format!("{}{}", self.committed, self.hypothesis).into()
165        }
166    }
167
168    /// Recent input levels in `0.0..=1.0`, oldest first, one per 80 ms of
169    /// audio. Peaks rise at once and fall back smoothly; background noise reads
170    /// as `0.0`.
171    pub fn levels(&self) -> impl ExactSizeIterator<Item = f32> + '_ {
172        self.meter.levels()
173    }
174
175    /// How far past the waveform's trailing edge the newest level sits at
176    /// `now`, in levels; see `LevelMeter::lead_at`.
177    pub(super) fn level_lead_at(&self, now: instant::Instant) -> Option<f32> {
178        self.meter.lead_at(now)
179    }
180
181    /// Start a session. Does nothing while one is running.
182    ///
183    /// When the recognizer or the input fails to start, emits
184    /// [`SpeechEvent::Error`] and stays idle.
185    pub fn start(&mut self, cx: &mut Context<Self>) {
186        if self.status.is_active() {
187            return;
188        }
189
190        let (Some(recognizer), Some(input)) = (self.active_recognizer(), self.input.clone()) else {
191            cx.emit(SpeechEvent::Error(SpeechError::Unsupported));
192            return;
193        };
194
195        self.next_session += 1;
196        let id = self.next_session;
197        let state = cx.weak_entity();
198        let sink = SpeechSink {
199            state: state.clone(),
200            session: id,
201        };
202        let recognition = match recognizer.start(sink, cx) {
203            Ok(recognition) => recognition,
204            Err(error) => {
205                cx.emit(SpeechEvent::Error(error));
206                return;
207            }
208        };
209
210        let format = recognizer.audio_format();
211        let sink = AudioSink { state, session: id };
212        let capture = match input.start(format, sink, cx) {
213            Ok(capture) => capture,
214            Err(error) => {
215                cx.emit(SpeechEvent::Error(error));
216                return;
217            }
218        };
219
220        self.status = SpeechStatus::Connecting;
221        self.committed.clear();
222        self.hypothesis = SharedString::default();
223        self.meter.reset(format.sample_rate(), format.channels());
224        self.session = Some(Session {
225            id,
226            capture: Some(capture),
227            recognition,
228            _stop_timeout: None,
229        });
230        cx.emit(SpeechEvent::Started);
231        cx.notify();
232    }
233
234    /// Stop capturing and wait for the final result, which arrives as
235    /// [`SpeechEvent::Final`].
236    pub fn stop(&mut self, cx: &mut Context<Self>) {
237        if !self.status.is_capturing() {
238            return;
239        }
240        let Some(session) = self.session.as_mut() else {
241            return;
242        };
243
244        session.capture = None;
245        session.recognition.finish(cx);
246
247        let id = session.id;
248        let timeout = self.stop_timeout;
249        session._stop_timeout = Some(cx.spawn(async move |this, cx| {
250            cx.background_executor().timer(timeout).await;
251            _ = this.update(cx, |this, cx| {
252                if this.is_session(id) {
253                    this.end(SpeechEvent::Final(this.transcript()), cx);
254                }
255            });
256        }));
257
258        self.status = SpeechStatus::Stopping;
259        cx.notify();
260    }
261
262    /// End the session at once and discard its transcript.
263    pub fn cancel(&mut self, cx: &mut Context<Self>) {
264        if !self.status.is_active() {
265            return;
266        }
267        self.committed.clear();
268        self.hypothesis = SharedString::default();
269        self.end(SpeechEvent::Cancelled, cx);
270    }
271
272    /// Start a session when idle, otherwise stop the running one.
273    pub fn toggle(&mut self, cx: &mut Context<Self>) {
274        if self.status.is_active() {
275            self.stop(cx);
276        } else {
277            self.start(cx);
278        }
279    }
280
281    fn active_recognizer(&self) -> Option<Rc<dyn SpeechRecognizer>> {
282        if let Some(recognizer) = &self.recognizer {
283            return Some(recognizer.clone());
284        }
285        if !self.system_fallback {
286            return None;
287        }
288        self.system_recognizer
289            .get_or_init(super::system_recognizer)
290            .clone()
291    }
292
293    pub(super) fn is_session(&self, id: usize) -> bool {
294        self.session
295            .as_ref()
296            .is_some_and(|session| session.id == id)
297    }
298
299    pub(super) fn on_ready(&mut self, cx: &mut Context<Self>) {
300        if self.status == SpeechStatus::Connecting {
301            self.status = SpeechStatus::Recording;
302            cx.notify();
303        }
304    }
305
306    pub(super) fn on_audio(&mut self, samples: &[i16], cx: &mut Context<Self>) {
307        if !self.status.is_capturing() {
308            return;
309        }
310        if let Some(session) = self.session.as_mut() {
311            session.recognition.push_audio(samples, cx);
312        }
313        // Redraw once per new level, not per push: levels are what the
314        // waveform shows.
315        if self.meter.push(samples) {
316            cx.notify();
317        }
318    }
319
320    pub(super) fn on_hypothesis(&mut self, text: SharedString, cx: &mut Context<Self>) {
321        self.hypothesis = text;
322        cx.emit(SpeechEvent::Partial(self.transcript()));
323        cx.notify();
324    }
325
326    pub(super) fn on_phrase(&mut self, text: SharedString, cx: &mut Context<Self>) {
327        self.committed.push_str(&text);
328        self.hypothesis = SharedString::default();
329        cx.emit(SpeechEvent::Partial(self.transcript()));
330        cx.notify();
331    }
332
333    pub(super) fn on_finish(&mut self, cx: &mut Context<Self>) {
334        self.end(SpeechEvent::Final(self.transcript()), cx);
335    }
336
337    pub(super) fn on_error(&mut self, error: SpeechError, cx: &mut Context<Self>) {
338        self.end(SpeechEvent::Error(error), cx);
339    }
340
341    /// Tear the session down, capture first, and report how it ended.
342    fn end(&mut self, event: SpeechEvent, cx: &mut Context<Self>) {
343        if let Some(mut session) = self.session.take() {
344            session.capture = None;
345        }
346        self.status = SpeechStatus::Idle;
347        self.meter = LevelMeter::new();
348        cx.emit(event);
349        cx.notify();
350    }
351}
352
353/// Apply `f` to `state` after the current update, if `session` is still its
354/// running session. Sinks go through this so that a recognizer or an input may
355/// report from anywhere, including from inside a call the state made.
356pub(super) fn defer_session_update(
357    state: WeakEntity<SpeechState>,
358    session: usize,
359    cx: &mut App,
360    f: impl FnOnce(&mut SpeechState, &mut Context<SpeechState>) + 'static,
361) {
362    cx.defer(move |cx| {
363        _ = state.update(cx, |state, cx| {
364            if state.is_session(session) {
365                f(state, cx);
366            }
367        });
368    });
369}
370
371#[cfg(test)]
372mod tests {
373    use std::{cell::RefCell, rc::Rc, time::Duration};
374
375    use gpui::{App, AppContext as _, Entity, Subscription, TestAppContext};
376
377    use super::*;
378    use crate::speech::{AudioFormat, AudioInput, AudioSink, RecognitionSession, SpeechSink};
379
380    #[derive(Default)]
381    struct Recorded {
382        sink: Option<SpeechSink>,
383        samples: usize,
384        finished: bool,
385        dropped: bool,
386    }
387
388    #[derive(Clone, Default)]
389    struct FakeRecognizer {
390        recorded: Rc<RefCell<Recorded>>,
391        fail_to_start: bool,
392    }
393
394    struct FakeSession(Rc<RefCell<Recorded>>);
395
396    impl SpeechRecognizer for FakeRecognizer {
397        fn start(
398            &self,
399            sink: SpeechSink,
400            _: &mut App,
401        ) -> Result<Box<dyn RecognitionSession>, SpeechError> {
402            if self.fail_to_start {
403                return Err(SpeechError::recognizer(anyhow::anyhow!("offline")));
404            }
405            *self.recorded.borrow_mut() = Recorded {
406                sink: Some(sink),
407                ..Default::default()
408            };
409            Ok(Box::new(FakeSession(self.recorded.clone())))
410        }
411    }
412
413    impl RecognitionSession for FakeSession {
414        fn push_audio(&mut self, samples: &[i16], _: &mut App) {
415            self.0.borrow_mut().samples += samples.len();
416        }
417
418        fn finish(&mut self, _: &mut App) {
419            self.0.borrow_mut().finished = true;
420        }
421    }
422
423    impl Drop for FakeSession {
424        fn drop(&mut self) {
425            self.0.borrow_mut().dropped = true;
426        }
427    }
428
429    #[derive(Clone, Default)]
430    struct FakeInput {
431        sink: Rc<RefCell<Option<AudioSink>>>,
432        capturing: Rc<RefCell<bool>>,
433    }
434
435    impl AudioInput for FakeInput {
436        fn start(
437            &self,
438            _: AudioFormat,
439            sink: AudioSink,
440            _: &mut App,
441        ) -> Result<Subscription, SpeechError> {
442            *self.sink.borrow_mut() = Some(sink);
443            *self.capturing.borrow_mut() = true;
444            let capturing = self.capturing.clone();
445            Ok(Subscription::new(move || *capturing.borrow_mut() = false))
446        }
447    }
448
449    struct Fixture {
450        state: Entity<SpeechState>,
451        recognizer: FakeRecognizer,
452        input: FakeInput,
453        events: Rc<RefCell<Vec<SpeechEvent>>>,
454        _subscription: Subscription,
455    }
456
457    impl Fixture {
458        fn new(recognizer: FakeRecognizer, cx: &mut TestAppContext) -> Self {
459            let input = FakeInput::default();
460            let state = cx.update(|cx| {
461                cx.new(|cx| {
462                    SpeechState::new(cx)
463                        .recognizer(recognizer.clone())
464                        .input(input.clone())
465                        .stop_timeout(Duration::from_secs(1))
466                })
467            });
468            let events = Rc::new(RefCell::new(Vec::new()));
469            let _subscription = cx.update(|cx| {
470                let events = events.clone();
471                cx.subscribe(&state, move |_, event: &SpeechEvent, _| {
472                    events.borrow_mut().push(event.clone());
473                })
474            });
475            Self {
476                state,
477                recognizer,
478                input,
479                events,
480                _subscription,
481            }
482        }
483
484        fn sink(&self) -> SpeechSink {
485            self.recognizer.recorded.borrow().sink.clone().unwrap()
486        }
487
488        fn audio(&self) -> AudioSink {
489            self.input.sink.borrow().clone().unwrap()
490        }
491
492        fn status(&self, cx: &mut TestAppContext) -> SpeechStatus {
493            cx.read(|cx| self.state.read(cx).status())
494        }
495
496        /// Events so far, as short labels.
497        fn take_events(&self) -> Vec<String> {
498            self.events
499                .borrow_mut()
500                .drain(..)
501                .map(|event| match event {
502                    SpeechEvent::Started => "started".into(),
503                    SpeechEvent::Partial(text) => format!("partial:{text}"),
504                    SpeechEvent::Final(text) => format!("final:{text}"),
505                    SpeechEvent::Cancelled => "cancelled".into(),
506                    SpeechEvent::Error(error) => format!("error:{error}"),
507                })
508                .collect()
509        }
510    }
511
512    #[gpui::test]
513    fn session_runs_from_start_to_final(cx: &mut TestAppContext) {
514        let f = Fixture::new(FakeRecognizer::default(), cx);
515
516        f.state.update(cx, |state, cx| state.start(cx));
517        assert_eq!(f.status(cx), SpeechStatus::Connecting);
518        assert!(*f.input.capturing.borrow());
519
520        cx.update(|cx| {
521            f.sink().ready(cx);
522            // Two levels' worth: 80 ms is 1 280 samples at 16 kHz.
523            f.audio().push(vec![i16::MAX / 2; 2_560], cx);
524            f.sink().hypothesis("hello", cx);
525        });
526        cx.run_until_parked();
527        assert_eq!(f.status(cx), SpeechStatus::Recording);
528        assert_eq!(f.recognizer.recorded.borrow().samples, 2_560);
529        cx.read(|cx| {
530            let state = f.state.read(cx);
531            assert_eq!(state.levels().len(), 2);
532            assert!(state.levels().next().unwrap() > 0.5);
533        });
534
535        cx.update(|cx| {
536            f.sink().phrase("Hello.", cx);
537            f.sink().hypothesis(" How", cx);
538        });
539        cx.run_until_parked();
540        assert_eq!(
541            cx.read(|cx| f.state.read(cx).transcript()),
542            SharedString::from("Hello. How")
543        );
544
545        f.state.update(cx, |state, cx| state.stop(cx));
546        assert_eq!(f.status(cx), SpeechStatus::Stopping);
547        assert!(!*f.input.capturing.borrow(), "stop releases the microphone");
548        assert!(f.recognizer.recorded.borrow().finished);
549
550        cx.update(|cx| {
551            f.sink().phrase(" How are you?", cx);
552            f.sink().finish(cx);
553        });
554        cx.run_until_parked();
555        assert_eq!(f.status(cx), SpeechStatus::Idle);
556        assert!(f.recognizer.recorded.borrow().dropped);
557        assert_eq!(
558            f.take_events(),
559            [
560                "started",
561                "partial:hello",
562                "partial:Hello.",
563                "partial:Hello. How",
564                "partial:Hello. How are you?",
565                "final:Hello. How are you?",
566            ]
567        );
568    }
569
570    #[gpui::test]
571    fn stop_ends_with_the_transcript_so_far_after_the_timeout(cx: &mut TestAppContext) {
572        let f = Fixture::new(FakeRecognizer::default(), cx);
573        f.state.update(cx, |state, cx| state.start(cx));
574        cx.update(|cx| f.sink().hypothesis("half a sentence", cx));
575        cx.run_until_parked();
576
577        f.state.update(cx, |state, cx| state.stop(cx));
578        cx.executor().advance_clock(Duration::from_millis(900));
579        assert_eq!(f.status(cx), SpeechStatus::Stopping);
580        cx.executor().advance_clock(Duration::from_millis(200));
581        cx.run_until_parked();
582
583        assert_eq!(f.status(cx), SpeechStatus::Idle);
584        assert_eq!(f.take_events().last().unwrap(), "final:half a sentence");
585    }
586
587    #[gpui::test]
588    fn cancel_discards_the_session_and_ignores_late_results(cx: &mut TestAppContext) {
589        let f = Fixture::new(FakeRecognizer::default(), cx);
590        f.state.update(cx, |state, cx| state.start(cx));
591        let stale = f.sink();
592        cx.update(|cx| stale.phrase("draft", cx));
593        cx.run_until_parked();
594
595        f.state.update(cx, |state, cx| state.cancel(cx));
596        assert_eq!(f.status(cx), SpeechStatus::Idle);
597        assert!(!*f.input.capturing.borrow());
598        assert!(f.recognizer.recorded.borrow().dropped);
599
600        // A new session must not pick up the old session's results.
601        f.state.update(cx, |state, cx| state.start(cx));
602        cx.update(|cx| {
603            stale.phrase("late", cx);
604            stale.finish(cx);
605        });
606        cx.run_until_parked();
607
608        assert_eq!(f.status(cx), SpeechStatus::Connecting);
609        assert_eq!(
610            cx.read(|cx| f.state.read(cx).transcript()),
611            SharedString::default()
612        );
613        assert_eq!(
614            f.take_events(),
615            ["started", "partial:draft", "cancelled", "started"]
616        );
617    }
618
619    #[gpui::test]
620    fn input_error_ends_the_session(cx: &mut TestAppContext) {
621        let f = Fixture::new(FakeRecognizer::default(), cx);
622        f.state.update(cx, |state, cx| state.start(cx));
623        cx.update(|cx| f.audio().error(SpeechError::NoInputDevice, cx));
624        cx.run_until_parked();
625
626        assert_eq!(f.status(cx), SpeechStatus::Idle);
627        assert!(f.recognizer.recorded.borrow().dropped);
628        assert_eq!(
629            f.take_events(),
630            ["started", "error:no audio input device is available"]
631        );
632    }
633
634    #[gpui::test]
635    fn recognizer_that_fails_to_start_leaves_the_state_idle(cx: &mut TestAppContext) {
636        let f = Fixture::new(
637            FakeRecognizer {
638                fail_to_start: true,
639                ..Default::default()
640            },
641            cx,
642        );
643        f.state.update(cx, |state, cx| state.start(cx));
644
645        assert_eq!(f.status(cx), SpeechStatus::Idle);
646        assert!(!*f.input.capturing.borrow(), "the input never starts");
647        assert_eq!(
648            f.take_events(),
649            ["error:speech recognition failed: offline"]
650        );
651    }
652
653    #[gpui::test]
654    fn state_without_a_recognizer_is_unsupported(cx: &mut TestAppContext) {
655        let state = cx.update(|cx| {
656            cx.new(|cx| {
657                SpeechState::new(cx)
658                    .input(FakeInput::default())
659                    .system_fallback(false)
660            })
661        });
662        cx.read(|cx| {
663            assert!(!state.read(cx).has_recognizer());
664            assert!(!state.read(cx).is_available(cx));
665        });
666    }
667}