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