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
10const DEFAULT_STOP_TIMEOUT: Duration = Duration::from_secs(3);
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
15pub enum SpeechStatus {
16 #[default]
18 Idle,
19 Connecting,
21 Recording,
23 Stopping,
25}
26
27impl SpeechStatus {
28 pub fn is_active(self) -> bool {
30 self != Self::Idle
31 }
32
33 pub fn is_capturing(self) -> bool {
35 matches!(self, Self::Connecting | Self::Recording)
36 }
37}
38
39#[derive(Debug, Clone)]
41pub enum SpeechEvent {
42 Started,
44 Partial(SharedString),
47 Final(SharedString),
49 Cancelled,
51 Error(SpeechError),
53}
54
55struct Session {
56 id: usize,
57 capture: Option<Subscription>,
59 recognition: Box<dyn RecognitionSession>,
60 _stop_timeout: Option<Task<()>>,
61}
62
63pub 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 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 pub fn recognizer(mut self, recognizer: impl SpeechRecognizer) -> Self {
111 self.recognizer = Some(Rc::new(recognizer));
112 self
113 }
114
115 pub fn input(mut self, input: impl AudioInput) -> Self {
117 self.input = Some(Rc::new(input));
118 self
119 }
120
121 pub fn system_fallback(mut self, system_fallback: bool) -> Self {
127 self.system_fallback = system_fallback;
128 self
129 }
130
131 pub fn stop_timeout(mut self, timeout: Duration) -> Self {
134 self.stop_timeout = timeout;
135 self
136 }
137
138 pub fn has_recognizer(&self) -> bool {
141 self.input.is_some() && self.active_recognizer().is_some()
142 }
143
144 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 pub fn status(&self) -> SpeechStatus {
155 self.status
156 }
157
158 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 pub fn levels(&self) -> impl ExactSizeIterator<Item = f32> + '_ {
172 self.meter.levels()
173 }
174
175 pub(super) fn level_lead_at(&self, now: instant::Instant) -> Option<f32> {
178 self.meter.lead_at(now)
179 }
180
181 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 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 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 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 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 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
353pub(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 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 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 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}