Skip to main content

active_call/media/
engine.rs

1use super::{
2    INTERNAL_SAMPLERATE,
3    asr_processor::AsrProcessor,
4    denoiser::NoiseReducer,
5    processor::Processor,
6    track::{
7        Track, TrackPacketSender,
8        tts::{SynthesisHandle, TtsTrack},
9    },
10    vad::{VADOption, VadProcessor, VadType},
11};
12use crate::{
13    CallOption, EouOption,
14    event::EventSender,
15    media::TrackId,
16    synthesis::{
17        AliyunTtsClient, DeepegramTtsClient, SynthesisClient, SynthesisOption, SynthesisType,
18        TencentCloudTtsBasicClient, TencentCloudTtsClient,
19    },
20    transcription::{
21        AliyunAsrClientBuilder, DeepgramAsrClientBuilder, TencentCloudAsrClientBuilder,
22        TranscriptionClient, TranscriptionOption, TranscriptionType,
23    },
24};
25
26#[cfg(feature = "offline")]
27use crate::{synthesis::SupertonicTtsClient, transcription::SensevoiceAsrClientBuilder};
28
29use anyhow::Result;
30use std::{collections::HashMap, future::Future, pin::Pin, sync::Arc};
31use tokio::sync::mpsc;
32use tokio_util::sync::CancellationToken;
33use tracing::debug;
34#[cfg(feature = "ringback-detection")]
35use tracing::warn;
36
37pub type FnCreateVadProcessor = fn(
38    token: CancellationToken,
39    event_sender: EventSender,
40    option: VADOption,
41) -> Result<Box<dyn Processor>>;
42
43pub type FnCreateEouProcessor = fn(
44    token: CancellationToken,
45    event_sender: EventSender,
46    option: EouOption,
47) -> Result<Box<dyn Processor>>;
48
49pub type FnCreateAsrClient = Box<
50    dyn Fn(
51            TrackId,
52            CancellationToken,
53            TranscriptionOption,
54            EventSender,
55        ) -> Pin<Box<dyn Future<Output = Result<Box<dyn TranscriptionClient>>> + Send>>
56        + Send
57        + Sync,
58>;
59pub type FnCreateTtsClient =
60    fn(streaming: bool, option: &SynthesisOption) -> Result<Box<dyn SynthesisClient>>;
61
62// Define hook types
63pub type CreateProcessorsHook = Box<
64    dyn Fn(
65            Arc<StreamEngine>,
66            TrackId,
67            CancellationToken,
68            EventSender,
69            TrackPacketSender,
70            CallOption,
71        ) -> Pin<Box<dyn Future<Output = Result<Vec<Box<dyn Processor>>>> + Send>>
72        + Send
73        + Sync,
74>;
75
76pub struct StreamEngine {
77    vad_creators: HashMap<VadType, FnCreateVadProcessor>,
78    eou_creators: HashMap<String, FnCreateEouProcessor>,
79    asr_creators: HashMap<TranscriptionType, FnCreateAsrClient>,
80    tts_creators: HashMap<SynthesisType, FnCreateTtsClient>,
81    create_processors_hook: Arc<CreateProcessorsHook>,
82}
83
84impl Default for StreamEngine {
85    fn default() -> Self {
86        let mut engine = Self::new();
87        engine.register_vad(VadType::Silero, VadProcessor::create);
88        engine.register_vad(VadType::Other("nop".to_string()), VadProcessor::create_nop);
89
90        engine.register_asr(
91            TranscriptionType::TencentCloud,
92            Box::new(TencentCloudAsrClientBuilder::create),
93        );
94        engine.register_asr(
95            TranscriptionType::Aliyun,
96            Box::new(AliyunAsrClientBuilder::create),
97        );
98        engine.register_asr(
99            TranscriptionType::Deepgram,
100            Box::new(DeepgramAsrClientBuilder::create),
101        );
102
103        #[cfg(feature = "offline")]
104        engine.register_asr(
105            TranscriptionType::Sensevoice,
106            Box::new(SensevoiceAsrClientBuilder::create),
107        );
108
109        engine.register_tts(SynthesisType::Aliyun, AliyunTtsClient::create);
110        engine.register_tts(SynthesisType::TencentCloud, TencentCloudTtsClient::create);
111        engine.register_tts(
112            SynthesisType::Other("tencent_basic".to_string()),
113            TencentCloudTtsBasicClient::create,
114        );
115        engine.register_tts(SynthesisType::Deepgram, DeepegramTtsClient::create);
116
117        #[cfg(feature = "offline")]
118        engine.register_tts(SynthesisType::Supertonic, SupertonicTtsClient::create);
119
120        engine
121    }
122}
123
124impl StreamEngine {
125    pub fn new() -> Self {
126        Self {
127            vad_creators: HashMap::new(),
128            asr_creators: HashMap::new(),
129            tts_creators: HashMap::new(),
130            eou_creators: HashMap::new(),
131            create_processors_hook: Arc::new(Box::new(Self::default_create_procesors_hook)),
132        }
133    }
134
135    pub fn register_vad(&mut self, vad_type: VadType, creator: FnCreateVadProcessor) -> &mut Self {
136        self.vad_creators.insert(vad_type, creator);
137        self
138    }
139
140    pub fn register_eou(&mut self, name: String, creator: FnCreateEouProcessor) -> &mut Self {
141        self.eou_creators.insert(name, creator);
142        self
143    }
144
145    pub fn register_asr(
146        &mut self,
147        asr_type: TranscriptionType,
148        creator: FnCreateAsrClient,
149    ) -> &mut Self {
150        self.asr_creators.insert(asr_type, creator);
151        self
152    }
153
154    pub fn register_tts(
155        &mut self,
156        tts_type: SynthesisType,
157        creator: FnCreateTtsClient,
158    ) -> &mut Self {
159        self.tts_creators.insert(tts_type, creator);
160        self
161    }
162
163    pub fn create_vad_processor(
164        &self,
165        token: CancellationToken,
166        event_sender: EventSender,
167        option: VADOption,
168    ) -> Result<Box<dyn Processor>> {
169        let creator = self.vad_creators.get(&option.r#type);
170        if let Some(creator) = creator {
171            creator(token, event_sender, option)
172        } else {
173            Err(anyhow::anyhow!("VAD type not found: {}", option.r#type))
174        }
175    }
176    pub fn create_eou_processor(
177        &self,
178        token: CancellationToken,
179        event_sender: EventSender,
180        option: EouOption,
181    ) -> Result<Box<dyn Processor>> {
182        let creator = self
183            .eou_creators
184            .get(&option.r#type.clone().unwrap_or_default());
185        if let Some(creator) = creator {
186            creator(token, event_sender, option)
187        } else {
188            Err(anyhow::anyhow!("EOU type not found: {:?}", option.r#type))
189        }
190    }
191
192    pub async fn create_asr_processor(
193        &self,
194        track_id: TrackId,
195        cancel_token: CancellationToken,
196        option: TranscriptionOption,
197        event_sender: EventSender,
198    ) -> Result<Box<dyn Processor>> {
199        let asr_client = match option.provider {
200            Some(ref provider) => {
201                let creator = self.asr_creators.get(&provider);
202                if let Some(creator) = creator {
203                    creator(track_id, cancel_token, option, event_sender).await?
204                } else {
205                    return Err(anyhow::anyhow!("ASR type not found: {}", provider));
206                }
207            }
208            None => return Err(anyhow::anyhow!("ASR type not found: {:?}", option.provider)),
209        };
210        Ok(Box::new(AsrProcessor { asr_client }))
211    }
212
213    pub async fn create_tts_client(
214        &self,
215        streaming: bool,
216        tts_option: &SynthesisOption,
217    ) -> Result<Box<dyn SynthesisClient>> {
218        match tts_option.provider {
219            Some(ref provider) => {
220                let creator = self.tts_creators.get(&provider);
221                if let Some(creator) = creator {
222                    creator(streaming, tts_option)
223                } else {
224                    Err(anyhow::anyhow!("TTS type not found: {}", provider))
225                }
226            }
227            None => Err(anyhow::anyhow!(
228                "TTS type not found: {:?}",
229                tts_option.provider
230            )),
231        }
232    }
233
234    pub async fn create_processors(
235        engine: Arc<StreamEngine>,
236        track: &dyn Track,
237        cancel_token: CancellationToken,
238        event_sender: EventSender,
239        packet_sender: TrackPacketSender,
240        option: &CallOption,
241    ) -> Result<Vec<Box<dyn Processor>>> {
242        (engine.clone().create_processors_hook)(
243            engine,
244            track.id().clone(),
245            cancel_token,
246            event_sender,
247            packet_sender,
248            option.clone(),
249        )
250        .await
251    }
252
253    pub async fn create_tts_track(
254        engine: Arc<StreamEngine>,
255        cancel_token: CancellationToken,
256        session_id: String,
257        track_id: TrackId,
258        ssrc: u32,
259        play_id: Option<String>,
260        streaming: bool,
261        tts_option: &SynthesisOption,
262    ) -> Result<(SynthesisHandle, Box<dyn Track>)> {
263        let (tx, rx) = mpsc::unbounded_channel();
264        let new_handle = SynthesisHandle::new(tx, play_id.clone(), ssrc);
265        let tts_client = engine.create_tts_client(streaming, tts_option).await?;
266        let sample_rate = tts_option.samplerate.unwrap_or(16000) as u32;
267        let tts_track = TtsTrack::new(track_id, session_id, streaming, play_id, rx, tts_client)
268            .with_ssrc(ssrc)
269            .with_sample_rate(sample_rate)
270            .with_cancel_token(cancel_token);
271        Ok((new_handle, Box::new(tts_track) as Box<dyn Track>))
272    }
273
274    pub fn with_processor_hook(&mut self, hook_fn: CreateProcessorsHook) -> &mut Self {
275        self.create_processors_hook = Arc::new(Box::new(hook_fn));
276        self
277    }
278
279    pub fn default_create_procesors_hook(
280        engine: Arc<StreamEngine>,
281        track_id: TrackId,
282        cancel_token: CancellationToken,
283        event_sender: EventSender,
284        packet_sender: TrackPacketSender,
285        option: CallOption,
286    ) -> Pin<Box<dyn Future<Output = Result<Vec<Box<dyn Processor>>>> + Send>> {
287        Box::pin(async move {
288            let mut processors = vec![];
289            debug!(%track_id, "Creating processors for track");
290
291            if let Some(realtime_option) = option.realtime {
292                debug!(%track_id, "Adding RealtimeProcessor");
293                let realtime_processor = crate::media::realtime_processor::RealtimeProcessor::new(
294                    track_id.clone(),
295                    cancel_token.child_token(),
296                    event_sender.clone(),
297                    packet_sender.clone(),
298                    realtime_option,
299                )?;
300                processors.push(Box::new(realtime_processor) as Box<dyn Processor>);
301                // In realtime mode, we usually don't need separate VAD or ASR processors
302                // as they are handled by the realtime service (OpenAI/Azure)
303                return Ok(processors);
304            }
305
306            match option.denoise {
307                Some(true) => {
308                    debug!(%track_id, "Adding NoiseReducer processor");
309                    let noise_reducer = NoiseReducer::new(INTERNAL_SAMPLERATE as usize);
310                    processors.push(Box::new(noise_reducer) as Box<dyn Processor>);
311                }
312                _ => {}
313            }
314            match option.vad {
315                Some(mut option) => {
316                    debug!(%track_id, "Adding VadProcessor processor type={:?}", option.r#type);
317                    option.samplerate = INTERNAL_SAMPLERATE;
318                    let vad_processor: Box<dyn Processor + 'static> = engine.create_vad_processor(
319                        cancel_token.child_token(),
320                        event_sender.clone(),
321                        option.to_owned(),
322                    )?;
323                    processors.push(vad_processor);
324                }
325                None => {}
326            }
327            match option.asr {
328                Some(mut option) => {
329                    debug!(%track_id, "Adding AsrProcessor processor provider={:?}", option.provider);
330                    option.samplerate = Some(INTERNAL_SAMPLERATE);
331                    let asr_processor = engine
332                        .create_asr_processor(
333                            track_id.clone(),
334                            cancel_token.child_token(),
335                            option.to_owned(),
336                            event_sender.clone(),
337                        )
338                        .await?;
339                    processors.push(asr_processor);
340                }
341                None => {}
342            }
343            match option.eou {
344                Some(ref option) => {
345                    let eou_processor = engine.create_eou_processor(
346                        cancel_token.child_token(),
347                        event_sender.clone(),
348                        option.to_owned(),
349                    )?;
350                    processors.push(eou_processor);
351                }
352                None => {}
353            }
354            match option.inactivity_timeout {
355                Some(timeout_secs) if timeout_secs > 0 => {
356                    let inactivity_processor = crate::media::inactivity::InactivityProcessor::new(
357                        track_id.clone(),
358                        std::time::Duration::from_secs(timeout_secs),
359                        event_sender.clone(),
360                        cancel_token.child_token(),
361                    );
362                    processors.push(Box::new(inactivity_processor) as Box<dyn Processor>);
363                }
364                _ => {}
365            }
366
367            #[cfg(feature = "ringback-detection")]
368            if let Some(ref ringback_opt) = option.ringback_detection {
369                if ringback_opt.enabled.unwrap_or(false) {
370                    debug!(%track_id, "Adding RingbackDetectionProcessor");
371                    match
372                        crate::media::ringback_detection::processor::RingbackDetectionProcessor::new(
373                            track_id.clone(),
374                            cancel_token.child_token(),
375                            event_sender.clone(),
376                            ringback_opt.clone(),
377                            None,
378                        ) {
379                        Ok(p) => processors.push(Box::new(p) as Box<dyn Processor>),
380                        Err(e) => warn!(%track_id, "Failed to create RingbackDetectionProcessor: {}", e),
381                    }
382                }
383            }
384
385            Ok(processors)
386        })
387    }
388}