Skip to main content

bt_broadcast_assistant/assistant/
event.rs

1// Copyright 2023 The Fuchsia Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style license that can be
3// found in the LICENSE file.
4
5use 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
55/// A stream of discovered broadcast sources.
56/// This stream polls the scan results from GATT client to discover
57/// available broadcast sources.
58pub 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    /// Polls the futures currently establishing periodic advertising syncs.
98    ///
99    /// Returns `Poll::Ready(())` if any future resolved (successfully
100    /// establishing a sync or failing), which indicates progress was made.
101    /// Returns `Poll::Pending` if no progress was made.
102    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    /// Polls the active periodic advertising sync report streams.
121    ///
122    /// Returns:
123    /// - `Poll::Ready(Some(event))` if progress was made and a completed
124    ///   `FoundBroadcastSource` event is ready to be returned.
125    /// - `Poll::Ready(None)` if progress was made (e.g., a report was processed
126    ///   but was incomplete, or a stream failed and was removed), but no event
127    ///   is ready yet.
128    /// - `Poll::Pending` if no progress was made.
129    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), // Progressed, but no event
139                }
140            }
141            Poll::Ready(Some((peer_id, sid, Err(_)))) => {
142                self.active_syncs.remove(&(peer_id, sid));
143                Poll::Ready(None) // Progressed, but no event
144            }
145            Poll::Ready(None) | Poll::Pending => Poll::Pending,
146        }
147    }
148
149    /// Returns the broadcast source if the scanned peer is a broadcast source.
150    /// Returns an error if parsing of the scan result data fails and None if
151    /// the scanned peer is not a broadcast source.
152    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        // See if it's appropriate to establish PA sync for this broadcast source.
258
259        // If we already have the endpoint data (BASE), we don't need to establish a
260        // sync.
261        if broadcast_source.endpoint.is_some() {
262            return None;
263        }
264
265        // If the platform doesn't support periodic advertising, we cannot sync.
266        let Some(ref pa) = self.periodic_advertising else {
267            return None;
268        };
269
270        // If we are already actively syncing (or establishing a sync) for this
271        // peer/SID, don't start another one.
272        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        // Scanned a broadcast source and its broadcast id.
415        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        // Found broadcast source event shouldn't have been sent since braodcast source
433        // information isn't complete.
434        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        // Poll again to let the PA sync transition from Establishing to Established.
438        assert!(stream.poll_next_unpin(&mut noop_cx).is_pending());
439
440        // Pretend somehow address, address type were filled out.
441        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        // Scanned broadcast source's BASE data:
449        #[rustfmt::skip]
450        let base_data = vec![
451            0x10, 0x20, 0x30, 0x02,               // presentation delay, num of subgroups
452            0x01, 0x03, 0x00, 0x00, 0x00, 0x00,   // num of bis, codec id (big #1)
453            0x00,                                 // codec specific config len
454            0x00,                                 // metadata len,
455            0x01, 0x00,                           // bis index, codec specific config len (big #1 / bis #1)
456            0x01, 0x02, 0x00, 0x00, 0x00, 0x00,   // num of bis, codec id (big #2)
457            0x00,                                 // codec specific config len
458            0x00,                                 // metadata len,
459            0x01, 0x03, 0x02, 0x05, 0x08,         // bis index, codec specific config len, codec frame blocks LTV (big #2 / bis #2)
460        ];
461
462        let advertising_data =
463            vec![AdvertisingDatum::ServiceData(BASIC_AUDIO_ANNOUNCEMENT_SERVICE, base_data)];
464
465        // Get the fake periodic advertising sync sender that was registered when we
466        // polled the scan result
467        let periodic_advertising_sender = pa
468            .get_sender(broadcast_source_pid)
469            .expect("should have registered periodic advertising sync sender");
470
471        // Send the PA report containing the BASE data
472        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        // Expect the stream to send out broadcast source found event since information
485        // is complete.
486        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        // Verify that the PA sync was stopped (removed from active_syncs) to conserve
498        // resources
499        assert!(stream.active_syncs.is_empty());
500
501        // Subsequent polls should be pending
502        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        // Mimick scan error.
510        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        // Entire stream should have terminated.
519        assert_matches!(stream.poll_next_unpin(&mut noop_cx), Poll::Ready(None));
520        assert_matches!(stream.is_terminated(), true);
521    }
522}