1use std::collections::HashMap;
22use std::sync::{Arc, Mutex as StdMutex};
23use std::time::Duration;
24
25use sipx_audio::mix::mix_into;
26use tokio::sync::Mutex;
27use tokio::task::JoinHandle;
28
29use crate::session::{MediaSession, Stop};
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
33#[non_exhaustive]
34pub enum ConferenceError {
35 #[error("conference mix interval must be at least 1 ms, got {0:?}")]
37 IntervalTooShort(Duration),
38}
39
40type Members = Arc<Mutex<HashMap<u64, Member>>>;
42
43const MOST_PENDING: usize = 4_000;
49
50struct Member {
51 session: Arc<MediaSession>,
52 pending: Vec<i16>,
54}
55
56#[derive(Debug)]
62struct Workers {
63 closed: bool,
64 collectors: HashMap<u64, JoinHandle<()>>,
65 mixer: Option<JoinHandle<()>>,
66}
67
68#[cfg(test)]
69#[derive(Debug, Default)]
70struct LifecycleHooks {
71 join_before_registration: StdMutex<Option<JoinRegistrationHook>>,
72 close_waiting_for_members: StdMutex<Option<tokio::sync::oneshot::Sender<()>>>,
73}
74
75#[cfg(test)]
76#[derive(Debug)]
77struct JoinRegistrationHook {
78 reached: tokio::sync::oneshot::Sender<()>,
79 release: std::sync::mpsc::Receiver<()>,
80}
81
82impl std::fmt::Debug for Member {
83 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
84 f.debug_struct("Member")
87 .field("pending", &self.pending.len())
88 .finish_non_exhaustive()
89 }
90}
91
92#[derive(Debug)]
98pub struct Conference {
99 members: Members,
100 next_id: std::sync::atomic::AtomicU64,
101 workers: StdMutex<Workers>,
102 samples_per_frame: usize,
103 stop: Arc<Stop>,
104 #[cfg(test)]
105 lifecycle_hooks: LifecycleHooks,
106}
107
108impl Conference {
109 pub fn new(samples_per_frame: usize, interval: Duration) -> Result<Self, ConferenceError> {
120 if interval < Duration::from_millis(1) {
121 return Err(ConferenceError::IntervalTooShort(interval));
122 }
123 let members: Members = Arc::new(Mutex::new(HashMap::new()));
124 let stop = Arc::new(Stop::default());
125 let mixer = tokio::spawn(mix_loop(
126 Arc::clone(&members),
127 samples_per_frame,
128 interval,
129 Arc::clone(&stop),
130 ));
131 Ok(Self {
132 members,
133 next_id: std::sync::atomic::AtomicU64::new(0),
134 workers: StdMutex::new(Workers {
135 closed: false,
136 collectors: HashMap::new(),
137 mixer: Some(mixer),
138 }),
139 samples_per_frame,
140 stop,
141 #[cfg(test)]
142 lifecycle_hooks: LifecycleHooks::default(),
143 })
144 }
145
146 pub fn narrowband() -> Result<Self, ConferenceError> {
153 Self::new(160, Duration::from_millis(20))
154 }
155
156 pub async fn join(&self, session: Arc<MediaSession>) -> u64 {
158 let id = self
159 .next_id
160 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
161
162 let mut participants = self.members.lock().await;
163 let mut workers = self.workers_lock();
164 if workers.closed {
165 return id;
166 }
167 participants.insert(
168 id,
169 Member {
170 session: Arc::clone(&session),
171 pending: Vec::new(),
172 },
173 );
174
175 let members = Arc::clone(&self.members);
180 let stop = Arc::clone(&self.stop);
181 let collector = tokio::spawn(async move {
182 loop {
183 let samples = tokio::select! {
184 () = stop.wait() => return,
185 samples = session.recv() => samples,
186 };
187 let Some(samples) = samples else {
188 return;
189 };
190 let mut members = members.lock().await;
191 let Some(member) = members.get_mut(&id) else {
192 return;
193 };
194 member.pending.extend_from_slice(&samples);
195 if member.pending.len() > MOST_PENDING {
196 let excess = member.pending.len() - MOST_PENDING;
197 member.pending.drain(..excess);
198 }
199 }
200 });
201 #[cfg(test)]
202 self.pause_join_before_registration();
203 workers.collectors.insert(id, collector);
204 id
205 }
206
207 pub async fn leave(&self, id: u64) {
212 let mut members = self.members.lock().await;
213 let collector = self.workers_lock().collectors.remove(&id);
214 members.remove(&id);
215 drop(members);
216 if let Some(collector) = collector {
217 collector.abort();
218 let _ = collector.await;
219 }
220 }
221
222 pub async fn len(&self) -> usize {
224 self.members.lock().await.len()
225 }
226
227 pub async fn is_empty(&self) -> bool {
229 self.len().await == 0
230 }
231
232 #[must_use]
234 pub fn samples_per_frame(&self) -> usize {
235 self.samples_per_frame
236 }
237
238 pub async fn close(&self) {
240 let mut members = self.lock_members_for_close().await;
245 let workers = self.shutdown();
246 members.clear();
247 drop(members);
248 for worker in workers {
249 let _ = worker.await;
253 }
254 }
255
256 fn workers_lock(&self) -> std::sync::MutexGuard<'_, Workers> {
260 match self.workers.lock() {
261 Ok(workers) => workers,
262 Err(poisoned) => poisoned.into_inner(),
263 }
264 }
265
266 async fn lock_members_for_close(&self) -> tokio::sync::MutexGuard<'_, HashMap<u64, Member>> {
267 #[cfg(not(test))]
268 {
269 self.members.lock().await
270 }
271 #[cfg(test)]
272 {
273 use std::future::Future as _;
274 use std::task::Poll;
275
276 let mut waiting = match self.lifecycle_hooks.close_waiting_for_members.lock() {
277 Ok(mut hook) => hook.take(),
278 Err(poisoned) => poisoned.into_inner().take(),
279 };
280 let mut lock = Box::pin(self.members.lock());
281 std::future::poll_fn(|cx| match lock.as_mut().poll(cx) {
282 Poll::Ready(members) => Poll::Ready(members),
283 Poll::Pending => {
284 if let Some(waiting) = waiting.take() {
285 let _ = waiting.send(());
288 }
289 Poll::Pending
290 }
291 })
292 .await
293 }
294 }
295
296 #[cfg(test)]
297 fn pause_join_before_registration(&self) {
298 let hook = match self.lifecycle_hooks.join_before_registration.lock() {
299 Ok(mut hook) => hook.take(),
300 Err(poisoned) => poisoned.into_inner().take(),
301 };
302 if let Some(hook) = hook {
303 let _ = hook.reached.send(());
306 let _ = hook.release.recv_timeout(Duration::from_secs(2));
307 }
308 }
309
310 fn shutdown(&self) -> Vec<JoinHandle<()>> {
312 let mut state = self.workers_lock();
313 state.closed = true;
314 self.stop.stop();
315 let mut workers = Vec::new();
316 if let Some(mixer) = state.mixer.take() {
317 mixer.abort();
318 workers.push(mixer);
319 }
320 workers.extend(state.collectors.drain().map(|(_, collector)| {
321 collector.abort();
322 collector
323 }));
324 workers
325 }
326}
327
328impl Drop for Conference {
329 fn drop(&mut self) {
330 drop(self.shutdown());
333 }
334}
335
336async fn mix_loop(members: Members, samples_per_frame: usize, interval: Duration, stop: Arc<Stop>) {
338 let mut tick = tokio::time::interval(interval);
339 tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
340
341 loop {
342 tokio::select! {
343 () = stop.wait() => return,
344 _ = tick.tick() => {}
345 }
346
347 let (ids, frames, sessions) = {
351 let mut members = members.lock().await;
352 if members.is_empty() {
353 continue;
354 }
355 let mut ids = Vec::with_capacity(members.len());
356 let mut frames = Vec::with_capacity(members.len());
357 let mut sessions = Vec::with_capacity(members.len());
358 for (id, member) in members.iter_mut() {
359 let take = member.pending.len().min(samples_per_frame);
360 let mut frame: Vec<i16> = member.pending.drain(..take).collect();
361 frame.resize(samples_per_frame, 0);
364 ids.push(*id);
365 frames.push(frame);
366 sessions.push(Arc::clone(&member.session));
367 }
368 (ids, frames, sessions)
369 };
370
371 for (index, session) in sessions.iter().enumerate() {
372 let mut mixed = vec![0i16; samples_per_frame];
376 for (other, frame) in frames.iter().enumerate() {
377 if other == index {
378 continue;
379 }
380 mix_into(&mut mixed, frame);
381 }
382 if !session.send(mixed).await {
383 tracing::debug!(id = ids.get(index), "a conference participant has gone");
386 }
387 }
388 }
389}
390
391#[cfg(test)]
392#[allow(
393 clippy::unwrap_used,
394 clippy::expect_used,
395 clippy::panic,
396 clippy::indexing_slicing
397)]
398mod tests {
399 use super::*;
400 use crate::session::{Codec, Config, MediaPort};
401
402 fn set_close_wait_hook(conference: &Conference) -> tokio::sync::oneshot::Receiver<()> {
403 let (waiting, reached) = tokio::sync::oneshot::channel();
404 match conference.lifecycle_hooks.close_waiting_for_members.lock() {
405 Ok(mut hook) => *hook = Some(waiting),
406 Err(poisoned) => *poisoned.into_inner() = Some(waiting),
407 }
408 reached
409 }
410
411 async fn wait_for_close_to_block(reached: tokio::sync::oneshot::Receiver<()>) {
412 tokio::time::timeout(Duration::from_secs(2), reached)
413 .await
414 .expect("close polls the members lock")
415 .expect("close reports its blocked lock poll");
416 }
417
418 #[tokio::test]
419 async fn cancelling_close_while_it_waits_leaves_no_half_closed_conference() {
420 let port = MediaPort::bind("127.0.0.1:0".parse().expect("valid"))
421 .await
422 .expect("binds");
423 let mut config = Config::new("127.0.0.1:9".parse().expect("valid"), Codec::Pcmu);
424 config.rtcp_interval = None;
425 let session = Arc::new(port.start(config).expect("valid media setup"));
426 let weak = Arc::downgrade(&session);
427 let conference = Arc::new(Conference::narrowband().expect("valid conference timing"));
428 conference.join(Arc::clone(&session)).await;
429
430 let members = conference.members.lock().await;
434 let close_waiting = set_close_wait_hook(&conference);
435 let closing = {
436 let conference = Arc::clone(&conference);
437 tokio::spawn(async move { conference.close().await })
438 };
439 wait_for_close_to_block(close_waiting).await;
440 assert!(
441 !closing.is_finished(),
442 "the close future is parked on the held members lock"
443 );
444 closing.abort();
445 let _ = closing.await;
446 assert!(
447 !conference.workers_lock().closed,
448 "cancellation before all locks are held must not half-close the conference"
449 );
450 drop(members);
451
452 drop(session);
453 conference.close().await;
454 assert!(
455 weak.upgrade().is_none(),
456 "a later close releases the participant rather than finding stranded state"
457 );
458 }
459
460 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
461 async fn close_cannot_pass_join_between_collector_spawn_and_registration() {
462 let port = MediaPort::bind("127.0.0.1:0".parse().expect("valid"))
463 .await
464 .expect("binds");
465 let mut config = Config::new("127.0.0.1:9".parse().expect("valid"), Codec::Pcmu);
466 config.rtcp_interval = None;
467 let session = Arc::new(port.start(config).expect("valid media setup"));
468 let weak = Arc::downgrade(&session);
469 let conference = Arc::new(Conference::narrowband().expect("valid conference timing"));
470
471 let (join_reached, reached) = tokio::sync::oneshot::channel();
472 let (release, join_release) = std::sync::mpsc::sync_channel(0);
473 {
474 let mut hook = match conference.lifecycle_hooks.join_before_registration.lock() {
475 Ok(hook) => hook,
476 Err(poisoned) => poisoned.into_inner(),
477 };
478 *hook = Some(JoinRegistrationHook {
479 reached: join_reached,
480 release: join_release,
481 });
482 }
483
484 let joining = {
485 let conference = Arc::clone(&conference);
486 let session = Arc::clone(&session);
487 tokio::spawn(async move { conference.join(session).await })
488 };
489 tokio::time::timeout(Duration::from_secs(2), reached)
490 .await
491 .expect("join reaches the registration boundary")
492 .expect("join reports the registration boundary");
493
494 let close_waiting = set_close_wait_hook(&conference);
495 let closing = {
496 let conference = Arc::clone(&conference);
497 tokio::spawn(async move { conference.close().await })
498 };
499 wait_for_close_to_block(close_waiting).await;
500 assert!(
501 !closing.is_finished(),
502 "close cannot drain workers while join owns the lifecycle transition"
503 );
504
505 release
506 .send(())
507 .expect("join is released to register its collector");
508 joining.await.expect("join task finishes");
509 closing.await.expect("close task finishes");
510 drop(session);
511
512 assert!(conference.is_empty().await);
513 assert!(
514 weak.upgrade().is_none(),
515 "close drains the collector registered by the serialized join"
516 );
517 }
518}