1use core::pin::Pin;
6use futures::stream::{
7 abortable, AbortHandle, FusedStream, FuturesUnordered, SelectAll, Stream, StreamExt,
8};
9use futures::FutureExt;
10use std::collections::HashMap;
11use std::future::Future;
12use std::sync::atomic::{AtomicBool, Ordering};
13use std::sync::Arc;
14use std::task::Poll;
15
16use bt_bap::types::{BroadcastAudioSourceEndpoint, BroadcastId};
17use bt_common::core::{AdvertisingSetId, PeriodicAdvertisingInterval};
18use bt_common::packet_encoding::Decodable;
19use bt_common::packet_encoding::Error as PacketError;
20use bt_common::PeerId;
21use bt_gatt::central::{AdvertisingDatum, ScanResult};
22use bt_gatt::periodic_advertising::{PeriodicAdvertising, SyncConfiguration, SyncReport};
23use bt_gatt::GattTypes;
24
25use crate::assistant::{
26 DiscoveredBroadcastSources, Error, BASIC_AUDIO_ANNOUNCEMENT_SERVICE,
27 BROADCAST_AUDIO_ANNOUNCEMENT_SERVICE,
28};
29use crate::types::BroadcastSource;
30
31#[derive(Debug)]
32pub enum Event {
33 FoundBroadcastSource {
34 peer: PeerId,
35 advertising_sid: AdvertisingSetId,
36 source: BroadcastSource,
37 },
38 CouldNotParseAdvertisingData {
39 peer: PeerId,
40 error: PacketError,
41 },
42}
43
44type PeriodicAdvertisingSyncResult<T> = Result<
45 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncStream,
46 bt_gatt::types::Error,
47>;
48
49type PeriodicAdvertisingFuture<T> =
50 Pin<Box<dyn Future<Output = (PeerId, u8, PeriodicAdvertisingSyncResult<T>)>>>;
51
52type PeriodicAdvertisingStream =
53 Pin<Box<dyn Stream<Item = (PeerId, u8, Result<SyncReport, bt_gatt::types::Error>)>>>;
54
55pub struct EventStream<T: bt_gatt::GattTypes> {
59 scan_result_stream: Pin<Box<<T as bt_gatt::GattTypes>::ScanResultStream>>,
60 terminated: bool,
61
62 broadcast_sources: Arc<DiscoveredBroadcastSources>,
63 broadcast_source_scan_started: Arc<AtomicBool>,
64
65 periodic_advertising: Option<T::PeriodicAdvertising>,
66
67 establishing_periodic_advertising_syncs: FuturesUnordered<PeriodicAdvertisingFuture<T>>,
68 active_periodic_advertising_sync_streams: SelectAll<PeriodicAdvertisingStream>,
69 active_syncs: HashMap<(PeerId, u8), Option<AbortHandle>>,
70}
71
72impl<T: bt_gatt::GattTypes> Unpin for EventStream<T> {}
73
74impl<T: bt_gatt::GattTypes> EventStream<T>
75where
76 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncStream: 'static,
77 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncFut: 'static,
78{
79 pub(crate) fn new(
80 scan_result_stream: T::ScanResultStream,
81 periodic_advertising: Option<T::PeriodicAdvertising>,
82 broadcast_sources: Arc<DiscoveredBroadcastSources>,
83 broadcast_source_scan_started: Arc<AtomicBool>,
84 ) -> Self {
85 Self {
86 scan_result_stream: Box::pin(scan_result_stream),
87 terminated: false,
88 broadcast_sources,
89 broadcast_source_scan_started,
90 periodic_advertising,
91 establishing_periodic_advertising_syncs: FuturesUnordered::new(),
92 active_periodic_advertising_sync_streams: SelectAll::new(),
93 active_syncs: HashMap::new(),
94 }
95 }
96
97 fn poll_establishing_syncs(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
103 if self.establishing_periodic_advertising_syncs.is_terminated() {
104 return Poll::Pending;
105 }
106
107 match self.establishing_periodic_advertising_syncs.poll_next_unpin(cx) {
108 Poll::Ready(Some((peer_id, sid, Ok(stream)))) => {
109 self.handle_established_sync(peer_id, sid, stream);
110 Poll::Ready(())
111 }
112 Poll::Ready(Some((peer_id, sid, Err(_)))) => {
113 self.active_syncs.remove(&(peer_id, sid));
114 Poll::Ready(())
115 }
116 Poll::Ready(None) | Poll::Pending => Poll::Pending,
117 }
118 }
119
120 fn poll_active_syncs(&mut self, cx: &mut std::task::Context<'_>) -> Poll<Option<Event>> {
130 if self.active_periodic_advertising_sync_streams.is_terminated() {
131 return Poll::Pending;
132 }
133
134 match self.active_periodic_advertising_sync_streams.poll_next_unpin(cx) {
135 Poll::Ready(Some((peer_id, sid, Ok(report)))) => {
136 match self.handle_periodic_advertising_report(peer_id, sid, report) {
137 Some(event) => Poll::Ready(Some(event)),
138 None => Poll::Ready(None), }
140 }
141 Poll::Ready(Some((peer_id, sid, Err(_)))) => {
142 self.active_syncs.remove(&(peer_id, sid));
143 Poll::Ready(None) }
145 Poll::Ready(None) | Poll::Pending => Poll::Pending,
146 }
147 }
148
149 fn try_into_broadcast_source(
153 scan_result: &ScanResult,
154 ) -> Result<Option<BroadcastSource>, PacketError> {
155 let mut source = None;
156 for datum in &scan_result.advertised {
157 match datum {
158 AdvertisingDatum::ServiceData(uuid, data)
159 if *uuid == BROADCAST_AUDIO_ANNOUNCEMENT_SERVICE =>
160 {
161 let bid = BroadcastId::decode(data.as_slice()).0?;
162 source.get_or_insert(BroadcastSource::default()).with_broadcast_id(bid);
163 }
164 AdvertisingDatum::BroadcastName(name) => {
165 source
166 .get_or_insert(BroadcastSource::default())
167 .with_broadcast_name(name.clone());
168 }
169 _ => {}
170 }
171 }
172 if let Some(src) = &mut source {
173 src.periodic_advertising_interval =
174 scan_result.periodic_advertising_interval.map(PeriodicAdvertisingInterval);
175 }
176 Ok(source)
177 }
178
179 fn handle_established_sync(
180 &mut self,
181 peer_id: PeerId,
182 sid: u8,
183 stream: <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncStream,
184 ) {
185 let mapped_stream = stream.map(move |report| (peer_id, sid, report));
186 let (abortable_stream, abort_handle) = abortable(mapped_stream);
187
188 self.active_periodic_advertising_sync_streams.push(Box::pin(abortable_stream));
189 self.active_syncs.insert((peer_id, sid), Some(abort_handle));
190 }
191
192 fn handle_periodic_advertising_report(
193 &mut self,
194 peer_id: PeerId,
195 sid: u8,
196 report: SyncReport,
197 ) -> Option<Event> {
198 let SyncReport::PeriodicAdvertisingReport(report) = report else {
199 return None;
200 };
201
202 let Some(base) = parse_base_from_advertising_data(&report.data) else {
203 return None;
204 };
205
206 let Ok(advertising_sid) = AdvertisingSetId::try_from(sid) else {
207 return None;
208 };
209
210 let (broadcast_source, changed) = self.broadcast_sources.merge_broadcast_source_data(
211 &(peer_id, advertising_sid),
212 &BroadcastSource::default().with_endpoint(base),
213 );
214
215 if broadcast_source.is_ready_to_add() && changed {
216 if let Some(Some(handle)) =
217 self.active_syncs.remove(&(peer_id, advertising_sid.value()))
218 {
219 handle.abort();
220 }
221 return Some(Event::FoundBroadcastSource {
222 peer: peer_id,
223 advertising_sid,
224 source: broadcast_source,
225 });
226 }
227 None
228 }
229
230 fn handle_scan_result(&mut self, scanned: ScanResult) -> Option<Event> {
231 let found_source = match Self::try_into_broadcast_source(&scanned) {
232 Ok(Some(src)) => src,
233 Ok(None) => return None,
234 Err(e) => {
235 return Some(Event::CouldNotParseAdvertisingData { peer: scanned.id, error: e })
236 }
237 };
238
239 let Some(raw_sid) = scanned.advertising_sid else {
240 return None;
241 };
242 let Ok(sid) = AdvertisingSetId::try_from(raw_sid) else {
243 return None;
244 };
245
246 let (broadcast_source, changed) =
247 self.broadcast_sources.merge_broadcast_source_data(&(scanned.id, sid), &found_source);
248
249 if broadcast_source.is_ready_to_add() && changed {
250 return Some(Event::FoundBroadcastSource {
251 peer: scanned.id,
252 advertising_sid: sid,
253 source: broadcast_source,
254 });
255 }
256
257 if broadcast_source.endpoint.is_some() {
262 return None;
263 }
264
265 let Some(ref pa) = self.periodic_advertising else {
267 return None;
268 };
269
270 let key = (scanned.id, sid.value());
273 if self.active_syncs.contains_key(&key) {
274 return None;
275 }
276
277 self.active_syncs.insert(key, None);
278 let fut = pa.sync_to_advertising_reports(
279 scanned.id,
280 sid.value(),
281 SyncConfiguration { filter_duplicates: true },
282 );
283 let mapped_fut = fut.map(move |res| (scanned.id, sid.value(), res));
284 self.establishing_periodic_advertising_syncs.push(Box::pin(mapped_fut));
285
286 None
287 }
288}
289
290fn parse_base_from_advertising_data(
291 data: &[AdvertisingDatum],
292) -> Option<BroadcastAudioSourceEndpoint> {
293 for datum in data {
294 let AdvertisingDatum::ServiceData(uuid, service_data) = datum else {
295 continue;
296 };
297 if *uuid != BASIC_AUDIO_ANNOUNCEMENT_SERVICE {
298 continue;
299 }
300 let (Ok(base), _) = BroadcastAudioSourceEndpoint::decode(service_data) else {
301 continue;
302 };
303 return Some(base);
304 }
305 None
306}
307
308impl<T: bt_gatt::GattTypes> Drop for EventStream<T> {
309 fn drop(&mut self) {
310 self.broadcast_source_scan_started.store(false, Ordering::Relaxed);
311 }
312}
313
314impl<T: bt_gatt::GattTypes> FusedStream for EventStream<T>
315where
316 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncStream: 'static,
317 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncFut: 'static,
318{
319 fn is_terminated(&self) -> bool {
320 self.terminated
321 }
322}
323
324impl<T: bt_gatt::GattTypes> Stream for EventStream<T>
325where
326 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncStream: 'static,
327 <<T as GattTypes>::PeriodicAdvertising as PeriodicAdvertising>::SyncFut: 'static,
328{
329 type Item = Result<Event, Error>;
330
331 fn poll_next(
332 mut self: std::pin::Pin<&mut Self>,
333 cx: &mut std::task::Context<'_>,
334 ) -> Poll<Option<Self::Item>> {
335 if self.terminated {
336 return Poll::Ready(None);
337 }
338
339 loop {
340 let mut progressed = false;
341
342 if self.poll_establishing_syncs(cx).is_ready() {
343 progressed = true;
344 }
345
346 match self.poll_active_syncs(cx) {
347 Poll::Ready(Some(event)) => return Poll::Ready(Some(Ok(event))),
348 Poll::Ready(None) => progressed = true,
349 Poll::Pending => {}
350 }
351
352 match self.scan_result_stream.poll_next_unpin(cx) {
353 Poll::Ready(Some(Ok(scanned))) => {
354 progressed = true;
355 if let Some(event) = self.handle_scan_result(scanned) {
356 return Poll::Ready(Some(Ok(event)));
357 }
358 }
359 Poll::Ready(None | Some(Err(_))) => {
360 self.terminated = true;
361 self.broadcast_source_scan_started.store(false, Ordering::Relaxed);
362 return Poll::Ready(Some(Err(Error::CentralScanTerminated)));
363 }
364 Poll::Pending => {}
365 }
366
367 if !progressed {
368 break;
369 }
370 }
371
372 Poll::Pending
373 }
374}
375
376#[cfg(test)]
377mod tests {
378 use super::*;
379
380 use assert_matches::assert_matches;
381
382 use bt_common::core::{AddressType, AdvertisingSetId};
383 use bt_gatt::central::{AdvertisingDatum, PeerName};
384 use bt_gatt::test_utils::{
385 FakePeriodicAdvertising, FakeTypes, ScannedResultStream, ScannedResultStreamController,
386 };
387 use bt_gatt::types::Error as BtGattError;
388 use bt_gatt::types::GattError;
389
390 fn setup_stream(
391 ) -> (EventStream<FakeTypes>, ScannedResultStreamController, FakePeriodicAdvertising) {
392 let fake_scan_result_stream = ScannedResultStream::new();
393 let controller = fake_scan_result_stream.controller();
394 let broadcast_sources = DiscoveredBroadcastSources::new();
395 let broadcast_source_scan_started = Arc::new(AtomicBool::new(false));
396 let pa = FakePeriodicAdvertising::default();
397
398 (
399 EventStream::<FakeTypes>::new(
400 fake_scan_result_stream,
401 Some(pa.clone()),
402 broadcast_sources,
403 broadcast_source_scan_started,
404 ),
405 controller,
406 pa,
407 )
408 }
409
410 #[test]
411 fn poll_found_broadcast_source_events() {
412 let (mut stream, scan_result_controller, pa) = setup_stream();
413
414 let broadcast_source_pid = PeerId(1005);
416
417 scan_result_controller.add_scanned_result(Ok(ScanResult {
418 id: broadcast_source_pid,
419 connectable: true,
420 name: PeerName::Unknown,
421 advertised: vec![
422 AdvertisingDatum::ServiceData(
423 BROADCAST_AUDIO_ANNOUNCEMENT_SERVICE,
424 vec![0x01, 0x02, 0x03],
425 ),
426 AdvertisingDatum::BroadcastName("Test Broadcast".to_string()),
427 ],
428 advertising_sid: Some(1),
429 periodic_advertising_interval: Some(0x0100),
430 }));
431
432 let mut noop_cx = futures::task::Context::from_waker(futures::task::noop_waker_ref());
435 assert!(stream.poll_next_unpin(&mut noop_cx).is_pending());
436
437 assert!(stream.poll_next_unpin(&mut noop_cx).is_pending());
439
440 let _ = stream.broadcast_sources.merge_broadcast_source_data(
442 &(broadcast_source_pid, AdvertisingSetId::try_from(1).unwrap()),
443 &BroadcastSource::default()
444 .with_address([1, 2, 3, 4, 5, 6])
445 .with_address_type(AddressType::Public),
446 );
447
448 #[rustfmt::skip]
450 let base_data = vec![
451 0x10, 0x20, 0x30, 0x02, 0x01, 0x03, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x01, 0x02, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x03, 0x02, 0x05, 0x08, ];
461
462 let advertising_data =
463 vec![AdvertisingDatum::ServiceData(BASIC_AUDIO_ANNOUNCEMENT_SERVICE, base_data)];
464
465 let periodic_advertising_sender = pa
468 .get_sender(broadcast_source_pid)
469 .expect("should have registered periodic advertising sync sender");
470
471 periodic_advertising_sender
473 .unbounded_send(Ok(SyncReport::PeriodicAdvertisingReport(
474 bt_gatt::periodic_advertising::PeriodicAdvertisingReport {
475 rssi: -50,
476 data: advertising_data,
477 event_counter: None,
478 subevent: None,
479 timestamp: 0,
480 },
481 )))
482 .unwrap();
483
484 let Poll::Ready(Some(Ok(event))) = stream.poll_next_unpin(&mut noop_cx) else {
487 panic!("should have received event");
488 };
489 assert_matches!(event, Event::FoundBroadcastSource { peer, advertising_sid, source } => {
490 assert_eq!(peer, broadcast_source_pid);
491 assert_eq!(advertising_sid, AdvertisingSetId::try_from(1).unwrap());
492 assert_eq!(source.periodic_advertising_interval, Some(PeriodicAdvertisingInterval(0x0100)));
493 assert_eq!(source.address, Some([1, 2, 3, 4, 5, 6]));
494 assert_eq!(source.broadcast_name, Some("Test Broadcast".to_string()));
495 });
496
497 assert!(stream.active_syncs.is_empty());
500
501 assert!(stream.poll_next_unpin(&mut noop_cx).is_pending());
503 }
504
505 #[test]
506 fn central_scan_stream_terminates() {
507 let (mut stream, scan_result_controller, _pa) = setup_stream();
508
509 scan_result_controller.add_scanned_result(Err(BtGattError::Gatt(GattError::InvalidPdu)));
511
512 let mut noop_cx = futures::task::Context::from_waker(futures::task::noop_waker_ref());
513 match stream.poll_next_unpin(&mut noop_cx) {
514 Poll::Ready(Some(Err(e))) => assert_matches!(e, Error::CentralScanTerminated),
515 _ => panic!("should have received central scan terminated error"),
516 }
517
518 assert_matches!(stream.poll_next_unpin(&mut noop_cx), Poll::Ready(None));
520 assert_matches!(stream.is_terminated(), true);
521 }
522}