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 ) -> 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 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}