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
62pub 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 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}