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
63pub 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 auto_hangup: Option<bool>,
264 ) -> Result<(SynthesisHandle, Box<dyn Track>)> {
265 let (tx, rx) = mpsc::unbounded_channel();
266 let new_handle = SynthesisHandle::new(tx, play_id.clone(), ssrc);
267 let tts_client = engine.create_tts_client(streaming, tts_option).await?;
268 let sample_rate = tts_option.samplerate.unwrap_or(16000) as u32;
269 let tts_track = TtsTrack::new(track_id, session_id, streaming, play_id, rx, tts_client)
270 .with_ssrc(ssrc)
271 .with_auto_hangup(auto_hangup)
272 .with_sample_rate(sample_rate)
273 .with_cancel_token(cancel_token);
274 Ok((new_handle, Box::new(tts_track) as Box<dyn Track>))
275 }
276
277 pub fn with_processor_hook(&mut self, hook_fn: CreateProcessorsHook) -> &mut Self {
278 self.create_processors_hook = Arc::new(Box::new(hook_fn));
279 self
280 }
281
282 pub fn default_create_procesors_hook(
283 engine: Arc<StreamEngine>,
284 track_id: TrackId,
285 cancel_token: CancellationToken,
286 event_sender: EventSender,
287 packet_sender: TrackPacketSender,
288 option: CallOption,
289 ) -> Pin<Box<dyn Future<Output = Result<Vec<Box<dyn Processor>>>> + Send>> {
290 Box::pin(async move {
291 let mut processors = vec![];
292 debug!(%track_id, "Creating processors for track");
293
294 if let Some(realtime_option) = option.realtime {
295 debug!(%track_id, "Adding RealtimeProcessor");
296 let realtime_processor = crate::media::realtime_processor::RealtimeProcessor::new(
297 track_id.clone(),
298 cancel_token.child_token(),
299 event_sender.clone(),
300 packet_sender.clone(),
301 realtime_option,
302 )?;
303 processors.push(Box::new(realtime_processor) as Box<dyn Processor>);
304 return Ok(processors);
307 }
308
309 match option.denoise {
310 Some(true) => {
311 debug!(%track_id, "Adding NoiseReducer processor");
312 let noise_reducer = NoiseReducer::new(INTERNAL_SAMPLERATE as usize);
313 processors.push(Box::new(noise_reducer) as Box<dyn Processor>);
314 }
315 _ => {}
316 }
317 match option.vad {
318 Some(mut option) => {
319 debug!(%track_id, "Adding VadProcessor processor type={:?}", option.r#type);
320 option.samplerate = INTERNAL_SAMPLERATE;
321 let vad_processor: Box<dyn Processor + 'static> = engine.create_vad_processor(
322 cancel_token.child_token(),
323 event_sender.clone(),
324 option.to_owned(),
325 )?;
326 processors.push(vad_processor);
327 }
328 None => {}
329 }
330 if let Some(agc_option) = option.agc.clone() {
331 debug!(%track_id, "Adding AutomaticGainControl processor");
332 let agc = AutomaticGainControl::new(INTERNAL_SAMPLERATE as u32, agc_option)?;
333 processors.push(Box::new(agc) as Box<dyn Processor>);
334 }
335 match option.asr {
336 Some(mut option) => {
337 debug!(%track_id, "Adding AsrProcessor processor provider={:?}", option.provider);
338 option.samplerate = Some(INTERNAL_SAMPLERATE);
339 let asr_processor = engine
340 .create_asr_processor(
341 track_id.clone(),
342 cancel_token.child_token(),
343 option.to_owned(),
344 event_sender.clone(),
345 )
346 .await?;
347 processors.push(asr_processor);
348 }
349 None => {}
350 }
351 match option.eou {
352 Some(ref option) => {
353 let eou_processor = engine.create_eou_processor(
354 cancel_token.child_token(),
355 event_sender.clone(),
356 option.to_owned(),
357 )?;
358 processors.push(eou_processor);
359 }
360 None => {}
361 }
362 match option.inactivity_timeout {
363 Some(timeout_secs) if timeout_secs > 0 => {
364 let inactivity_processor = crate::media::inactivity::InactivityProcessor::new(
365 track_id.clone(),
366 std::time::Duration::from_secs(timeout_secs),
367 event_sender.clone(),
368 cancel_token.child_token(),
369 );
370 processors.push(Box::new(inactivity_processor) as Box<dyn Processor>);
371 }
372 _ => {}
373 }
374
375 #[cfg(feature = "ringback-detection")]
376 if let Some(ref ringback_opt) = option.ringback_detection {
377 if ringback_opt.enabled.unwrap_or(false) {
378 debug!(%track_id, "Adding RingbackDetectionProcessor");
379 match
380 crate::media::ringback_detection::processor::RingbackDetectionProcessor::new(
381 track_id.clone(),
382 cancel_token.child_token(),
383 event_sender.clone(),
384 ringback_opt.clone(),
385 None,
386 ) {
387 Ok(p) => processors.push(Box::new(p) as Box<dyn Processor>),
388 Err(e) => warn!(%track_id, "Failed to create RingbackDetectionProcessor: {}", e),
389 }
390 }
391 }
392
393 Ok(processors)
394 })
395 }
396}