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