use super::PubSubError;
use super::Result;
const WILDCARD_TOPIC_ID: u16 = 0xFFFF;
type PubSubCallback = fn(topic_id: u16, data: &[u8]) -> bool;
#[derive(Clone)]
struct Subscriber {
id: usize,
topic_id: u16,
callback: PubSubCallback,
active: bool,
last_active: u64,
}
pub struct SubscriberManager {
max_topics: usize,
max_subscribers_per_topic: usize,
subscribers: alloc::vec::Vec<alloc::vec::Vec<Option<Subscriber>>>,
wildcard_subscribers: alloc::vec::Vec<Option<Subscriber>>,
next_subscription_id: usize,
global_subscribers: alloc::vec::Vec<usize>, topic_map: alloc::collections::BTreeMap<&'static str, u16>,
id_map: alloc::vec::Vec<Option<&'static str>>,
expected_seq_nums: alloc::vec::Vec<u32>,
}
impl Clone for SubscriberManager {
fn clone(&self) -> Self {
Self {
max_topics: self.max_topics,
max_subscribers_per_topic: self.max_subscribers_per_topic,
subscribers: self.subscribers.clone(),
wildcard_subscribers: self.wildcard_subscribers.clone(),
next_subscription_id: self.next_subscription_id,
global_subscribers: self.global_subscribers.clone(),
topic_map: self.topic_map.clone(),
id_map: self.id_map.clone(),
expected_seq_nums: self.expected_seq_nums.clone(),
}
}
}
impl SubscriberManager {
pub fn new(max_topics: usize, max_subscribers_per_topic: usize) -> Result<Self> {
if max_topics == 0 || max_subscribers_per_topic == 0 {
return Err(PubSubError::InvalidParameter);
}
let mut subscribers = alloc::vec::Vec::with_capacity(max_topics);
for _ in 0..max_topics {
let mut topic_subscribers = alloc::vec::Vec::with_capacity(max_subscribers_per_topic);
for _ in 0..max_subscribers_per_topic {
topic_subscribers.push(None);
}
subscribers.push(topic_subscribers);
}
let mut wildcard_subscribers = alloc::vec::Vec::with_capacity(max_subscribers_per_topic);
for _ in 0..max_subscribers_per_topic {
wildcard_subscribers.push(None);
}
let id_map = alloc::vec::Vec::with_capacity(max_topics);
let mut expected_seq_nums = alloc::vec::Vec::with_capacity(max_topics);
for _ in 0..max_topics {
expected_seq_nums.push(1);
}
Ok(Self {
max_topics,
max_subscribers_per_topic,
subscribers,
wildcard_subscribers,
next_subscription_id: 1, global_subscribers: alloc::vec::Vec::new(),
topic_map: alloc::collections::BTreeMap::new(),
id_map,
expected_seq_nums,
})
}
pub fn subscribe(&mut self, topic_id: u16, callback: PubSubCallback) -> Result<usize> {
let subscription_id = self.next_subscription_id;
self.next_subscription_id += 1;
let subscriber = Subscriber {
id: subscription_id,
topic_id,
callback,
active: true,
last_active: Self::get_current_time(),
};
const WILDCARD_TOPIC_ID: u16 = 0xFFFF;
if topic_id == WILDCARD_TOPIC_ID {
let index = match self.wildcard_subscribers.iter().position(|s| s.is_none()) {
Some(idx) => idx,
None => {
if self.wildcard_subscribers.len() >= self.max_subscribers_per_topic {
return Err(PubSubError::ResourceExhausted);
}
self.wildcard_subscribers.push(None);
self.wildcard_subscribers.len() - 1
}
};
self.wildcard_subscribers[index] = Some(subscriber);
self.global_subscribers
.push(WILDCARD_TOPIC_ID as usize * 1000 + index);
} else {
if topic_id as usize >= self.max_topics {
return Err(PubSubError::InvalidParameter);
}
let topic_subscribers = &mut self.subscribers[topic_id as usize];
let index = match topic_subscribers.iter().position(|s| s.is_none()) {
Some(idx) => idx,
None => {
if topic_subscribers.len() >= self.max_subscribers_per_topic {
return Err(PubSubError::ResourceExhausted);
}
topic_subscribers.push(None);
topic_subscribers.len() - 1
}
};
topic_subscribers[index] = Some(subscriber);
self.global_subscribers
.push(topic_id as usize * 1000 + index); }
Ok(subscription_id)
}
pub fn unsubscribe(&mut self, subscription_id: usize) -> Result<()> {
for topic_subscribers in &mut self.subscribers {
for subscriber in topic_subscribers {
if let Some(ref mut s) = subscriber {
if s.id == subscription_id {
s.active = false;
*subscriber = None;
return Ok(());
}
}
}
}
for subscriber in &mut self.wildcard_subscribers {
if let Some(ref mut s) = subscriber {
if s.id == subscription_id {
s.active = false;
*subscriber = None;
return Ok(());
}
}
}
Err(PubSubError::SubscriptionNotFound)
}
pub fn handle_data(&mut self, topic_id: u16, data: &[u8]) -> Result<()> {
if topic_id as usize >= self.max_topics {
return Err(PubSubError::InvalidParameter);
}
let current_time = Self::get_current_time();
let topic_subscribers = &mut self.subscribers[topic_id as usize];
for subscriber in topic_subscribers {
if let Some(ref mut s) = subscriber {
if s.active {
s.last_active = current_time;
let continue_flag = (s.callback)(topic_id, data);
if !continue_flag {
s.active = false;
*subscriber = None;
}
}
}
}
for subscriber in &mut self.wildcard_subscribers {
if let Some(ref mut s) = subscriber {
if s.active {
s.last_active = current_time;
let continue_flag = (s.callback)(topic_id, data);
if !continue_flag {
s.active = false;
*subscriber = None;
}
}
}
}
Ok(())
}
pub fn cleanup_inactive(&mut self, timeout: u64) -> Result<()> {
let current_time = Self::get_current_time();
for topic_subscribers in &mut self.subscribers {
for subscriber in topic_subscribers {
if let Some(ref mut s) = subscriber {
if s.active {
if current_time - s.last_active > timeout {
s.active = false;
*subscriber = None;
}
}
}
}
}
for subscriber in &mut self.wildcard_subscribers {
if let Some(ref mut s) = subscriber {
if s.active {
if current_time - s.last_active > timeout {
s.active = false;
*subscriber = None;
}
}
}
}
Ok(())
}
fn get_current_time() -> u64 {
#[cfg(feature = "posix")]
{
let now = std::time::SystemTime::now();
let duration = now.duration_since(std::time::UNIX_EPOCH).unwrap();
return duration.as_millis() as u64;
}
#[cfg(feature = "baremetal")]
{
return 0u64;
}
#[cfg(not(any(feature = "posix", feature = "baremetal")))]
{
return 0u64;
}
}
pub fn get_subscriber_count(&self, topic_id: u16) -> Result<usize> {
if topic_id == WILDCARD_TOPIC_ID {
let count = self
.wildcard_subscribers
.iter()
.filter(|s| s.is_some() && s.as_ref().unwrap().active)
.count();
return Ok(count);
}
if topic_id as usize >= self.max_topics {
return Err(PubSubError::InvalidParameter);
}
let count = self.subscribers[topic_id as usize]
.iter()
.filter(|s| s.is_some() && s.as_ref().unwrap().active)
.count();
Ok(count)
}
pub fn register_topic(&mut self, topic_name: &'static str, topic_id: u16) -> Result<()> {
if (topic_id as usize) >= self.max_topics {
return Err(PubSubError::InvalidParameter);
}
if self.topic_map.contains_key(topic_name) {
return Err(PubSubError::InvalidParameter);
}
if (topic_id as usize) < self.id_map.len() && self.id_map[topic_id as usize].is_some() {
return Err(PubSubError::InvalidParameter);
}
self.topic_map.insert(topic_name, topic_id);
while (topic_id as usize) >= self.id_map.len() {
self.id_map.push(None);
}
self.id_map[topic_id as usize] = Some(topic_name);
Ok(())
}
pub fn get_topic_id(&self, topic_name: &str) -> Option<u16> {
self.topic_map.get(topic_name).copied()
}
pub fn get_topic_name(&self, topic_id: u16) -> Option<&'static str> {
if (topic_id as usize) < self.id_map.len() {
self.id_map[topic_id as usize]
} else {
None
}
}
pub fn get_total_subscriber_count(&self) -> usize {
let mut count = 0;
for topic_subscribers in &self.subscribers {
count += topic_subscribers
.iter()
.filter(|s| s.is_some() && s.as_ref().unwrap().active)
.count();
}
count += self
.wildcard_subscribers
.iter()
.filter(|s| s.is_some() && s.as_ref().unwrap().active)
.count();
count
}
pub fn check_seq_num(&mut self, topic_id: u16, seq_num: u32) -> Option<u32> {
if topic_id as usize >= self.max_topics {
return None;
}
while topic_id as usize >= self.expected_seq_nums.len() {
self.expected_seq_nums.push(1);
}
let expected_seq_num = self.expected_seq_nums[topic_id as usize];
if seq_num < expected_seq_num {
return None;
} else if seq_num > expected_seq_num {
return Some(expected_seq_num);
} else {
self.expected_seq_nums[topic_id as usize] = expected_seq_num + 1;
return None;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_subscribe_unsubscribe() {
let mut manager = SubscriberManager::new(32, 16).unwrap();
let callback = |_topic_id: u16, _data: &[u8]| -> bool { true };
let subscription_id = manager.subscribe(0, callback).unwrap();
assert!(subscription_id > 0);
assert_eq!(manager.get_subscriber_count(0).unwrap(), 1);
manager.unsubscribe(subscription_id).unwrap();
assert_eq!(manager.get_subscriber_count(0).unwrap(), 0);
}
#[test]
fn test_handle_data() {
static mut CALLBACK_CALLED: bool = false;
fn test_callback(_topic_id: u16, data: &[u8]) -> bool {
unsafe {
CALLBACK_CALLED = true;
}
assert_eq!(data, b"test data");
true
}
let mut manager = SubscriberManager::new(32, 16).unwrap();
manager.subscribe(0, test_callback).unwrap();
manager.handle_data(0, b"test data").unwrap();
unsafe {
assert!(CALLBACK_CALLED);
}
}
#[test]
fn test_cleanup_inactive() {
let mut manager = SubscriberManager::new(32, 16).unwrap();
let callback = |_topic_id: u16, _data: &[u8]| -> bool { true };
manager.subscribe(0, callback).unwrap();
assert_eq!(manager.get_subscriber_count(0).unwrap(), 1);
manager.subscribe(1, callback).unwrap();
assert_eq!(manager.get_subscriber_count(1).unwrap(), 1);
assert_eq!(manager.get_subscriber_count(0).unwrap(), 1);
assert_eq!(manager.get_subscriber_count(1).unwrap(), 1);
assert_eq!(manager.get_total_subscriber_count(), 2);
}
}