Skip to main content

active_call/media/
engine.rs

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