use super::Error;
use super::service_info::Subscriber;
use super::subscription_manager::{SUBSCRIBERS_PER_GROUP, SubscriptionHandle};
use crate::e2e::E2EKey;
use crate::protocol::{Header, Message};
use crate::traits::{PayloadWireFormat, WireFormat};
use crate::transport::{E2ERegistryHandle, SharedHandle, TransportSocket};
#[cfg(test)]
use alloc::sync::Arc;
use core::marker::PhantomData;
use core::net::SocketAddrV4;
use heapless::Vec as HeaplessVec;
const _: () = assert!(
SUBSCRIBERS_PER_GROUP >= 1,
"SUBSCRIBERS_PER_GROUP must be >= 1 for the publish snapshot to fit any subscribers"
);
pub struct EventPublisher<R, S, H, T>
where
R: E2ERegistryHandle,
S: SubscriptionHandle,
T: TransportSocket + 'static,
H: SharedHandle<T>,
{
subscriptions: S,
socket: H,
e2e_registry: R,
_phantom: PhantomData<fn() -> T>,
}
impl<R, S, H, T> EventPublisher<R, S, H, T>
where
R: E2ERegistryHandle,
S: SubscriptionHandle,
T: TransportSocket + 'static,
H: SharedHandle<T>,
{
pub fn new(subscriptions: S, socket: H, e2e_registry: R) -> Self {
Self {
subscriptions,
socket,
e2e_registry,
_phantom: PhantomData,
}
}
#[allow(clippy::too_many_lines)]
pub async fn publish_event_with_buffers<P: PayloadWireFormat>(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
message: &Message<P>,
msg_buf: &mut [u8],
protected_buf: &mut [u8],
) -> Result<usize, Error> {
let mut subscribers: HeaplessVec<SocketAddrV4, SUBSCRIBERS_PER_GROUP> = HeaplessVec::new();
let mut visit = |sub: &Subscriber| {
let _ = subscribers.push(sub.address);
};
let _total = self
.subscriptions
.for_each_subscriber(service_id, instance_id, event_group_id, &mut visit)
.await;
if subscribers.is_empty() {
crate::log::trace!(
"No subscribers for service 0x{:04X}, instance {}, event group 0x{:04X}",
service_id,
instance_id,
event_group_id
);
return Ok(0);
}
let required_size = message.required_size();
if required_size > msg_buf.len() {
crate::log::error!(
"Message size ({} bytes) exceeds msg_buf.len() ({}); dropping publish",
required_size,
msg_buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
let mut message_length = message.encode_to_slice(msg_buf)?;
{
let key = E2EKey::from_message_id(message.header().message_id());
if self.e2e_registry.contains_key(&key) {
let upper_header: [u8; 8] = msg_buf[8..16].try_into().expect("upper header slice");
let result = self.e2e_registry.protect(
key,
&msg_buf[16..message_length],
upper_header,
protected_buf,
);
match result {
Some(Ok(protected_len)) => {
if 16 + protected_len > msg_buf.len() {
crate::log::error!(
"E2E-protected datagram ({} bytes, header + protected payload) \
exceeds msg_buf.len() ({}); dropping publish",
16 + protected_len,
msg_buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
#[allow(clippy::cast_possible_truncation)]
let new_length: u32 = 8 + protected_len as u32;
msg_buf[4..8].copy_from_slice(&new_length.to_be_bytes());
msg_buf[16..16 + protected_len]
.copy_from_slice(&protected_buf[..protected_len]);
message_length = 16 + protected_len;
}
Some(Err(e @ crate::e2e::Error::BufferTooSmall { .. })) => {
crate::log::error!(
"E2E protect error (buffer too small): {:?}; dropping publish",
e
);
return Err(Error::Capacity("udp_buffer"));
}
None => unreachable!("contains_key was true"),
}
}
}
let datagram = &msg_buf[..message_length];
let mut sent_count = 0usize;
let mut last_err: Option<crate::transport::TransportError> = None;
for addr in &subscribers {
match self.socket.get().send_to(datagram, *addr).await {
Ok(()) => {
sent_count += 1;
crate::log::trace!(
"Sent event to subscriber {} ({} bytes)",
addr,
message_length
);
}
Err(e) => {
crate::log::error!("Failed to send event to subscriber {}: {:?}", addr, e);
last_err = Some(e);
}
}
}
crate::log::debug!(
"Published event to {}/{} subscribers for service 0x{:04X}",
sent_count,
subscribers.len(),
service_id
);
if sent_count == 0 {
return Err(Error::Transport(
last_err.unwrap_or(crate::transport::TransportError::Unsupported),
));
}
Ok(sent_count)
}
#[cfg(feature = "_alloc")]
pub async fn publish_event<P: PayloadWireFormat>(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
message: &Message<P>,
) -> Result<usize, Error>
where
for<'a> S::ForEachFuture<'a>: Send,
{
let mut msg_buf = alloc::vec![0u8; crate::UDP_BUFFER_SIZE];
let mut protected_buf = alloc::vec![0u8; crate::UDP_BUFFER_SIZE];
self.publish_event_with_buffers(
service_id,
instance_id,
event_group_id,
message,
&mut msg_buf,
&mut protected_buf,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn publish_raw_event_with_buffers(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
event_id: u16,
request_id: u32,
protocol_version: u8,
interface_version: u8,
payload: &[u8],
buf: &mut [u8],
) -> Result<usize, Error> {
let mut subscribers: HeaplessVec<SocketAddrV4, SUBSCRIBERS_PER_GROUP> = HeaplessVec::new();
let mut visit = |sub: &Subscriber| {
let _ = subscribers.push(sub.address);
};
let _total = self
.subscriptions
.for_each_subscriber(service_id, instance_id, event_group_id, &mut visit)
.await;
if subscribers.is_empty() {
return Ok(0);
}
if buf.len() < 16 {
crate::log::error!(
"raw event buffer ({} bytes) too small for the 16-byte SOME/IP header; dropping publish",
buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
if payload.len() > buf.len().saturating_sub(16) {
crate::log::error!(
"raw event payload ({} bytes) + 16-byte header exceeds buf.len() ({}); dropping publish",
payload.len(),
buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
let header = Header::new_event(
service_id,
event_id,
request_id,
protocol_version,
interface_version,
payload.len(),
);
let header_len = header.encode_to_slice(buf)?;
let Some(total_len) = header_len.checked_add(payload.len()) else {
crate::log::error!(
"raw event length computation overflowed usize (header_len={}, payload.len()={}); dropping publish",
header_len,
payload.len()
);
return Err(Error::Capacity("udp_buffer"));
};
if total_len > buf.len() {
crate::log::error!(
"raw event ({} bytes) exceeds buf.len() ({}); dropping publish",
total_len,
buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
buf[header_len..total_len].copy_from_slice(payload);
let datagram = &buf[..total_len];
let mut sent_count = 0usize;
let mut last_err: Option<crate::transport::TransportError> = None;
for addr in &subscribers {
match self.socket.get().send_to(datagram, *addr).await {
Ok(()) => {
sent_count += 1;
}
Err(e) => {
crate::log::error!("Failed to send raw event to {}: {:?}", addr, e);
last_err = Some(e);
}
}
}
if sent_count == 0 {
return Err(Error::Transport(
last_err.unwrap_or(crate::transport::TransportError::Unsupported),
));
}
Ok(sent_count)
}
#[cfg(feature = "_alloc")]
#[allow(clippy::too_many_arguments)]
pub async fn publish_raw_event(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
event_id: u16,
request_id: u32,
protocol_version: u8,
interface_version: u8,
payload: &[u8],
) -> Result<usize, Error>
where
for<'a> S::ForEachFuture<'a>: Send,
{
let mut buf = alloc::vec![0u8; crate::UDP_BUFFER_SIZE];
self.publish_raw_event_with_buffers(
service_id,
instance_id,
event_group_id,
event_id,
request_id,
protocol_version,
interface_version,
payload,
&mut buf,
)
.await
}
pub async fn has_subscribers(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
) -> bool
where
for<'a> S::ForEachFuture<'a>: Send,
{
self.subscriber_count(service_id, instance_id, event_group_id)
.await
> 0
}
pub async fn register_subscriber(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: core::net::SocketAddrV4,
) -> Result<(), crate::server::SubscribeError> {
self.subscriptions
.subscribe(service_id, instance_id, event_group_id, subscriber_addr)
.await
}
pub async fn remove_subscriber(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
subscriber_addr: core::net::SocketAddrV4,
) {
self.subscriptions
.unsubscribe(service_id, instance_id, event_group_id, subscriber_addr)
.await;
}
pub async fn subscriber_count(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
) -> usize
where
for<'a> S::ForEachFuture<'a>: Send,
{
let mut visit = |_: &Subscriber| {};
self.subscriptions
.for_each_subscriber(service_id, instance_id, event_group_id, &mut visit)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn publish_raw_event_to_with_buffers(
&self,
target: SocketAddrV4,
service_id: u16,
instance_id: u16,
event_group_id: u16,
event_id: u16,
request_id: u32,
protocol_version: u8,
interface_version: u8,
payload: &[u8],
buf: &mut [u8],
) -> Result<usize, Error> {
let mut is_subscribed = false;
{
let mut visit = |sub: &Subscriber| {
if sub.address == target {
is_subscribed = true;
}
};
self.subscriptions
.for_each_subscriber(service_id, instance_id, event_group_id, &mut visit)
.await;
}
if !is_subscribed {
return Ok(0);
}
if buf.len() < 16 {
crate::log::error!(
"raw event buffer ({} bytes) too small for the 16-byte SOME/IP header; dropping publish",
buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
if payload.len() > buf.len().saturating_sub(16) {
crate::log::error!(
"raw event payload ({} bytes) + 16-byte header exceeds buf.len() ({}); dropping publish",
payload.len(),
buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
let header = Header::new_event(
service_id,
event_id,
request_id,
protocol_version,
interface_version,
payload.len(),
);
let header_len = header.encode_to_slice(buf)?;
let Some(total_len) = header_len.checked_add(payload.len()) else {
crate::log::error!(
"raw event length computation overflowed usize (header_len={}, payload.len()={}); dropping publish",
header_len,
payload.len()
);
return Err(Error::Capacity("udp_buffer"));
};
if total_len > buf.len() {
crate::log::error!(
"raw event ({} bytes) exceeds buf.len() ({}); dropping publish",
total_len,
buf.len()
);
return Err(Error::Capacity("udp_buffer"));
}
buf[header_len..total_len].copy_from_slice(payload);
let datagram = &buf[..total_len];
match self.socket.get().send_to(datagram, target).await {
Ok(()) => Ok(1),
Err(e) => {
crate::log::error!("Failed to send raw event to {}: {:?}", target, e);
Err(Error::Transport(e))
}
}
}
#[cfg(feature = "_alloc")]
#[allow(clippy::too_many_arguments)]
pub async fn publish_raw_event_to(
&self,
target: SocketAddrV4,
service_id: u16,
instance_id: u16,
event_group_id: u16,
event_id: u16,
request_id: u32,
protocol_version: u8,
interface_version: u8,
payload: &[u8],
) -> Result<usize, Error>
where
for<'a> S::ForEachFuture<'a>: Send,
{
let mut buf = alloc::vec![0u8; crate::UDP_BUFFER_SIZE];
self.publish_raw_event_to_with_buffers(
target,
service_id,
instance_id,
event_group_id,
event_id,
request_id,
protocol_version,
interface_version,
payload,
&mut buf,
)
.await
}
pub async fn subscriber_addresses(
&self,
service_id: u16,
instance_id: u16,
event_group_id: u16,
) -> HeaplessVec<SocketAddrV4, SUBSCRIBERS_PER_GROUP>
where
for<'a> S::ForEachFuture<'a>: Send,
{
let mut addrs: HeaplessVec<SocketAddrV4, SUBSCRIBERS_PER_GROUP> = HeaplessVec::new();
let mut visit = |sub: &Subscriber| {
let _ = addrs.push(sub.address);
};
self.subscriptions
.for_each_subscriber(service_id, instance_id, event_group_id, &mut visit)
.await;
addrs
}
}
#[cfg(all(test, feature = "server-tokio"))]
mod tests {
use super::*;
use crate::UDP_BUFFER_SIZE;
use crate::e2e::E2ERegistry;
use crate::protocol::sd::test_support::{TestPayload, empty_sd_header};
use crate::server::SubscriptionManager;
use crate::tokio_transport::TokioSocket;
use std::net::{Ipv4Addr, SocketAddrV4};
use std::sync::Mutex;
use std::vec;
use std::vec::Vec;
use tokio::net::UdpSocket;
use tokio::sync::RwLock;
type TestEventPublisher = EventPublisher<
Arc<Mutex<E2ERegistry>>,
Arc<RwLock<SubscriptionManager>>,
Arc<TokioSocket>,
TokioSocket,
>;
fn test_registry() -> Arc<Mutex<E2ERegistry>> {
Arc::new(Mutex::new(E2ERegistry::new()))
}
async fn bind_tokio_socket() -> Arc<TokioSocket> {
use crate::transport::{SocketOptions, TransportFactory};
let factory = crate::tokio_transport::TokioTransport;
Arc::new(
factory
.bind(
SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0),
&SocketOptions::new(),
)
.await
.expect("bind tokio socket for test"),
)
}
async fn make_publisher(
subscriptions: Arc<RwLock<SubscriptionManager>>,
) -> (TestEventPublisher, Arc<TokioSocket>) {
let socket = bind_tokio_socket().await;
let publisher = EventPublisher::new(subscriptions, Arc::clone(&socket), test_registry());
(publisher, socket)
}
fn make_test_message() -> Message<TestPayload> {
Message::new_sd(0x0001, &empty_sd_header())
}
#[tokio::test]
async fn test_event_publisher_creation() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let socket = bind_tokio_socket().await;
let publisher = EventPublisher::new(subscriptions, socket, test_registry());
assert!(std::mem::size_of_val(&publisher) > 0);
}
#[tokio::test]
async fn test_publish_event_no_subscribers() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(subscriptions).await;
let msg = make_test_message();
let count = publisher.publish_event(0x5B, 1, 0x01, &msg).await.unwrap();
assert_eq!(count, 0);
}
#[tokio::test]
async fn test_publish_event_with_subscriber() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let receiver = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let core::net::SocketAddr::V4(recv_addr) = receiver.local_addr().unwrap() else {
panic!("expected v4 source address");
};
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, recv_addr).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
let msg = make_test_message();
let count = publisher.publish_event(0x5B, 1, 0x01, &msg).await.unwrap();
assert_eq!(count, 1);
let mut buf = [0u8; 1024];
let (len, _) = tokio::time::timeout(
std::time::Duration::from_secs(2),
receiver.recv_from(&mut buf),
)
.await
.expect("timeout receiving event")
.unwrap();
assert!(len > 0);
}
#[tokio::test]
async fn test_publish_raw_event_no_subscribers() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(subscriptions).await;
let count = publisher
.publish_raw_event(0x5B, 1, 0x01, 0x8001, 0x0001, 0x01, 0x01, &[0xAA, 0xBB])
.await
.unwrap();
assert_eq!(count, 0);
}
#[tokio::test]
async fn test_publish_raw_event_exceeds_udp_buffer_returns_capacity_error() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, addr).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
let too_big = vec![0u8; UDP_BUFFER_SIZE];
let err = publisher
.publish_raw_event(0x5B, 1, 0x01, 0x8001, 0x0001, 0x01, 0x01, &too_big)
.await
.expect_err("oversize payload must error, not report Ok(0)");
match err {
Error::Capacity(tag) => assert_eq!(tag, "udp_buffer"),
other => panic!("expected Error::Capacity(\"udp_buffer\"), got {other:?}"),
}
}
#[tokio::test]
async fn publish_event_returns_err_when_every_send_fails() {
use crate::transport::{IoErrorKind, ReceivedDatagram, TransportError, TransportSocket};
use core::future::{Future, Ready, ready};
use core::pin::Pin;
use core::task::{Context, Poll};
struct AlwaysFailSocket;
struct AlwaysFailSend;
impl Future for AlwaysFailSend {
type Output = Result<(), TransportError>;
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
Poll::Ready(Err(TransportError::Io(IoErrorKind::NetworkUnreachable)))
}
}
impl TransportSocket for AlwaysFailSocket {
type SendFuture<'a> = AlwaysFailSend;
type RecvFuture<'a> = Ready<Result<ReceivedDatagram, TransportError>>;
fn send_to<'a>(&'a self, _buf: &'a [u8], _t: SocketAddrV4) -> Self::SendFuture<'a> {
AlwaysFailSend
}
fn recv_from<'a>(&'a self, _buf: &'a mut [u8]) -> Self::RecvFuture<'a> {
ready(Err(TransportError::Unsupported))
}
fn local_addr(&self) -> Result<SocketAddrV4, TransportError> {
Ok(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0))
}
fn join_multicast_v4(&self, _g: Ipv4Addr, _i: Ipv4Addr) -> Result<(), TransportError> {
Ok(())
}
fn leave_multicast_v4(&self, _g: Ipv4Addr, _i: Ipv4Addr) -> Result<(), TransportError> {
Ok(())
}
}
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, addr).unwrap();
}
#[allow(
clippy::type_complexity,
reason = "tests reasonably spell out the full type for clarity"
)]
let publisher: EventPublisher<
Arc<Mutex<E2ERegistry>>,
Arc<RwLock<SubscriptionManager>>,
Arc<AlwaysFailSocket>,
AlwaysFailSocket,
> = EventPublisher::new(subscriptions, Arc::new(AlwaysFailSocket), test_registry());
let msg = make_test_message();
let err = publisher
.publish_event(0x5B, 1, 0x01, &msg)
.await
.expect_err("total-failure path must surface Err, not Ok(0)");
match err {
Error::Transport(TransportError::Io(IoErrorKind::NetworkUnreachable)) => {}
other => panic!(
"expected Transport(Io(NetworkUnreachable)) from total-failure send, got {other:?}"
),
}
}
#[tokio::test]
async fn publish_raw_event_returns_err_when_every_send_fails() {
use crate::transport::{IoErrorKind, ReceivedDatagram, TransportError, TransportSocket};
use core::future::{Future, Ready, ready};
use core::pin::Pin;
use core::task::{Context, Poll};
struct AlwaysFailSocket;
struct AlwaysFailSend;
impl Future for AlwaysFailSend {
type Output = Result<(), TransportError>;
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
Poll::Ready(Err(TransportError::Io(IoErrorKind::ConnectionRefused)))
}
}
impl TransportSocket for AlwaysFailSocket {
type SendFuture<'a> = AlwaysFailSend;
type RecvFuture<'a> = Ready<Result<ReceivedDatagram, TransportError>>;
fn send_to<'a>(&'a self, _buf: &'a [u8], _t: SocketAddrV4) -> Self::SendFuture<'a> {
AlwaysFailSend
}
fn recv_from<'a>(&'a self, _buf: &'a mut [u8]) -> Self::RecvFuture<'a> {
ready(Err(TransportError::Unsupported))
}
fn local_addr(&self) -> Result<SocketAddrV4, TransportError> {
Ok(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0))
}
fn join_multicast_v4(&self, _g: Ipv4Addr, _i: Ipv4Addr) -> Result<(), TransportError> {
Ok(())
}
fn leave_multicast_v4(&self, _g: Ipv4Addr, _i: Ipv4Addr) -> Result<(), TransportError> {
Ok(())
}
}
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, addr).unwrap();
}
#[allow(
clippy::type_complexity,
reason = "tests reasonably spell out the full type for clarity"
)]
let publisher: EventPublisher<
Arc<Mutex<E2ERegistry>>,
Arc<RwLock<SubscriptionManager>>,
Arc<AlwaysFailSocket>,
AlwaysFailSocket,
> = EventPublisher::new(subscriptions, Arc::new(AlwaysFailSocket), test_registry());
let err = publisher
.publish_raw_event(0x5B, 1, 0x01, 0x8001, 0x0001, 0x01, 0x01, &[0xAA, 0xBB])
.await
.expect_err("total-failure path must surface Err, not Ok(0)");
match err {
Error::Transport(TransportError::Io(IoErrorKind::ConnectionRefused)) => {}
other => panic!("expected Transport(Io(ConnectionRefused)), got {other:?}"),
}
}
#[tokio::test]
async fn publish_event_pre_encode_exceeds_udp_buffer_returns_capacity_error() {
use crate::RawPayload;
use crate::protocol::{Header, MessageId, MessageType, MessageTypeField, ReturnCode};
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, addr).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
let message_id = MessageId::new_from_service_and_method(0x1234, 0x5678);
let payload_len = UDP_BUFFER_SIZE - 16 + 1 ;
let payload_bytes = vec![0u8; payload_len];
let payload = RawPayload::from_payload_bytes(message_id, &payload_bytes).unwrap();
let header = Header::new(
message_id,
0x0001_0001,
0x01,
0x01,
MessageTypeField::new(MessageType::Request, false),
ReturnCode::Ok,
payload_bytes.len(),
);
let message = Message::new(header, payload);
assert!(
message.required_size() > UDP_BUFFER_SIZE,
"fixture must exceed cap",
);
let err = publisher
.publish_event(0x5B, 1, 0x01, &message)
.await
.expect_err("oversize message must error, not report Ok(_)");
match err {
Error::Capacity(tag) => assert_eq!(tag, "udp_buffer"),
other => panic!("expected Error::Capacity(\"udp_buffer\"), got {other:?}"),
}
}
#[tokio::test]
async fn test_publish_event_e2e_protected_exceeds_udp_buffer_returns_capacity_error() {
use crate::RawPayload;
use crate::e2e::{E2EProfile, Profile4Config};
use crate::protocol::MessageId;
let message_id = MessageId::new_from_service_and_method(0x5B, 0x8001);
let key = E2EKey::from_message_id(message_id);
let mut reg = E2ERegistry::new();
reg.register(key, E2EProfile::Profile4(Profile4Config::new(0, 15)))
.expect("E2E registry has capacity for one entry");
let e2e_registry = Arc::new(Mutex::new(reg));
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999))
.unwrap();
}
let socket = bind_tokio_socket().await;
let publisher = EventPublisher::new(subscriptions, socket, e2e_registry);
let payload_len = UDP_BUFFER_SIZE - 16; let payload_bytes = vec![0u8; payload_len];
let payload = RawPayload::from_payload_bytes(message_id, &payload_bytes).unwrap();
let header = Header::new_event(
message_id.service_id(),
message_id.method_id(),
0x0001_0001,
0x01,
0x01,
payload_bytes.len(),
);
let message = Message::new(header, payload);
assert!(
message.required_size() <= UDP_BUFFER_SIZE,
"fixture's raw size must fit the cap so the pre-encode check passes and \
we actually exercise the post-protect guard",
);
let err = publisher
.publish_event(0x5B, 1, 0x01, &message)
.await
.expect_err("E2E-protected oversize message must error, not report Ok(n)");
match err {
Error::Capacity(tag) => assert_eq!(tag, "udp_buffer"),
other => panic!("expected Error::Capacity(\"udp_buffer\"), got {other:?}"),
}
}
#[tokio::test]
async fn test_publish_raw_event_with_subscriber() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let receiver = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let core::net::SocketAddr::V4(recv_addr) = receiver.local_addr().unwrap() else {
panic!("expected v4 source address");
};
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, recv_addr).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
let payload = [0xDE, 0xAD];
let count = publisher
.publish_raw_event(0x5B, 1, 0x01, 0x8001, 0x0001, 0x01, 0x01, &payload)
.await
.unwrap();
assert_eq!(count, 1);
let mut buf = [0u8; 1024];
let (len, _) = tokio::time::timeout(
std::time::Duration::from_secs(2),
receiver.recv_from(&mut buf),
)
.await
.expect("timeout receiving raw event")
.unwrap();
assert_eq!(len, 18);
assert_eq!(&buf[16..18], &payload);
}
#[tokio::test]
async fn test_subscriber_count() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let addr1 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9001);
let addr2 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9002);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, addr1).unwrap();
mgr.subscribe(0x5B, 1, 0x01, addr2).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 2);
}
#[tokio::test]
async fn test_has_subscribers() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
assert!(!publisher.has_subscribers(0x5B, 1, 0x01).await);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 1, 0x01, SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9001))
.unwrap();
}
assert!(publisher.has_subscribers(0x5B, 1, 0x01).await);
}
const ADDR_A: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9001);
const ADDR_B: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9002);
const ADDR_C: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9003);
#[tokio::test]
async fn register_subscriber_adds_to_manager() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
assert!(!publisher.has_subscribers(0x5B, 1, 0x01).await);
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
assert!(publisher.has_subscribers(0x5B, 1, 0x01).await);
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 1);
}
#[tokio::test]
async fn register_subscriber_is_idempotent_on_repeat() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 1);
}
#[tokio::test]
async fn register_subscriber_separates_different_eventgroups() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher
.register_subscriber(0x5B, 1, 0x02, ADDR_A)
.await
.unwrap();
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 1);
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x02).await, 1);
assert!(publisher.has_subscribers(0x5B, 1, 0x01).await);
assert!(publisher.has_subscribers(0x5B, 1, 0x02).await);
}
#[tokio::test]
async fn remove_subscriber_happy_path() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
assert!(publisher.has_subscribers(0x5B, 1, 0x01).await);
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_A).await;
assert!(!publisher.has_subscribers(0x5B, 1, 0x01).await);
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 0);
}
#[tokio::test]
async fn remove_subscriber_leaves_siblings_alone() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_B)
.await
.unwrap();
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_C)
.await
.unwrap();
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 3);
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_B).await;
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 2);
let mgr = subscriptions.read().await;
let subscribers = mgr.get_subscribers(0x5B, 1, 0x01);
let addrs: Vec<_> = subscribers.iter().map(|s| s.address).collect();
assert!(addrs.contains(&ADDR_A));
assert!(addrs.contains(&ADDR_C));
assert!(!addrs.contains(&ADDR_B));
}
#[tokio::test]
async fn remove_subscriber_nonexistent_is_noop() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_A).await;
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 0);
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_B).await;
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 1);
publisher.remove_subscriber(0x99, 1, 0x01, ADDR_A).await;
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 1);
}
#[tokio::test]
async fn remove_subscriber_all_then_has_subscribers_false() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_B)
.await
.unwrap();
assert!(publisher.has_subscribers(0x5B, 1, 0x01).await);
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_A).await;
assert!(publisher.has_subscribers(0x5B, 1, 0x01).await);
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_B).await;
assert!(!publisher.has_subscribers(0x5B, 1, 0x01).await);
}
#[tokio::test]
async fn register_and_remove_roundtrip_preserves_idempotence() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(Arc::clone(&subscriptions)).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_A).await;
publisher
.register_subscriber(0x5B, 1, 0x01, ADDR_A)
.await
.unwrap();
assert_eq!(publisher.subscriber_count(0x5B, 1, 0x01).await, 1);
}
#[tokio::test]
async fn test_publish_raw_event_to_targets_one_subscriber() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let rx_a = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let rx_b = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let core::net::SocketAddr::V4(addr_a) = rx_a.local_addr().unwrap() else {
panic!("expected v4");
};
let core::net::SocketAddr::V4(addr_b) = rx_b.local_addr().unwrap() else {
panic!("expected v4");
};
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 2, 0x01, addr_a).unwrap();
mgr.subscribe(0x5B, 2, 0x01, addr_b).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
let payload = [0xAB, 0xCD];
let sent = publisher
.publish_raw_event_to(addr_a, 0x5B, 2, 0x01, 0x8001, 0x0001, 0x01, 0x01, &payload)
.await
.unwrap();
assert_eq!(sent, 1);
let mut buf = [0u8; 64];
let (len, _) =
tokio::time::timeout(std::time::Duration::from_secs(2), rx_a.recv_from(&mut buf))
.await
.expect("addr_a should receive")
.unwrap();
assert_eq!(len, 18);
assert_eq!(&buf[16..18], &payload);
let nothing = tokio::time::timeout(
std::time::Duration::from_millis(300),
rx_b.recv_from(&mut buf),
)
.await;
assert!(
nothing.is_err(),
"addr_b must not receive a targeted publish"
);
}
#[tokio::test]
async fn test_publish_raw_event_to_unsubscribed_target_sends_nothing() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _) = make_publisher(subscriptions).await;
let bogus = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
let sent = publisher
.publish_raw_event_to(bogus, 0x5B, 2, 0x01, 0x8001, 0x0001, 0x01, 0x01, &[0u8])
.await
.unwrap();
assert_eq!(sent, 0);
}
#[tokio::test]
async fn test_subscriber_addresses_lists_all_under_key() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let a = SocketAddrV4::new(Ipv4Addr::new(192, 168, 11, 101), 30682);
let b = SocketAddrV4::new(Ipv4Addr::new(192, 168, 11, 102), 30682);
{
let mut mgr = subscriptions.write().await;
mgr.subscribe(0x5B, 2, 0x01, a).unwrap();
mgr.subscribe(0x5B, 2, 0x01, b).unwrap();
}
let (publisher, _) = make_publisher(subscriptions).await;
let addrs = publisher.subscriber_addresses(0x5B, 2, 0x01).await;
assert_eq!(addrs.len(), 2);
assert!(addrs.contains(&a));
assert!(addrs.contains(&b));
}
#[allow(dead_code)]
fn assert_publisher_futures_are_send(
publisher: &TestEventPublisher,
message: &Message<TestPayload>,
) {
fn assert_send<T: Send>(_: T) {}
assert_send(publisher.register_subscriber(0x5B, 1, 0x01, ADDR_A));
assert_send(publisher.remove_subscriber(0x5B, 1, 0x01, ADDR_A));
assert_send(publisher.publish_event(0x5B, 1, 0x01, message));
assert_send(publisher.publish_raw_event(0x5B, 1, 0x01, 0x8001, 0x0001, 0x01, 0x01, &[0u8]));
assert_send(publisher.publish_raw_event_to(
ADDR_A,
0x5B,
1,
0x01,
0x8001,
0x0001,
0x01,
0x01,
&[0u8],
));
assert_send(publisher.has_subscribers(0x5B, 1, 0x01));
assert_send(publisher.subscriber_count(0x5B, 1, 0x01));
assert_send(publisher.subscriber_addresses(0x5B, 1, 0x01));
}
#[allow(dead_code)]
fn assert_publisher_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<TestEventPublisher>();
assert_send_sync::<Arc<TestEventPublisher>>();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn publisher_work_is_spawnable_on_a_multi_thread_runtime() {
let subscriptions = Arc::new(RwLock::new(SubscriptionManager::new()));
let (publisher, _socket) = make_publisher(Arc::clone(&subscriptions)).await;
let publisher = Arc::new(publisher);
let receiver = UdpSocket::bind("127.0.0.1:0")
.await
.expect("bind receiver socket");
let receiver_addr = match receiver.local_addr().expect("receiver local_addr") {
std::net::SocketAddr::V4(a) => a,
other => panic!("expected an IPv4 receiver address, got {other}"),
};
publisher
.register_subscriber(0x5B, 1, 0x01, receiver_addr)
.await
.expect("register subscriber");
let spawned = Arc::clone(&publisher);
let sent = tokio::spawn(async move {
assert!(spawned.has_subscribers(0x5B, 1, 0x01).await);
assert_eq!(spawned.subscriber_count(0x5B, 1, 0x01).await, 1);
assert_eq!(spawned.subscriber_addresses(0x5B, 1, 0x01).await.len(), 1);
let msg = make_test_message();
spawned
.publish_event(0x5B, 1, 0x01, &msg)
.await
.expect("publish from a spawned task")
})
.await
.expect("spawned publish task must not panic");
assert_eq!(sent, 1);
}
}