use super::service_info::Subscriber;
use core::future::Future;
use core::net::SocketAddrV4;
use heapless::{Vec as HeaplessVec, index_map::FnvIndexMap};
#[cfg(feature = "server-tokio")]
use std::sync::Arc;
#[cfg(feature = "server-tokio")]
use tokio::sync::RwLock;
#[cfg(feature = "bare_metal")]
const DEFAULT_EVENT_GROUPS: usize = 4;
#[cfg(not(feature = "bare_metal"))]
const DEFAULT_EVENT_GROUPS: usize = 32;
#[cfg(feature = "bare_metal")]
const DEFAULT_SUBSCRIBERS: usize = 1;
#[cfg(not(feature = "bare_metal"))]
const DEFAULT_SUBSCRIBERS: usize = 16;
const EVENT_GROUPS_CAP: usize = crate::from_env_or(
option_env!("SIMPLE_SOMEIP_MAX_OFFERS"),
DEFAULT_EVENT_GROUPS,
)
.next_power_of_two();
pub(crate) const SUBSCRIBERS_PER_GROUP: usize =
crate::from_env_or(option_env!("SIMPLE_SOMEIP_MAX_SUBS"), DEFAULT_SUBSCRIBERS);
const _: () = assert!(
SUBSCRIBERS_PER_GROUP >= 1,
"SUBSCRIBERS_PER_GROUP must be >= 1: a value of 0 would crash subscribe() on first push"
);
const _: () = assert!(
EVENT_GROUPS_CAP.is_power_of_two(),
"EVENT_GROUPS_CAP must be a power of two for heapless::FnvIndexMap"
);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SubscribeError {
SubscribersPerGroupFull,
EventGroupsFull,
}
impl core::fmt::Display for SubscribeError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::SubscribersPerGroupFull => write!(
f,
"subscribers-per-group at capacity ({SUBSCRIBERS_PER_GROUP})"
),
Self::EventGroupsFull => {
write!(f, "event-group map at capacity ({EVENT_GROUPS_CAP})")
}
}
}
}
type SubscribersList = HeaplessVec<Subscriber, SUBSCRIBERS_PER_GROUP>;
#[derive(Debug)]
pub struct SubscriptionManager {
subscriptions: FnvIndexMap<(u16, u16, u16), SubscribersList, EVENT_GROUPS_CAP>,
}
impl SubscriptionManager {
#[must_use]
pub const fn new() -> Self {
Self {
subscriptions: FnvIndexMap::new(),
}
}
pub fn subscribe(
&mut self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Result<(), SubscribeError> {
let key = (service_id, instance_id, event_group_id);
if let Some(subscribers) = self.subscriptions.get_mut(&key) {
if subscribers.iter().any(|s| s.address == subscriber_addr) {
crate::log::debug!(
"Subscriber {} already subscribed for service 0x{:04X}, instance {}, \
event group 0x{:04X}; skipping duplicate",
subscriber_addr,
service_id,
instance_id,
event_group_id
);
return Ok(());
}
let subscriber =
Subscriber::new(subscriber_addr, service_id, instance_id, event_group_id);
if subscribers.push(subscriber).is_err() {
crate::log::warn!(
"Subscribers-per-group at capacity ({}); dropping new subscriber {} \
for service 0x{:04X}, instance {}, event group 0x{:04X}",
SUBSCRIBERS_PER_GROUP,
subscriber_addr,
service_id,
instance_id,
event_group_id
);
return Err(SubscribeError::SubscribersPerGroupFull);
}
crate::log::info!(
"Subscriber {} added for service 0x{:04X}, instance {}, event group 0x{:04X}",
subscriber_addr,
service_id,
instance_id,
event_group_id
);
return Ok(());
}
let mut list = SubscribersList::new();
list.push(Subscriber::new(
subscriber_addr,
service_id,
instance_id,
event_group_id,
))
.expect(
"new SubscribersList must accept the first subscriber; \
SUBSCRIBERS_PER_GROUP must be >= 1",
);
if self.subscriptions.insert(key, list).is_err() {
crate::log::warn!(
"Event-group map at capacity ({}); dropping subscriber {} for new group \
service 0x{:04X}, instance {}, event group 0x{:04X}",
EVENT_GROUPS_CAP,
subscriber_addr,
service_id,
instance_id,
event_group_id
);
return Err(SubscribeError::EventGroupsFull);
}
crate::log::info!(
"Subscriber {} added for service 0x{:04X}, instance {}, event group 0x{:04X}",
subscriber_addr,
service_id,
instance_id,
event_group_id
);
Ok(())
}
pub fn unsubscribe(
&mut self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) {
let key = (service_id, instance_id, event_group_id);
if let Some(subscribers) = self.subscriptions.get_mut(&key) {
subscribers.retain(|s| s.address != subscriber_addr);
if subscribers.is_empty() {
self.subscriptions.remove(&key);
}
crate::log::info!(
"Removed subscriber {} from service 0x{:04X}, instance {}, event group 0x{:04X}",
subscriber_addr,
service_id,
instance_id,
event_group_id
);
}
}
#[cfg(feature = "_alloc")]
#[must_use]
pub fn get_subscribers(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
) -> alloc::vec::Vec<Subscriber> {
let key = (service_id, instance_id, event_group_id);
self.subscriptions
.get(&key)
.map(|list| list.iter().cloned().collect())
.unwrap_or_default()
}
#[must_use]
pub fn subscription_count(&self) -> usize {
self.subscriptions.values().map(|v| v.len()).sum()
}
}
impl Default for SubscriptionManager {
fn default() -> Self {
Self::new()
}
}
pub trait SubscriptionHandle: Clone + 'static {
type SubscribeFuture<'a>: Future<Output = Result<(), SubscribeError>> + 'a
where
Self: 'a;
type UnsubscribeFuture<'a>: Future<Output = ()> + 'a
where
Self: 'a;
type ForEachFuture<'a>: Future<Output = usize> + 'a
where
Self: 'a;
fn subscribe(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Self::SubscribeFuture<'_>;
fn unsubscribe(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Self::UnsubscribeFuture<'_>;
fn for_each_subscriber<'a>(
&'a self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
f: &'a mut (dyn FnMut(&Subscriber) + Send),
) -> Self::ForEachFuture<'a>;
}
#[cfg(feature = "server-tokio")]
impl SubscriptionHandle for Arc<RwLock<SubscriptionManager>> {
type SubscribeFuture<'a> = core::pin::Pin<
alloc::boxed::Box<dyn Future<Output = Result<(), SubscribeError>> + Send + 'a>,
>;
type UnsubscribeFuture<'a> =
core::pin::Pin<alloc::boxed::Box<dyn Future<Output = ()> + Send + 'a>>;
type ForEachFuture<'a> =
core::pin::Pin<alloc::boxed::Box<dyn Future<Output = usize> + Send + 'a>>;
fn subscribe(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Self::SubscribeFuture<'_> {
let this = self.clone();
alloc::boxed::Box::pin(async move {
this.write()
.await
.subscribe(service_id, instance_id, event_group_id, subscriber_addr)
})
}
fn unsubscribe(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Self::UnsubscribeFuture<'_> {
let this = self.clone();
alloc::boxed::Box::pin(async move {
this.write().await.unsubscribe(
service_id,
instance_id,
event_group_id,
subscriber_addr,
);
})
}
fn for_each_subscriber<'a>(
&'a self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
f: &'a mut (dyn FnMut(&Subscriber) + Send),
) -> Self::ForEachFuture<'a> {
let this = self.clone();
alloc::boxed::Box::pin(async move {
let guard = this.read().await;
let key = (service_id, instance_id, event_group_id);
match guard.subscriptions.get(&key) {
Some(list) => {
for sub in list {
f(sub);
}
list.len()
}
None => 0,
}
})
}
}
#[cfg(feature = "bare_metal")]
pub mod bare_metal_subscription_impl {
use super::{SubscribeError, Subscriber, SubscriptionHandle, SubscriptionManager};
use core::cell::RefCell;
use core::net::SocketAddrV4;
use embassy_sync::blocking_mutex::Mutex;
use embassy_sync::blocking_mutex::raw::CriticalSectionRawMutex;
pub type StaticSubscriptionStorage =
Mutex<CriticalSectionRawMutex, RefCell<SubscriptionManager>>;
#[derive(Clone, Copy)]
pub struct StaticSubscriptionHandle(&'static StaticSubscriptionStorage);
impl StaticSubscriptionHandle {
#[must_use]
pub const fn new(storage: &'static StaticSubscriptionStorage) -> Self {
Self(storage)
}
}
impl SubscriptionHandle for StaticSubscriptionHandle {
type SubscribeFuture<'a> = core::future::Ready<Result<(), SubscribeError>>;
type UnsubscribeFuture<'a> = core::future::Ready<()>;
type ForEachFuture<'a> = core::future::Ready<usize>;
fn subscribe(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Self::SubscribeFuture<'_> {
let storage = self.0;
core::future::ready(storage.lock(|cell| {
cell.borrow_mut().subscribe(
service_id,
instance_id,
event_group_id,
subscriber_addr,
)
}))
}
fn unsubscribe(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: SocketAddrV4,
) -> Self::UnsubscribeFuture<'_> {
let storage = self.0;
storage.lock(|cell| {
cell.borrow_mut().unsubscribe(
service_id,
instance_id,
event_group_id,
subscriber_addr,
);
});
core::future::ready(())
}
fn for_each_subscriber<'a>(
&'a self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
f: &'a mut (dyn FnMut(&Subscriber) + Send),
) -> Self::ForEachFuture<'a> {
let storage = self.0;
core::future::ready(storage.lock(|cell| {
let guard = cell.borrow();
let key = (service_id, instance_id, event_group_id);
match guard.subscriptions.get(&key) {
Some(list) => {
for sub in list {
f(sub);
}
list.len()
}
None => 0,
}
}))
}
}
}
#[cfg(feature = "bare_metal")]
pub use bare_metal_subscription_impl::{StaticSubscriptionHandle, StaticSubscriptionStorage};
#[cfg(test)]
mod tests {
use super::*;
use std::net::Ipv4Addr;
use std::vec::Vec;
#[test]
fn test_subscription_management() {
let mut manager = SubscriptionManager::new();
let addr = SocketAddrV4::new(Ipv4Addr::new(192, 168, 1, 1), 8080);
manager.subscribe(0x5B, 1, 0x01, addr).unwrap();
assert_eq!(manager.subscription_count(), 1);
let subscribers = manager.get_subscribers(0x5B, 1, 0x01);
assert_eq!(subscribers.len(), 1);
assert_eq!(subscribers[0].address, addr);
manager.unsubscribe(0x5B, 1, 0x01, addr);
assert_eq!(manager.subscription_count(), 0);
}
#[test]
fn test_duplicate_subscriber_refresh() {
let mut manager = SubscriptionManager::new();
let addr = SocketAddrV4::new(Ipv4Addr::new(192, 168, 1, 1), 8080);
manager.subscribe(0x5B, 1, 0x01, addr).unwrap();
assert_eq!(manager.subscription_count(), 1);
manager.subscribe(0x5B, 1, 0x01, addr).unwrap();
assert_eq!(manager.subscription_count(), 1);
}
#[test]
fn test_unsubscribe_nonexistent_key() {
let mut manager = SubscriptionManager::new();
let addr = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 9000);
manager.unsubscribe(0x99, 1, 0x01, addr);
assert_eq!(manager.subscription_count(), 0);
}
#[test]
fn test_get_subscribers_empty() {
let manager = SubscriptionManager::new();
assert!(manager.get_subscribers(0x99, 1, 0x01).is_empty());
}
#[test]
fn test_default_impl() {
let manager = SubscriptionManager::default();
assert_eq!(manager.subscription_count(), 0);
}
#[test]
fn subscribers_per_group_capacity_overflow() {
let mut manager = SubscriptionManager::new();
for i in 0..SUBSCRIBERS_PER_GROUP {
let addr =
SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8000 + u16::try_from(i).unwrap());
manager.subscribe(0x5B, 1, 0x01, addr).unwrap();
}
assert_eq!(manager.subscription_count(), SUBSCRIBERS_PER_GROUP);
let extra = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 9999);
assert_eq!(
manager.subscribe(0x5B, 1, 0x01, extra),
Err(SubscribeError::SubscribersPerGroupFull),
);
assert_eq!(manager.subscription_count(), SUBSCRIBERS_PER_GROUP);
let subs = manager.get_subscribers(0x5B, 1, 0x01);
assert!(subs.iter().all(|s| s.address != extra));
}
#[test]
fn event_groups_capacity_overflow() {
let mut manager = SubscriptionManager::new();
let addr = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8000);
for i in 0..EVENT_GROUPS_CAP {
let eg = u16::try_from(i).unwrap();
manager.subscribe(0x5B, 1, eg, addr).unwrap();
}
assert_eq!(manager.subscription_count(), EVENT_GROUPS_CAP);
let overflow_eg = u16::try_from(EVENT_GROUPS_CAP).unwrap();
assert_eq!(
manager.subscribe(0x5B, 1, overflow_eg, addr),
Err(SubscribeError::EventGroupsFull),
);
assert_eq!(manager.subscription_count(), EVENT_GROUPS_CAP);
assert!(manager.get_subscribers(0x5B, 1, overflow_eg).is_empty());
}
#[test]
fn unsubscribe_one_of_multiple_leaves_group_intact() {
let mut manager = SubscriptionManager::new();
let a1 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8001);
let a2 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 8002);
manager.subscribe(0x5B, 1, 0x01, a1).unwrap();
manager.subscribe(0x5B, 1, 0x01, a2).unwrap();
assert_eq!(manager.subscription_count(), 2);
manager.unsubscribe(0x5B, 1, 0x01, a1);
assert_eq!(manager.subscription_count(), 1);
let subs = manager.get_subscribers(0x5B, 1, 0x01);
assert_eq!(subs.len(), 1);
assert_eq!(subs[0].address, a2);
}
#[test]
fn unsubscribe_address_not_in_existing_group_is_noop() {
let mut manager = SubscriptionManager::new();
let a1 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8001);
let a2 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 8002);
manager.subscribe(0x5B, 1, 0x01, a1).unwrap();
manager.unsubscribe(0x5B, 1, 0x01, a2);
assert_eq!(manager.subscription_count(), 1);
assert_eq!(manager.get_subscribers(0x5B, 1, 0x01)[0].address, a1);
}
#[test]
fn get_subscribers_returns_all_in_group() {
let mut manager = SubscriptionManager::new();
let addrs: Vec<SocketAddrV4> = (0..4)
.map(|i| SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, i + 1), 8000 + u16::from(i)))
.collect();
for &a in &addrs {
manager.subscribe(0x5B, 1, 0x01, a).unwrap();
}
let subs = manager.get_subscribers(0x5B, 1, 0x01);
assert_eq!(subs.len(), 4);
for &a in &addrs {
assert!(subs.iter().any(|s| s.address == a));
}
}
#[test]
fn subscription_count_spans_multiple_event_groups() {
let mut manager = SubscriptionManager::new();
let a = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8000);
manager.subscribe(0x5B, 1, 0x01, a).unwrap();
manager.subscribe(0x5B, 1, 0x02, a).unwrap();
manager.subscribe(0x5C, 1, 0x01, a).unwrap();
assert_eq!(manager.subscription_count(), 3);
}
#[test]
fn subscribe_error_display() {
use std::string::ToString;
assert!(
SubscribeError::SubscribersPerGroupFull
.to_string()
.contains("subscribers-per-group"),
);
assert!(
SubscribeError::EventGroupsFull
.to_string()
.contains("event-group"),
);
}
#[cfg(feature = "server-tokio")]
mod tokio_handle {
use super::*;
use std::sync::Arc;
use tokio::sync::RwLock;
#[tokio::test]
async fn for_each_subscriber_visits_all() {
let handle: Arc<RwLock<SubscriptionManager>> =
Arc::new(RwLock::new(SubscriptionManager::new()));
let a1 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8001);
let a2 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 8002);
handle.subscribe(0x5B, 1, 0x01, a1).await.unwrap();
handle.subscribe(0x5B, 1, 0x01, a2).await.unwrap();
let mut visited = Vec::new();
let count = {
let mut visit = |s: &Subscriber| visited.push(s.address);
handle.for_each_subscriber(0x5B, 1, 0x01, &mut visit).await
};
assert_eq!(count, 2);
assert!(visited.contains(&a1));
assert!(visited.contains(&a2));
}
#[tokio::test]
async fn for_each_subscriber_empty_group_returns_zero() {
let handle: Arc<RwLock<SubscriptionManager>> =
Arc::new(RwLock::new(SubscriptionManager::new()));
let mut visit = |_: &Subscriber| {};
let count = handle.for_each_subscriber(0x5B, 1, 0x01, &mut visit).await;
assert_eq!(count, 0);
}
#[tokio::test]
async fn for_each_subscriber_reflects_unsubscribe() {
let handle: Arc<RwLock<SubscriptionManager>> =
Arc::new(RwLock::new(SubscriptionManager::new()));
let a1 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8001);
let a2 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 8002);
handle.subscribe(0x5B, 1, 0x01, a1).await.unwrap();
handle.subscribe(0x5B, 1, 0x01, a2).await.unwrap();
handle.unsubscribe(0x5B, 1, 0x01, a1).await;
let mut visited = Vec::new();
let count = {
let mut visit = |s: &Subscriber| visited.push(s.address);
handle.for_each_subscriber(0x5B, 1, 0x01, &mut visit).await
};
assert_eq!(count, 1);
assert_eq!(visited, [a2]);
}
}
#[cfg(feature = "bare_metal")]
mod static_handle {
use super::*;
use crate::server::{StaticSubscriptionHandle, StaticSubscriptionStorage};
use core::cell::RefCell;
use embassy_sync::blocking_mutex::Mutex as BlockingMutex;
use embassy_sync::blocking_mutex::raw::CriticalSectionRawMutex;
fn block_on_sync<F: core::future::Future>(fut: F) -> F::Output {
use core::pin::pin;
use core::task::{Context, Poll, Waker};
let mut fut = pin!(fut);
let waker = Waker::noop();
let mut cx = Context::from_waker(waker);
match fut.as_mut().poll(&mut cx) {
Poll::Ready(v) => v,
Poll::Pending => panic!(
"StaticSubscriptionHandle methods must complete \
synchronously (no .await inside the lock); got Pending"
),
}
}
#[test]
fn static_subscription_handle_full_contract() {
let storage: &'static StaticSubscriptionStorage =
std::boxed::Box::leak(std::boxed::Box::new(BlockingMutex::<
CriticalSectionRawMutex,
RefCell<SubscriptionManager>,
>::new(RefCell::new(
SubscriptionManager::new(),
))));
let handle = StaticSubscriptionHandle::new(storage);
let a1 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 1), 8001);
let a2 = SocketAddrV4::new(Ipv4Addr::new(10, 0, 0, 2), 8002);
block_on_sync(handle.subscribe(0x5B, 1, 0x01, a1)).unwrap();
block_on_sync(handle.subscribe(0x5B, 1, 0x01, a2)).unwrap();
let mut visited: std::vec::Vec<SocketAddrV4> = std::vec::Vec::new();
let count = {
let mut visit = |s: &Subscriber| visited.push(s.address);
block_on_sync(handle.for_each_subscriber(0x5B, 1, 0x01, &mut visit))
};
assert_eq!(count, 2);
assert!(visited.contains(&a1));
assert!(visited.contains(&a2));
block_on_sync(handle.unsubscribe(0x5B, 1, 0x01, a1));
visited.clear();
let count = {
let mut visit = |s: &Subscriber| visited.push(s.address);
block_on_sync(handle.for_each_subscriber(0x5B, 1, 0x01, &mut visit))
};
assert_eq!(count, 1);
assert_eq!(visited, [a2]);
}
}
}