use std::{cell::RefCell, rc::Rc};
use anyhow::anyhow;
use block2::RcBlock;
use gpui::{App, SharedString, Task};
use objc2::{AnyThread as _, rc::Retained, runtime::NSObjectProtocol as _, sel};
use objc2_av_foundation::{AVAuthorizationStatus, AVCaptureDevice, AVMediaTypeAudio};
use objc2_avf_audio::{AVAudioCommonFormat, AVAudioFormat, AVAudioPCMBuffer};
use objc2_foundation::{NSBundle, NSError, NSLocale, NSString};
use objc2_speech::{
SFSpeechAudioBufferRecognitionRequest, SFSpeechRecognitionResult, SFSpeechRecognitionTask,
SFSpeechRecognizer, SFSpeechRecognizerAuthorizationStatus,
};
use smol::channel::{Receiver, Sender, unbounded};
use crate::speech::{AudioFormat, RecognitionSession, SpeechError, SpeechRecognizer, SpeechSink};
const USAGE_DESCRIPTION_KEY: &str = "NSSpeechRecognitionUsageDescription";
const NO_SPEECH_DOMAIN: &str = "kAFAssistantErrorDomain";
const NO_SPEECH_CODE: isize = 1110;
pub(super) struct PlatformRecognizer {
recognizer: Option<Retained<SFSpeechRecognizer>>,
}
impl PlatformRecognizer {
pub(super) fn new(locale: Option<SharedString>) -> Self {
if !has_usage_description() {
tracing::warn!(
"speech recognition is unavailable: the application's Info.plist has no \
`{USAGE_DESCRIPTION_KEY}`, and asking for access without it terminates the process"
);
return Self { recognizer: None };
}
let recognizer = match &locale {
Some(locale) => {
let locale = NSLocale::initWithLocaleIdentifier(
NSLocale::alloc(),
&NSString::from_str(locale),
);
unsafe { SFSpeechRecognizer::initWithLocale(SFSpeechRecognizer::alloc(), &locale) }
}
None => unsafe { SFSpeechRecognizer::init(SFSpeechRecognizer::alloc()) },
};
if recognizer.is_none() {
tracing::warn!(
"speech recognition is unavailable: no recognizer for locale {}",
locale.as_deref().unwrap_or("of the system")
);
}
Self { recognizer }
}
fn on_device(&self) -> Option<&Retained<SFSpeechRecognizer>> {
self.recognizer
.as_ref()
.filter(|recognizer| unsafe {
recognizer.isAvailable() && recognizer.supportsOnDeviceRecognition()
})
}
}
impl SpeechRecognizer for PlatformRecognizer {
fn audio_format(&self) -> AudioFormat {
AudioFormat::default()
}
fn is_available(&self, _: &App) -> bool {
self.on_device().is_some() && !is_speech_denied(speech_authorization())
}
fn start(
&self,
sink: SpeechSink,
cx: &mut App,
) -> Result<Box<dyn RecognitionSession>, SpeechError> {
let recognizer = self.on_device().ok_or(SpeechError::Unsupported)?.clone();
let authorization = speech_authorization();
if is_speech_denied(authorization) {
return Err(SpeechError::PermissionDenied);
}
let (events, rx) = unbounded();
let mut recognition = Recognition::new(recognizer, sink, events.clone())?;
if authorization == SFSpeechRecognizerAuthorizationStatus::Authorized {
recognition.begin(cx);
} else {
request_authorization(events);
}
let recognition = Rc::new(RefCell::new(recognition));
let task = cx.spawn({
let recognition = recognition.clone();
async move |cx| run(recognition, rx, cx).await
});
Ok(Box::new(Session {
recognition,
_task: task,
}))
}
}
pub(in crate::speech) fn is_microphone_denied() -> bool {
let Some(audio) = (unsafe { AVMediaTypeAudio }) else {
return false;
};
let status = unsafe { AVCaptureDevice::authorizationStatusForMediaType(audio) };
status == AVAuthorizationStatus::Denied || status == AVAuthorizationStatus::Restricted
}
fn has_usage_description() -> bool {
NSBundle::mainBundle()
.objectForInfoDictionaryKey(&NSString::from_str(USAGE_DESCRIPTION_KEY))
.is_some()
}
fn speech_authorization() -> SFSpeechRecognizerAuthorizationStatus {
unsafe { SFSpeechRecognizer::authorizationStatus() }
}
fn is_speech_denied(status: SFSpeechRecognizerAuthorizationStatus) -> bool {
status == SFSpeechRecognizerAuthorizationStatus::Denied
|| status == SFSpeechRecognizerAuthorizationStatus::Restricted
}
fn request_authorization(events: Sender<Event>) {
let handler = RcBlock::new(move |status: SFSpeechRecognizerAuthorizationStatus| {
_ = events.try_send(Event::Authorized(
status == SFSpeechRecognizerAuthorizationStatus::Authorized,
));
});
unsafe { SFSpeechRecognizer::requestAuthorization(&handler) };
}
enum Event {
Authorized(bool),
Result {
text: String,
is_final: bool,
ends_utterance: bool,
},
Error {
domain: String,
code: isize,
message: String,
},
}
async fn run(
recognition: Rc<RefCell<Recognition>>,
events: Receiver<Event>,
cx: &mut gpui::AsyncApp,
) {
while let Ok(event) = events.recv().await {
let done = cx.update(|cx| recognition.borrow_mut().on_event(event, cx));
if done {
break;
}
}
}
struct Session {
recognition: Rc<RefCell<Recognition>>,
_task: Task<()>,
}
impl RecognitionSession for Session {
fn push_audio(&mut self, samples: &[i16], _: &mut App) {
self.recognition.borrow_mut().push_audio(samples);
}
fn finish(&mut self, _: &mut App) {
self.recognition.borrow_mut().finish();
}
}
impl Drop for Session {
fn drop(&mut self) {
if let Some(task) = &self.recognition.borrow().task {
unsafe { task.cancel() };
}
}
}
struct Recognition {
recognizer: Retained<SFSpeechRecognizer>,
request: Retained<SFSpeechAudioBufferRecognitionRequest>,
format: Retained<AVAudioFormat>,
task: Option<Retained<SFSpeechRecognitionTask>>,
sink: SpeechSink,
events: Sender<Event>,
pending: Vec<i16>,
finishing: bool,
done: bool,
utterance: Option<String>,
hypothesis: String,
last_committed: Option<char>,
}
impl Recognition {
fn new(
recognizer: Retained<SFSpeechRecognizer>,
sink: SpeechSink,
events: Sender<Event>,
) -> Result<Self, SpeechError> {
let format = AudioFormat::default();
let format = unsafe {
AVAudioFormat::initWithCommonFormat_sampleRate_channels_interleaved(
AVAudioFormat::alloc(),
AVAudioCommonFormat::PCMFormatInt16,
format.sample_rate().into(),
format.channels().into(),
false,
)
}
.ok_or_else(|| SpeechError::recognizer(anyhow!("cannot create the audio format")))?;
let request = unsafe {
let request = SFSpeechAudioBufferRecognitionRequest::new();
request.setShouldReportPartialResults(true);
request.setRequiresOnDeviceRecognition(true);
if request.respondsToSelector(sel!(setAddsPunctuation:)) {
request.setAddsPunctuation(true);
}
request
};
Ok(Self {
recognizer,
request,
format,
task: None,
sink,
events,
pending: Vec::new(),
finishing: false,
done: false,
utterance: None,
hypothesis: String::new(),
last_committed: None,
})
}
fn begin(&mut self, cx: &mut App) {
let events = self.events.clone();
let handler = RcBlock::new(
move |result: *mut SFSpeechRecognitionResult, error: *mut NSError| {
let (result, error) = unsafe { (result.as_ref(), error.as_ref()) };
if let Some(result) = result {
let event = unsafe {
Event::Result {
text: result.bestTranscription().formattedString().to_string(),
is_final: result.isFinal(),
ends_utterance: result.speechRecognitionMetadata().is_some(),
}
};
_ = events.try_send(event);
}
if let Some(error) = error {
_ = events.try_send(Event::Error {
domain: error.domain().to_string(),
code: error.code(),
message: error.localizedDescription().to_string(),
});
}
},
);
let task = unsafe {
self.recognizer
.recognitionTaskWithRequest_resultHandler(&self.request, &handler)
};
self.task = Some(task);
let pending = std::mem::take(&mut self.pending);
self.append(&pending);
if self.finishing {
unsafe { self.request.endAudio() };
}
self.sink.ready(cx);
}
fn push_audio(&mut self, samples: &[i16]) {
if self.done || self.finishing {
return;
}
if self.task.is_some() {
self.append(samples);
} else {
self.pending.extend_from_slice(samples);
}
}
fn finish(&mut self) {
if self.finishing {
return;
}
self.finishing = true;
if self.task.is_some() {
unsafe { self.request.endAudio() };
}
}
fn append(&self, samples: &[i16]) {
let Ok(frames) = u32::try_from(samples.len()) else {
return;
};
if frames == 0 {
return;
}
let Some(buffer) = (unsafe {
AVAudioPCMBuffer::initWithPCMFormat_frameCapacity(
AVAudioPCMBuffer::alloc(),
&self.format,
frames,
)
}) else {
return;
};
unsafe {
let channel = (*buffer.int16ChannelData()).as_ptr();
std::ptr::copy_nonoverlapping(samples.as_ptr(), channel, samples.len());
buffer.setFrameLength(frames);
self.request.appendAudioPCMBuffer(&buffer);
}
}
fn on_event(&mut self, event: Event, cx: &mut App) -> bool {
if self.done {
return true;
}
match event {
Event::Authorized(true) => self.begin(cx),
Event::Authorized(false) => {
self.sink.error(SpeechError::PermissionDenied, cx);
self.done = true;
}
Event::Result {
text,
is_final,
ends_utterance,
} => self.on_result(text, is_final, ends_utterance, cx),
Event::Error {
domain,
code,
message,
} => {
if self.finishing && domain == NO_SPEECH_DOMAIN && code == NO_SPEECH_CODE {
let hypothesis = std::mem::take(&mut self.hypothesis);
self.commit(&hypothesis, cx);
self.sink.finish(cx);
} else {
self.sink.error(
SpeechError::recognizer(anyhow!("{message} ({domain} {code})")),
cx,
);
}
self.done = true;
}
}
self.done
}
fn on_result(&mut self, text: String, is_final: bool, ends_utterance: bool, cx: &mut App) {
if let Some(utterance) = self.utterance.take()
&& starts_over(&utterance, &text)
{
self.commit(&utterance, cx);
}
if is_final {
self.hypothesis.clear();
self.commit(&text, cx);
self.sink.finish(cx);
self.done = true;
return;
}
if ends_utterance {
self.utterance = Some(text.clone());
}
self.sink.hypothesis(self.joined(&text), cx);
self.hypothesis = text;
}
fn commit(&mut self, text: &str, cx: &mut App) {
if text.is_empty() {
return;
}
let text = self.joined(text);
self.last_committed = text.chars().next_back();
self.sink.phrase(text, cx);
}
fn joined(&self, text: &str) -> String {
match (self.last_committed, text.chars().next()) {
(Some(before), Some(after)) if needs_space(before, after) => format!(" {text}"),
_ => text.to_string(),
}
}
}
fn starts_over(utterance: &str, text: &str) -> bool {
let common = utterance
.chars()
.zip(text.chars())
.take_while(|(a, b)| a == b)
.count();
common * 2 < utterance.chars().count()
}
fn needs_space(before: char, after: char) -> bool {
!before.is_whitespace()
&& !after.is_whitespace()
&& !matches!(after, ',' | '.' | '?' | '!' | ';' | ':' | ')')
&& !is_unspaced_script(before)
&& !is_unspaced_script(after)
}
fn is_unspaced_script(c: char) -> bool {
matches!(c,
'\u{3000}'..='\u{30FF}'
| '\u{3400}'..='\u{4DBF}'
| '\u{4E00}'..='\u{9FFF}'
| '\u{F900}'..='\u{FAFF}'
| '\u{FF00}'..='\u{FFEF}'
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_revised_utterance_carries_on() {
assert!(!starts_over("Hello word", "Hello world, how"));
assert!(!starts_over("Hello world", "Hello world. How are you"));
assert!(!starts_over("今天天气", "今天天气很好"));
}
#[test]
fn a_reset_result_starts_over() {
assert!(starts_over("Hello world.", "How"));
assert!(starts_over("今天天气很好。", "明天"));
}
#[test]
fn phrases_are_spaced_by_script() {
assert!(needs_space('.', 'H'));
assert!(!needs_space('d', ','));
assert!(!needs_space('。', '明'));
assert!(!needs_space('好', 'O'));
}
}