Skip to main content

socketioxide_core/adapter/
mod.rs

1//! The adapter module contains the [`CoreAdapter`] trait and other related types.
2//!
3//! It is used to implement communication between socket.io servers to share messages and state.
4//!
5//! The [`CoreLocalAdapter`] provide a local implementation that will allow any implementors to apply local
6//! operations (`broadcast_with_ack`, `broadcast`, `rooms`, etc...).
7//!
8use std::{
9    borrow::Cow,
10    collections::{HashMap, HashSet, hash_map, hash_set},
11    error::Error as StdError,
12    future::{self, Future},
13    hash::Hash,
14    slice,
15    sync::{Arc, RwLock},
16    time::Duration,
17};
18
19use engineioxide_core::{Sid, Str};
20use futures_core::{FusedStream, Stream};
21use serde::{Deserialize, Serialize, de::DeserializeOwned};
22use smallvec::SmallVec;
23
24use crate::{Uid, Value, packet::Packet, parser::Parse};
25use errors::{AdapterError, BroadcastError, SocketError};
26
27pub mod errors;
28#[cfg(feature = "remote-adapter")]
29pub mod heartbeat;
30#[cfg(feature = "remote-adapter")]
31pub mod remote_packet;
32#[cfg(feature = "remote-adapter")]
33pub mod stream;
34
35/// A room identifier
36pub type Room = Cow<'static, str>;
37
38/// Flags that can be used to modify the behavior of the broadcast methods.
39#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
40pub enum BroadcastFlags {
41    /// Broadcast only to the current server
42    Local = 0x01,
43    /// Broadcast to all clients except the sender
44    Broadcast = 0x02,
45    /// The event may be dropped if the client is not ready to receive it
46    /// (e.g. the connection is buffering or not connected).
47    /// This is useful for events that are not critical, like position updates in a game.
48    /// See [socket.io volatile events](https://socket.io/docs/v4/emitting-events/#volatile-events).
49    Volatile = 0x04,
50}
51
52/// Options that can be used to modify the behavior of the broadcast methods.
53#[derive(Clone, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
54pub struct BroadcastOptions {
55    /// The flags to apply to the broadcast represented as a bitflag.
56    flags: u8,
57    /// The rooms to broadcast to.
58    pub rooms: SmallVec<[Room; 4]>,
59    /// The rooms to exclude from the broadcast.
60    pub except: SmallVec<[Room; 4]>,
61    /// The socket id of the sender.
62    pub sid: Option<Sid>,
63    /// The target server id can be used to optimize the broadcast.
64    /// More specifically when we use broadcasting to apply a single action on a remote socket.
65    /// We now the server_id of the remote socket, so we can send the action directly to the server.
66    pub server_id: Option<Uid>,
67}
68impl BroadcastOptions {
69    /// Add any flags to the options.
70    pub fn add_flag(&mut self, flag: BroadcastFlags) {
71        self.flags |= flag as u8;
72    }
73    /// Check if the options have a flag.
74    pub fn has_flag(&self, flag: BroadcastFlags) -> bool {
75        self.flags & flag as u8 == flag as u8
76    }
77
78    /// get the flags of the options.
79    pub fn flags(&self) -> u8 {
80        self.flags
81    }
82
83    /// Set the socket id of the sender.
84    pub fn new(sid: Sid) -> Self {
85        Self {
86            sid: Some(sid),
87            ..Default::default()
88        }
89    }
90    /// Create a new broadcast options from a remote socket data.
91    pub fn new_remote(data: &RemoteSocketData) -> Self {
92        Self {
93            sid: Some(data.id),
94            server_id: Some(data.server_id),
95            ..Default::default()
96        }
97    }
98
99    /// Check if the selected options are local to the current server.
100    #[inline]
101    pub fn is_local(&self, uid: Uid) -> bool {
102        let target_sock_is_local = !self.has_flag(BroadcastFlags::Broadcast)
103            && self.server_id == Some(uid)
104            && self.rooms.is_empty()
105            && self.sid.is_some();
106        self.has_flag(BroadcastFlags::Local) || target_sock_is_local
107    }
108}
109
110/// A trait for types that can be used as a room parameter.
111///
112/// [`String`], [`Vec<String>`], [`Vec<&str>`], [`&'static str`](str) and const arrays are implemented by default.
113pub trait RoomParam: Send + 'static {
114    /// The type of the iterator returned by `into_room_iter`.
115    type IntoIter: Iterator<Item = Room>;
116
117    /// Convert `self` into an iterator of rooms.
118    fn into_room_iter(self) -> Self::IntoIter;
119}
120
121impl RoomParam for Room {
122    type IntoIter = std::iter::Once<Room>;
123    #[inline(always)]
124    fn into_room_iter(self) -> Self::IntoIter {
125        std::iter::once(self)
126    }
127}
128impl RoomParam for String {
129    type IntoIter = std::iter::Once<Room>;
130    #[inline(always)]
131    fn into_room_iter(self) -> Self::IntoIter {
132        std::iter::once(Cow::Owned(self))
133    }
134}
135impl RoomParam for Vec<String> {
136    type IntoIter = std::iter::Map<std::vec::IntoIter<String>, fn(String) -> Room>;
137    #[inline(always)]
138    fn into_room_iter(self) -> Self::IntoIter {
139        self.into_iter().map(Cow::Owned)
140    }
141}
142impl RoomParam for Vec<&'static str> {
143    type IntoIter = std::iter::Map<std::vec::IntoIter<&'static str>, fn(&'static str) -> Room>;
144    #[inline(always)]
145    fn into_room_iter(self) -> Self::IntoIter {
146        self.into_iter().map(Cow::Borrowed)
147    }
148}
149
150impl RoomParam for Vec<Room> {
151    type IntoIter = std::vec::IntoIter<Room>;
152    #[inline(always)]
153    fn into_room_iter(self) -> Self::IntoIter {
154        self.into_iter()
155    }
156}
157impl RoomParam for &'static str {
158    type IntoIter = std::iter::Once<Room>;
159    #[inline(always)]
160    fn into_room_iter(self) -> Self::IntoIter {
161        std::iter::once(Cow::Borrowed(self))
162    }
163}
164impl<const COUNT: usize> RoomParam for [&'static str; COUNT] {
165    type IntoIter =
166        std::iter::Map<std::array::IntoIter<&'static str, COUNT>, fn(&'static str) -> Room>;
167
168    #[inline(always)]
169    fn into_room_iter(self) -> Self::IntoIter {
170        self.into_iter().map(Cow::Borrowed)
171    }
172}
173impl RoomParam for &'static [&'static str] {
174    type IntoIter =
175        std::iter::Map<std::slice::Iter<'static, &'static str>, fn(&'static &'static str) -> Room>;
176
177    #[inline(always)]
178    fn into_room_iter(self) -> Self::IntoIter {
179        self.iter().map(|i| Cow::Borrowed(*i))
180    }
181}
182impl<const COUNT: usize> RoomParam for [String; COUNT] {
183    type IntoIter = std::iter::Map<std::array::IntoIter<String, COUNT>, fn(String) -> Room>;
184    #[inline(always)]
185    fn into_room_iter(self) -> Self::IntoIter {
186        self.into_iter().map(Cow::Owned)
187    }
188}
189impl RoomParam for Sid {
190    type IntoIter = std::iter::Once<Room>;
191    #[inline(always)]
192    fn into_room_iter(self) -> Self::IntoIter {
193        std::iter::once(Cow::Owned(self.to_string()))
194    }
195}
196
197/// A item yield by the ack stream.
198pub type AckStreamItem<E> = (Sid, Result<Value, E>);
199/// The [`SocketEmitter`] will be implemented by the socketioxide library.
200/// It is simply used as an abstraction to allow the adapter to communicate
201/// with the socket server without the need to depend on the socketioxide lib.
202pub trait SocketEmitter: Send + Sync + 'static {
203    /// An error that can occur when sending data an acknowledgment.
204    type AckError: StdError + Send + Serialize + DeserializeOwned + 'static;
205    /// A stream that emits the acknowledgments of multiple sockets.
206    type AckStream: Stream<Item = AckStreamItem<Self::AckError>> + FusedStream + Send + 'static;
207
208    /// Get all the socket ids in the namespace.
209    fn get_all_sids(&self, filter: impl Fn(&Sid) -> bool) -> Vec<Sid>;
210    /// Get the socket data that match the list of socket ids.
211    fn get_remote_sockets(&self, sids: BroadcastIter<'_>) -> Vec<RemoteSocketData>;
212    /// Send data to the list of socket ids.
213    fn send_many(&self, sids: BroadcastIter<'_>, data: Value) -> Result<(), Vec<SocketError>>;
214    /// Send data to the list of socket ids with volatile semantics.
215    /// Errors are silently discarded; packets may be dropped if the
216    /// transport is not ready.
217    fn send_many_volatile(&self, sids: BroadcastIter<'_>, data: Value);
218    /// Send data to the list of socket ids and get a stream of acks and the number of expected acks.
219    fn send_many_with_ack(
220        &self,
221        sids: BroadcastIter<'_>,
222        packet: Packet,
223        timeout: Option<Duration>,
224    ) -> (Self::AckStream, u32);
225    /// Disconnect all the sockets in the list.
226    /// TODO: take a [`BroadcastIter`]. Currently it is impossible because it may create deadlocks
227    /// with Adapter::del_all call.
228    fn disconnect_many(&self, sids: Vec<Sid>) -> Result<(), Vec<SocketError>>;
229    /// Get the path of the namespace.
230    fn path(&self) -> &Str;
231    /// Get the parser of the namespace.
232    fn parser(&self) -> impl Parse;
233    /// Get the unique server id.
234    fn server_id(&self) -> Uid;
235    /// Get the default configured ack timeout.
236    fn ack_timeout(&self) -> Duration;
237}
238
239/// For static namespaces, the init response will be managed by the user.
240/// However, for dynamic namespaces, the socket.io client will manage the response.
241/// As it does not know the type of the response, the spawnable trait is used to spawn the response.
242/// Without the client having to know the type of the response.
243pub trait Spawnable {
244    /// Spawn the response. Implementors should spawn the future with `tokio::spawn` if it is an async function.
245    /// They should also print a `tracing::error` log in case of an error.
246    fn spawn(self);
247}
248impl Spawnable for () {
249    fn spawn(self) {}
250}
251
252/// A trait to add a "defined" bound to adapter types.
253/// This allow the socket io library to implement function given a *defined* adapter
254/// and not a generic `A: Adapter`.
255///
256/// This is useful to force the user to handle potential init response type [`CoreAdapter::InitRes`].
257pub trait DefinedAdapter {}
258
259/// An adapter is responsible for managing the state of the namespace.
260/// This adapter can be implemented to share the state between multiple servers.
261///
262/// A [`CoreLocalAdapter`] instance will be given when constructing this type, it will allow
263/// you to manipulate local sockets (emitting, fetching data, broadcasting).
264pub trait CoreAdapter<E: SocketEmitter>: Sized + Send + Sync + 'static {
265    /// An error that can occur when using the adapter.
266    type Error: StdError + Into<AdapterError> + Send + 'static;
267    /// A shared state between all the namespace [`CoreAdapter`].
268    /// This can be used to share a connection for example.
269    type State: Send + Sync + 'static;
270    /// A stream that emits the acknowledgments of multiple sockets.
271    type AckStream: Stream<Item = AckStreamItem<E::AckError>> + FusedStream + Send + 'static;
272    /// A named result type for the initialization of the adapter.
273    type InitRes: Spawnable + Send;
274
275    /// Creates a new adapter with the given state and local adapter.
276    ///
277    /// The state is used to share a common state between all your adapters. E.G. a connection to a remote system.
278    /// The local adapter is used to manipulate the local sockets.
279    fn new(state: &Self::State, local: CoreLocalAdapter<E>) -> Self;
280
281    /// Initializes the adapter. The on_success callback should be called when the adapter ready.
282    fn init(self: Arc<Self>, on_success: impl FnOnce() + Send + 'static) -> Self::InitRes;
283
284    /// Closes the adapter.
285    fn close(&self) -> impl Future<Output = Result<(), Self::Error>> + Send {
286        future::ready(Ok(()))
287    }
288
289    /// Returns the number of servers.
290    fn server_count(&self) -> impl Future<Output = Result<u16, Self::Error>> + Send {
291        future::ready(Ok(1))
292    }
293
294    /// Broadcasts the packet to the sockets that match the [`BroadcastOptions`].
295    fn broadcast(
296        &self,
297        packet: Packet,
298        opts: BroadcastOptions,
299    ) -> impl Future<Output = Result<(), BroadcastError>> + Send {
300        future::ready(
301            self.get_local()
302                .broadcast(packet, opts)
303                .map_err(BroadcastError::from),
304        )
305    }
306
307    /// Broadcasts the packet to the sockets that match the [`BroadcastOptions`]
308    /// and return a stream of ack responses.
309    ///
310    /// This method does not have default implementation because GAT cannot have default impls.
311    /// <https://github.com/rust-lang/rust/issues/29661>
312    fn broadcast_with_ack(
313        &self,
314        packet: Packet,
315        opts: BroadcastOptions,
316        timeout: Option<Duration>,
317    ) -> impl Future<Output = Result<Self::AckStream, Self::Error>> + Send;
318
319    /// Adds the sockets that match the [`BroadcastOptions`] to the rooms.
320    fn add_sockets(
321        &self,
322        opts: BroadcastOptions,
323        rooms: impl RoomParam,
324    ) -> impl Future<Output = Result<(), Self::Error>> + Send {
325        self.get_local().add_sockets(opts, rooms);
326        future::ready(Ok(()))
327    }
328
329    /// Removes the sockets that match the [`BroadcastOptions`] from the rooms.
330    fn del_sockets(
331        &self,
332        opts: BroadcastOptions,
333        rooms: impl RoomParam,
334    ) -> impl Future<Output = Result<(), Self::Error>> + Send {
335        self.get_local().del_sockets(opts, rooms);
336        future::ready(Ok(()))
337    }
338
339    /// Disconnects the sockets that match the [`BroadcastOptions`].
340    fn disconnect_socket(
341        &self,
342        opts: BroadcastOptions,
343    ) -> impl Future<Output = Result<(), BroadcastError>> + Send {
344        future::ready(
345            self.get_local()
346                .disconnect_socket(opts)
347                .map_err(BroadcastError::Socket),
348        )
349    }
350
351    /// Fetches rooms that match the [`BroadcastOptions`]
352    fn rooms(
353        &self,
354        opts: BroadcastOptions,
355    ) -> impl Future<Output = Result<Vec<Room>, Self::Error>> + Send {
356        future::ready(Ok(self.get_local().rooms(opts).into_iter().collect()))
357    }
358
359    /// Fetches remote sockets that match the [`BroadcastOptions`].
360    fn fetch_sockets(
361        &self,
362        opts: BroadcastOptions,
363    ) -> impl Future<Output = Result<Vec<RemoteSocketData>, Self::Error>> + Send {
364        future::ready(Ok(self.get_local().fetch_sockets(opts)))
365    }
366
367    /// Returns the local adapter. Used to enable default behaviors.
368    fn get_local(&self) -> &CoreLocalAdapter<E>;
369
370    //TODO: implement
371    // fn server_side_emit(&self, packet: Packet, opts: BroadcastOptions) -> Result<u64, Error>;
372    // fn persist_session(&self, sid: i64);
373    // fn restore_session(&self, sid: i64) -> Session;
374}
375
376/// The default adapter. Store the state in memory.
377pub struct CoreLocalAdapter<E> {
378    rooms: RwLock<HashMap<Room, HashSet<Sid>>>,
379    sockets: RwLock<HashMap<Sid, HashSet<Room>>>,
380    emitter: E,
381}
382
383impl<E: SocketEmitter> CoreLocalAdapter<E> {
384    /// Create a new local adapter with the given sockets interface.
385    pub fn new(emitter: E) -> Self {
386        Self {
387            rooms: RwLock::new(HashMap::new()),
388            sockets: RwLock::new(HashMap::new()),
389            emitter,
390        }
391    }
392
393    /// Clears all the rooms and sockets.
394    pub fn close(&self) {
395        let mut rooms = self.rooms.write().unwrap();
396        rooms.clear();
397        rooms.shrink_to_fit();
398    }
399
400    /// Adds the socket to all the rooms.
401    pub fn add_all(&self, sid: Sid, rooms: impl RoomParam) {
402        let mut rooms_map = self.rooms.write().unwrap();
403        let mut socket_map = self.sockets.write().unwrap();
404        for room in rooms.into_room_iter() {
405            rooms_map.entry(room.clone()).or_default().insert(sid);
406            socket_map.entry(sid).or_default().insert(room);
407        }
408    }
409
410    /// Removes the socket from the rooms.
411    pub fn del(&self, sid: Sid, rooms: impl RoomParam) {
412        let mut rooms_map = self.rooms.write().unwrap();
413        let mut socket_map = self.sockets.write().unwrap();
414        for room in rooms.into_room_iter() {
415            remove_and_clean_entry(rooms_map.entry(room.clone()), &sid, || {
416                socket_map.entry(sid).and_modify(|r| {
417                    r.remove(&room);
418                });
419            });
420        }
421    }
422
423    /// Removes the socket from all the rooms.
424    pub fn del_all(&self, sid: Sid) {
425        let mut rooms_map = self.rooms.write().unwrap();
426        if let Some(rooms) = self.sockets.write().unwrap().remove(&sid) {
427            for room in rooms {
428                remove_and_clean_entry(rooms_map.entry(room.clone()), &sid, || ());
429            }
430        }
431    }
432
433    /// Broadcasts the packet to the sockets that match the [`BroadcastOptions`].
434    pub fn broadcast(
435        &self,
436        packet: Packet,
437        opts: BroadcastOptions,
438    ) -> Result<(), Vec<SocketError>> {
439        let room_map = self.rooms.read().unwrap();
440        let sids = self.apply_opts(&opts, &room_map);
441
442        if sids.is_empty() {
443            return Ok(());
444        }
445
446        let is_volatile = opts.has_flag(BroadcastFlags::Volatile);
447        let data = self.emitter.parser().encode(packet);
448        if is_volatile {
449            self.emitter.send_many_volatile(sids, data);
450            Ok(())
451        } else {
452            self.emitter.send_many(sids, data)
453        }
454    }
455
456    /// Broadcasts the packet to the sockets that match the [`BroadcastOptions`] and return a stream of ack responses.
457    /// Also returns the number of local expected aknowledgements to know when to stop waiting.
458    pub fn broadcast_with_ack(
459        &self,
460        packet: Packet,
461        opts: BroadcastOptions,
462        timeout: Option<Duration>,
463    ) -> (E::AckStream, u32) {
464        let room_map = self.rooms.read().unwrap();
465        let sids = self.apply_opts(&opts, &room_map);
466        // We cannot pre-serialize the packet because we need to change the ack id.
467        self.emitter.send_many_with_ack(sids, packet, timeout)
468    }
469
470    /// Returns the sockets ids that match the [`BroadcastOptions`].
471    pub fn sockets(&self, opts: BroadcastOptions) -> Vec<Sid> {
472        self.apply_opts(&opts, &self.rooms.read().unwrap())
473            .collect()
474    }
475
476    /// Returns the sockets ids that match the [`BroadcastOptions`].
477    pub fn fetch_sockets(&self, opts: BroadcastOptions) -> Vec<RemoteSocketData> {
478        let rooms = self.rooms.read().unwrap();
479        let sids = self.apply_opts(&opts, &rooms);
480        self.emitter.get_remote_sockets(sids)
481    }
482
483    /// Returns the rooms of the socket.
484    pub fn socket_rooms(&self, sid: Sid) -> HashSet<Room> {
485        self.sockets
486            .read()
487            .unwrap()
488            .get(&sid)
489            .cloned()
490            .unwrap_or_default()
491    }
492
493    /// Adds the sockets that match the [`BroadcastOptions`] to the rooms.
494    pub fn add_sockets(&self, opts: BroadcastOptions, rooms: impl RoomParam) {
495        let rooms: Vec<Room> = rooms.into_room_iter().collect();
496        let mut room_map = self.rooms.write().unwrap();
497        let mut socket_map = self.sockets.write().unwrap();
498        // Here we have to collect sids, because we are going to modify the rooms map.
499        let sids = self.apply_opts(&opts, &room_map).collect::<Vec<_>>();
500        for sid in &sids {
501            let entry = socket_map.entry(*sid).or_default();
502            for room in &rooms {
503                entry.insert(room.clone());
504            }
505        }
506        for room in rooms {
507            let entry = room_map.entry(room).or_default();
508            for sid in &sids {
509                entry.insert(*sid);
510            }
511        }
512    }
513
514    /// Removes the sockets that match the [`BroadcastOptions`] from the rooms.
515    pub fn del_sockets(&self, opts: BroadcastOptions, rooms: impl RoomParam) {
516        let rooms: Vec<Room> = rooms.into_room_iter().collect();
517        let mut rooms_map = self.rooms.write().unwrap();
518        let mut socket_map = self.sockets.write().unwrap();
519        let sids = self.apply_opts(&opts, &rooms_map).collect::<Vec<_>>();
520        for room in rooms {
521            for sid in &sids {
522                remove_and_clean_entry(socket_map.entry(*sid), &room, || ());
523                remove_and_clean_entry(rooms_map.entry(room.clone()), sid, || ());
524            }
525        }
526    }
527
528    /// Disconnects the sockets that match the [`BroadcastOptions`].
529    pub fn disconnect_socket(&self, opts: BroadcastOptions) -> Result<(), Vec<SocketError>> {
530        let sids = self
531            .apply_opts(&opts, &self.rooms.read().unwrap())
532            .collect();
533        self.emitter.disconnect_many(sids)
534    }
535
536    /// Returns all the matching rooms
537    pub fn rooms(&self, opts: BroadcastOptions) -> HashSet<Room> {
538        let rooms = self.rooms.read().unwrap();
539        let sockets = self.sockets.read().unwrap();
540        let sids = self.apply_opts(&opts, &rooms);
541        sids.filter_map(|id| sockets.get(&id))
542            .flatten()
543            .cloned()
544            .collect()
545    }
546
547    /// Get the namespace path.
548    pub fn path(&self) -> &Str {
549        self.emitter.path()
550    }
551
552    /// Get the parser of the namespace.
553    pub fn parser(&self) -> impl Parse + '_ {
554        self.emitter.parser()
555    }
556    /// Get the unique server identifier
557    pub fn server_id(&self) -> Uid {
558        self.emitter.server_id()
559    }
560    /// Get the default configured ack timeout.
561    pub fn ack_timeout(&self) -> Duration {
562        self.emitter.ack_timeout()
563    }
564}
565
566/// The default broadcast iterator.
567/// Extract, flatten and filter a list of sid from a room list
568struct BroadcastRooms<'a> {
569    rooms: slice::Iter<'a, Room>,
570    rooms_map: &'a HashMap<Room, HashSet<Sid>>,
571    except: HashSet<Sid>,
572    flatten_iter: Option<hash_set::Iter<'a, Sid>>,
573}
574impl<'a> BroadcastRooms<'a> {
575    fn new(
576        rooms: &'a [Room],
577        rooms_map: &'a HashMap<Room, HashSet<Sid>>,
578        except: HashSet<Sid>,
579    ) -> Self {
580        BroadcastRooms {
581            rooms: rooms.iter(),
582            rooms_map,
583            except,
584            flatten_iter: None,
585        }
586    }
587}
588impl Iterator for BroadcastRooms<'_> {
589    type Item = Sid;
590    fn next(&mut self) -> Option<Self::Item> {
591        loop {
592            match self.flatten_iter.as_mut().and_then(Iterator::next) {
593                Some(sid) if !self.except.contains(sid) => return Some(*sid),
594                Some(_) => continue,
595                None => self.flatten_iter = None,
596            }
597
598            let room = self.rooms.next()?;
599            self.flatten_iter = self.rooms_map.get(room).map(HashSet::iter);
600        }
601    }
602}
603
604impl<E: SocketEmitter> CoreLocalAdapter<E> {
605    /// Applies the given `opts` and return the sockets that match.
606    fn apply_opts<'a>(
607        &self,
608        opts: &'a BroadcastOptions,
609        rooms: &'a HashMap<Room, HashSet<Sid>>,
610    ) -> BroadcastIter<'a> {
611        let is_broadcast = opts.has_flag(BroadcastFlags::Broadcast);
612
613        let mut except = get_except_sids(&opts.except, rooms);
614        // In case of broadcast flag + if the sender is set,
615        // we should not broadcast to it.
616        if is_broadcast {
617            //FIXME(1.88): switch to if let chains when available
618            if let Some(sid) = opts.sid {
619                except.insert(sid);
620            }
621        }
622
623        if !opts.rooms.is_empty() {
624            let iter = BroadcastRooms::new(&opts.rooms, rooms, except);
625            InnerBroadcastIter::BroadcastRooms(iter).into()
626        } else if is_broadcast {
627            let sids = self.emitter.get_all_sids(|id| !except.contains(id));
628            InnerBroadcastIter::GlobalBroadcast(sids.into_iter()).into()
629        } else if let Some(id) = opts.sid {
630            InnerBroadcastIter::Single(id).into()
631        } else {
632            InnerBroadcastIter::None.into()
633        }
634    }
635}
636
637#[inline]
638fn get_except_sids(except: &[Room], rooms: &HashMap<Room, HashSet<Sid>>) -> HashSet<Sid> {
639    let mut except_sids = HashSet::new();
640    for room in except {
641        if let Some(sockets) = rooms.get(room) {
642            except_sids.extend(sockets);
643        }
644    }
645    except_sids
646}
647
648/// Remove a field from a HashSet value and remove it if empty.
649/// Call `cleanup` fn if the entry exists
650#[inline]
651fn remove_and_clean_entry<K, T: Hash + Eq>(
652    entry: hash_map::Entry<'_, K, HashSet<T>>,
653    el: &T,
654    cleanup: impl FnOnce(),
655) {
656    //TODO: use hashmap raw entry when stabilized to avoid entry clone.
657    // https://github.com/rust-lang/rust/issues/56167
658    match entry {
659        hash_map::Entry::Occupied(mut entry) => {
660            entry.get_mut().remove(el);
661            if entry.get().is_empty() {
662                entry.remove_entry();
663            }
664            cleanup();
665        }
666        hash_map::Entry::Vacant(_) => (),
667    }
668}
669
670/// An iterator that yields the socket ids that match the broadcast options.
671/// Used with the [`SocketEmitter`] interface.
672pub struct BroadcastIter<'a> {
673    inner: InnerBroadcastIter<'a>,
674}
675enum InnerBroadcastIter<'a> {
676    BroadcastRooms(BroadcastRooms<'a>),
677    GlobalBroadcast(<Vec<Sid> as IntoIterator>::IntoIter),
678    Single(Sid),
679    None,
680}
681impl BroadcastIter<'_> {
682    fn is_empty(&self) -> bool {
683        matches!(self.inner, InnerBroadcastIter::None)
684    }
685}
686impl<'a> From<InnerBroadcastIter<'a>> for BroadcastIter<'a> {
687    fn from(inner: InnerBroadcastIter<'a>) -> Self {
688        BroadcastIter { inner }
689    }
690}
691
692impl Iterator for BroadcastIter<'_> {
693    type Item = Sid;
694
695    #[inline(always)]
696    fn next(&mut self) -> Option<Self::Item> {
697        self.inner.next()
698    }
699}
700impl Iterator for InnerBroadcastIter<'_> {
701    type Item = Sid;
702
703    fn next(&mut self) -> Option<Self::Item> {
704        match self {
705            InnerBroadcastIter::BroadcastRooms(inner) => inner.next(),
706            InnerBroadcastIter::GlobalBroadcast(inner) => inner.next(),
707            InnerBroadcastIter::Single(sid) => {
708                let sid = *sid;
709                *self = InnerBroadcastIter::None;
710                Some(sid)
711            }
712            InnerBroadcastIter::None => None,
713        }
714    }
715}
716
717/// Represent the data of a remote socket.
718#[derive(Debug, Serialize, Deserialize, PartialEq, Eq, Default, Clone)]
719pub struct RemoteSocketData {
720    /// The id of the remote socket.
721    pub id: Sid,
722    /// The server id this socket is connected to.
723    pub server_id: Uid,
724    /// The namespace this socket is connected to.
725    pub ns: Str,
726}
727
728#[cfg(test)]
729mod test {
730
731    use smallvec::smallvec;
732    use std::{
733        array,
734        pin::Pin,
735        task::{Context, Poll},
736    };
737
738    use super::*;
739
740    struct StubSockets {
741        sockets: HashSet<Sid>,
742        path: Str,
743    }
744    impl StubSockets {
745        fn new(sockets: &[Sid]) -> Self {
746            let sockets = HashSet::from_iter(sockets.iter().copied());
747            Self {
748                sockets,
749                path: Str::from("/"),
750            }
751        }
752    }
753
754    struct StubAckStream;
755    impl Stream for StubAckStream {
756        type Item = (Sid, Result<Value, StubError>);
757        fn poll_next(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<Option<Self::Item>> {
758            Poll::Ready(None)
759        }
760    }
761    impl FusedStream for StubAckStream {
762        fn is_terminated(&self) -> bool {
763            true
764        }
765    }
766    #[derive(Debug, Serialize, Deserialize)]
767    struct StubError;
768    impl std::fmt::Display for StubError {
769        fn fmt(&self, _: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
770            Ok(())
771        }
772    }
773    impl std::error::Error for StubError {}
774
775    impl SocketEmitter for StubSockets {
776        type AckError = StubError;
777        type AckStream = StubAckStream;
778        fn get_all_sids(&self, filter: impl Fn(&Sid) -> bool) -> Vec<Sid> {
779            self.sockets
780                .iter()
781                .copied()
782                .filter(|id| filter(id))
783                .collect()
784        }
785
786        fn get_remote_sockets(&self, sids: BroadcastIter<'_>) -> Vec<RemoteSocketData> {
787            sids.map(|id| RemoteSocketData {
788                id,
789                server_id: Uid::ZERO,
790                ns: self.path.clone(),
791            })
792            .collect()
793        }
794
795        fn send_many(&self, _: BroadcastIter<'_>, _: Value) -> Result<(), Vec<SocketError>> {
796            Ok(())
797        }
798
799        fn send_many_volatile(&self, _: BroadcastIter<'_>, _: Value) {}
800
801        fn send_many_with_ack(
802            &self,
803            _: BroadcastIter<'_>,
804            _: Packet,
805            _: Option<Duration>,
806        ) -> (Self::AckStream, u32) {
807            (StubAckStream, 0)
808        }
809
810        fn disconnect_many(&self, _: Vec<Sid>) -> Result<(), Vec<SocketError>> {
811            Ok(())
812        }
813
814        fn path(&self) -> &Str {
815            &self.path
816        }
817        fn parser(&self) -> impl Parse {
818            crate::parser::test::StubParser
819        }
820        fn server_id(&self) -> Uid {
821            Uid::ZERO
822        }
823        fn ack_timeout(&self) -> Duration {
824            Duration::ZERO
825        }
826    }
827
828    fn create_adapter<const S: usize>(sockets: [Sid; S]) -> CoreLocalAdapter<StubSockets> {
829        CoreLocalAdapter::new(StubSockets::new(&sockets))
830    }
831
832    #[test]
833    fn add_all() {
834        let socket = Sid::new();
835        let adapter = create_adapter([socket]);
836        adapter.add_all(socket, ["room1", "room2"]);
837        let rooms_map = adapter.rooms.read().unwrap();
838        let socket_map = adapter.sockets.read().unwrap();
839        assert_eq!(rooms_map.len(), 2);
840        assert_eq!(socket_map.len(), 1);
841        assert_eq!(rooms_map.get("room1").unwrap().len(), 1);
842        assert_eq!(rooms_map.get("room2").unwrap().len(), 1);
843
844        let rooms = socket_map.get(&socket).unwrap();
845        assert!(rooms.contains("room1"));
846        assert!(rooms.contains("room2"));
847    }
848
849    #[test]
850    fn del() {
851        let socket = Sid::new();
852        let adapter = create_adapter([socket]);
853        adapter.add_all(socket, ["room1", "room2"]);
854        {
855            let rooms_map = adapter.rooms.read().unwrap();
856            assert_eq!(rooms_map.len(), 2);
857            assert_eq!(rooms_map.get("room1").unwrap().len(), 1);
858            assert_eq!(rooms_map.get("room2").unwrap().len(), 1);
859            let socket_map = adapter.sockets.read().unwrap();
860            let rooms = socket_map.get(&socket).unwrap();
861            assert!(rooms.contains("room1"));
862            assert!(rooms.contains("room2"));
863        }
864        adapter.del(socket, "room1");
865        let rooms_map = adapter.rooms.read().unwrap();
866        let socket_map = adapter.sockets.read().unwrap();
867        assert_eq!(rooms_map.len(), 1);
868        assert!(rooms_map.get("room1").is_none());
869        assert_eq!(rooms_map.get("room2").unwrap().len(), 1);
870        assert_eq!(socket_map.get(&socket).unwrap().len(), 1);
871    }
872    #[test]
873    fn del_all() {
874        let socket = Sid::new();
875        let adapter = create_adapter([socket]);
876        adapter.add_all(socket, ["room1", "room2"]);
877
878        {
879            let rooms_map = adapter.rooms.read().unwrap();
880            assert_eq!(rooms_map.len(), 2);
881            assert_eq!(rooms_map.get("room1").unwrap().len(), 1);
882            assert_eq!(rooms_map.get("room2").unwrap().len(), 1);
883
884            let socket_map = adapter.sockets.read().unwrap();
885            let rooms = socket_map.get(&socket).unwrap();
886            assert!(rooms.contains("room1"));
887            assert!(rooms.contains("room2"));
888        }
889
890        adapter.del_all(socket);
891
892        {
893            let rooms_map = adapter.rooms.read().unwrap();
894            assert_eq!(rooms_map.len(), 0);
895
896            let socket_map = adapter.sockets.read().unwrap();
897            assert!(socket_map.get(&socket).is_none());
898        }
899    }
900
901    #[test]
902    fn socket_room() {
903        let sid1 = Sid::new();
904        let sid2 = Sid::new();
905        let sid3 = Sid::new();
906        let adapter = create_adapter([sid1, sid2, sid3]);
907        adapter.add_all(sid1, ["room1", "room2"]);
908        adapter.add_all(sid2, ["room1"]);
909        adapter.add_all(sid3, ["room2"]);
910        assert!(adapter.socket_rooms(sid1).contains(&Cow::Borrowed("room1")));
911        assert!(adapter.socket_rooms(sid1).contains(&Cow::Borrowed("room2")));
912        assert_eq!(
913            adapter.socket_rooms(sid2).into_iter().collect::<Vec<_>>(),
914            ["room1"]
915        );
916        assert_eq!(
917            adapter.socket_rooms(sid3).into_iter().collect::<Vec<_>>(),
918            ["room2"]
919        );
920    }
921
922    #[test]
923    fn add_socket() {
924        let socket = Sid::new();
925        let adapter = create_adapter([socket]);
926        adapter.add_all(socket, ["room1"]);
927
928        let mut opts = BroadcastOptions::new(socket);
929        opts.rooms = smallvec!["room1".into()];
930        adapter.add_sockets(opts, "room2");
931        let rooms_map = adapter.rooms.read().unwrap();
932
933        assert_eq!(rooms_map.len(), 2);
934        assert!(rooms_map.get("room1").unwrap().contains(&socket));
935        assert!(rooms_map.get("room2").unwrap().contains(&socket));
936    }
937
938    #[test]
939    fn del_socket() {
940        let socket = Sid::new();
941        let adapter = create_adapter([socket]);
942        adapter.add_all(socket, ["room1"]);
943
944        let mut opts = BroadcastOptions::new(socket);
945        opts.rooms = smallvec!["room1".into()];
946        adapter.add_sockets(opts, "room2");
947
948        {
949            let rooms_map = adapter.rooms.read().unwrap();
950
951            assert_eq!(rooms_map.len(), 2);
952            assert!(rooms_map.get("room1").unwrap().contains(&socket));
953            assert!(rooms_map.get("room2").unwrap().contains(&socket));
954        }
955
956        let mut opts = BroadcastOptions::new(socket);
957        opts.rooms = smallvec!["room1".into()];
958        adapter.del_sockets(opts, "room2");
959
960        {
961            let rooms_map = adapter.rooms.read().unwrap();
962
963            assert_eq!(rooms_map.len(), 1);
964            assert!(rooms_map.get("room1").unwrap().contains(&socket));
965            assert!(rooms_map.get("room2").is_none());
966        }
967    }
968
969    #[test]
970    fn sockets() {
971        let socket0 = Sid::new();
972        let socket1 = Sid::new();
973        let socket2 = Sid::new();
974        let adapter = create_adapter([socket0, socket1, socket2]);
975        adapter.add_all(socket0, ["room1", "room2"]);
976        adapter.add_all(socket1, ["room1", "room3"]);
977        adapter.add_all(socket2, ["room2", "room3"]);
978
979        let mut opts = BroadcastOptions {
980            rooms: smallvec!["room1".into()],
981            ..Default::default()
982        };
983        let sockets = adapter.sockets(opts.clone());
984        assert_eq!(sockets.len(), 2);
985        assert!(sockets.contains(&socket0));
986        assert!(sockets.contains(&socket1));
987
988        opts.rooms = smallvec!["room2".into()];
989        let sockets = adapter.sockets(opts.clone());
990        assert_eq!(sockets.len(), 2);
991        assert!(sockets.contains(&socket0));
992        assert!(sockets.contains(&socket2));
993
994        opts.rooms = smallvec!["room3".into()];
995        let sockets = adapter.sockets(opts.clone());
996        assert_eq!(sockets.len(), 2);
997        assert!(sockets.contains(&socket1));
998        assert!(sockets.contains(&socket2));
999    }
1000
1001    #[test]
1002    fn disconnect_socket() {
1003        let socket0 = Sid::new();
1004        let socket1 = Sid::new();
1005        let socket2 = Sid::new();
1006        let adapter = create_adapter([socket0, socket1, socket2]);
1007        adapter.add_all(socket0, ["room1", "room2", "room4"]);
1008        adapter.add_all(socket1, ["room1", "room3", "room5"]);
1009        adapter.add_all(socket2, ["room2", "room3", "room6"]);
1010
1011        let mut opts = BroadcastOptions::new(socket0);
1012        opts.rooms = smallvec!["room5".into()];
1013        adapter.disconnect_socket(opts).unwrap();
1014
1015        let mut opts = BroadcastOptions::default();
1016        opts.rooms.push("room2".into());
1017        let sockets = adapter.sockets(opts.clone());
1018        assert_eq!(sockets.len(), 2);
1019        assert!(sockets.contains(&socket2));
1020        assert!(sockets.contains(&socket0));
1021    }
1022    #[test]
1023    fn disconnect_empty_opts() {
1024        let adapter = create_adapter([]);
1025        let opts = BroadcastOptions::default();
1026        adapter.disconnect_socket(opts).unwrap();
1027    }
1028    #[test]
1029    fn rooms() {
1030        let socket0 = Sid::new();
1031        let socket1 = Sid::new();
1032        let socket2 = Sid::new();
1033        let adapter = create_adapter([socket0, socket1, socket2]);
1034        adapter.add_all(socket0, ["room1", "room2", "room4"]);
1035        adapter.add_all(socket1, ["room1", "room3", "room5"]);
1036        adapter.add_all(socket2, ["room2", "room3", "room6"]);
1037
1038        let mut opts = BroadcastOptions::new(socket0);
1039        opts.rooms = smallvec!["room5".into()];
1040        opts.add_flag(BroadcastFlags::Broadcast);
1041        let rooms = adapter.rooms(opts);
1042        assert_eq!(rooms.len(), 3);
1043        assert!(rooms.contains(&Cow::Borrowed("room1")));
1044        assert!(rooms.contains(&Cow::Borrowed("room3")));
1045        assert!(rooms.contains(&Cow::Borrowed("room5")));
1046
1047        let mut opts = BroadcastOptions::default();
1048        opts.rooms.push("room2".into());
1049        let rooms = adapter.rooms(opts.clone());
1050        assert_eq!(rooms.len(), 5);
1051        assert!(rooms.contains(&Cow::Borrowed("room1")));
1052        assert!(rooms.contains(&Cow::Borrowed("room2")));
1053        assert!(rooms.contains(&Cow::Borrowed("room3")));
1054        assert!(rooms.contains(&Cow::Borrowed("room4")));
1055        assert!(rooms.contains(&Cow::Borrowed("room6")));
1056    }
1057
1058    #[test]
1059    fn apply_opts() {
1060        let mut sockets: [Sid; 3] = array::from_fn(|_| Sid::new());
1061        sockets.sort();
1062        let adapter = create_adapter(sockets);
1063
1064        adapter.add_all(sockets[0], ["room1", "room2"]);
1065        adapter.add_all(sockets[1], ["room1", "room3"]);
1066        adapter.add_all(sockets[2], ["room1", "room2", "room3"]);
1067
1068        // socket 2 is the sender
1069        let mut opts = BroadcastOptions::new(sockets[2]);
1070        opts.rooms = smallvec!["room1".into()];
1071        opts.except = smallvec!["room2".into()];
1072        let sids = adapter
1073            .apply_opts(&opts, &adapter.rooms.read().unwrap())
1074            .collect::<Vec<_>>();
1075        assert_eq!(sids, [sockets[1]]);
1076
1077        let mut opts = BroadcastOptions::new(sockets[2]);
1078        opts.add_flag(BroadcastFlags::Broadcast);
1079        let mut sids = adapter
1080            .apply_opts(&opts, &adapter.rooms.read().unwrap())
1081            .collect::<Vec<_>>();
1082        sids.sort();
1083        assert_eq!(sids, [sockets[0], sockets[1]]);
1084
1085        let mut opts = BroadcastOptions::new(sockets[2]);
1086        opts.add_flag(BroadcastFlags::Broadcast);
1087        opts.except = smallvec!["room2".into()];
1088        let sids = adapter
1089            .apply_opts(&opts, &adapter.rooms.read().unwrap())
1090            .collect::<Vec<_>>();
1091        assert_eq!(sids.len(), 1);
1092
1093        let opts = BroadcastOptions::new(sockets[2]);
1094        let sids = adapter
1095            .apply_opts(&opts, &adapter.rooms.read().unwrap())
1096            .collect::<Vec<_>>();
1097        assert_eq!(sids.len(), 1);
1098        assert_eq!(sids[0], sockets[2]);
1099
1100        let opts = BroadcastOptions::new(Sid::new());
1101        let sids = adapter
1102            .apply_opts(&opts, &adapter.rooms.read().unwrap())
1103            .collect::<Vec<_>>();
1104        assert_eq!(sids.len(), 1);
1105    }
1106
1107    #[test]
1108    fn test_is_local_opts() {
1109        let server_id = Uid::new();
1110        let remote = RemoteSocketData {
1111            id: Sid::new(),
1112            server_id,
1113            ns: "/".into(),
1114        };
1115        let opts = BroadcastOptions::new_remote(&remote);
1116        assert!(opts.is_local(server_id));
1117        assert!(!opts.is_local(Uid::new()));
1118        let opts = BroadcastOptions::new(Sid::new());
1119        assert!(!opts.is_local(Uid::new()));
1120    }
1121}