use std::{
borrow::ToOwned,
collections::{HashMap, VecDeque},
future,
net::{Ipv4Addr, SocketAddr, SocketAddrV4},
sync::{Arc, Mutex},
task::Poll,
};
use tokio::{
select,
sync::{
mpsc::{self, Receiver, Sender},
oneshot,
},
};
use tracing::{debug, error, info, trace, warn};
use crate::{
client::{
ClientUpdate, DiscoveryMessage,
service_registry::{ServiceEndpointInfo, ServiceInstanceId, ServiceRegistry},
session::{SessionTracker, SessionVerdict, TransportKind},
socket_manager::{ReceivedMessage, SocketManager},
},
e2e::E2ERegistry,
protocol::{self, Message},
traits::PayloadWireFormat,
};
use super::error::Error;
pub(super) enum ControlMessage<P: PayloadWireFormat> {
SetInterface(Ipv4Addr, oneshot::Sender<Result<(), Error>>),
BindDiscovery(oneshot::Sender<Result<(), Error>>),
UnbindDiscovery(oneshot::Sender<Result<(), Error>>),
SendSD(
SocketAddrV4,
P::SdHeader,
oneshot::Sender<Result<(), Error>>,
),
AddEndpoint(
u16,
u16,
SocketAddrV4,
u16,
oneshot::Sender<Result<(), Error>>,
),
RemoveEndpoint(u16, u16, oneshot::Sender<Result<(), Error>>),
SendToService {
service_id: u16,
instance_id: u16,
message: Message<P>,
send_complete: oneshot::Sender<Result<(), Error>>,
response: oneshot::Sender<Result<P, Error>>,
},
Subscribe {
service_id: u16,
instance_id: u16,
major_version: u8,
ttl: u32,
event_group_id: u16,
client_port: u16,
response: oneshot::Sender<Result<(), Error>>,
},
QueryRebootFlag(oneshot::Sender<crate::protocol::sd::RebootFlag>),
#[cfg(test)]
ForceSdSessionWrappedForTest(bool, oneshot::Sender<()>),
}
impl<P: PayloadWireFormat> std::fmt::Debug for ControlMessage<P> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::SetInterface(addr, _) => f.debug_tuple("SetInterface").field(addr).finish(),
Self::BindDiscovery(_) => f.write_str("BindDiscovery"),
Self::UnbindDiscovery(_) => f.write_str("UnbindDiscovery"),
Self::SendSD(addr, header, _) => {
f.debug_tuple("SendSD").field(addr).field(header).finish()
}
Self::AddEndpoint(sid, iid, addr, local_port, _) => f
.debug_tuple("AddEndpoint")
.field(sid)
.field(iid)
.field(addr)
.field(local_port)
.finish(),
Self::RemoveEndpoint(sid, iid, _) => f
.debug_tuple("RemoveEndpoint")
.field(sid)
.field(iid)
.finish(),
Self::SendToService {
service_id,
instance_id,
message,
..
} => f
.debug_struct("SendToService")
.field("service_id", service_id)
.field("instance_id", instance_id)
.field("message", message)
.finish_non_exhaustive(),
Self::Subscribe {
service_id,
instance_id,
event_group_id,
..
} => f
.debug_struct("Subscribe")
.field("service_id", service_id)
.field("instance_id", instance_id)
.field("event_group_id", event_group_id)
.finish_non_exhaustive(),
Self::QueryRebootFlag(_) => f.write_str("QueryRebootFlag"),
#[cfg(test)]
Self::ForceSdSessionWrappedForTest(b, _) => f
.debug_tuple("ForceSdSessionWrappedForTest")
.field(b)
.finish(),
}
}
}
impl<P: PayloadWireFormat> ControlMessage<P> {
pub fn set_interface(interface: Ipv4Addr) -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(receiver, Self::SetInterface(interface, sender))
}
pub fn bind_discovery() -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(receiver, Self::BindDiscovery(sender))
}
pub fn unbind_discovery() -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(receiver, Self::UnbindDiscovery(sender))
}
pub fn send_sd(
socket_addr: SocketAddrV4,
header: P::SdHeader,
) -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(receiver, Self::SendSD(socket_addr, header, sender))
}
pub fn add_endpoint(
service_id: u16,
instance_id: u16,
addr: SocketAddrV4,
local_port: u16,
) -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(
receiver,
Self::AddEndpoint(service_id, instance_id, addr, local_port, sender),
)
}
pub fn remove_endpoint(
service_id: u16,
instance_id: u16,
) -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(
receiver,
Self::RemoveEndpoint(service_id, instance_id, sender),
)
}
#[allow(clippy::type_complexity)]
pub fn send_to_service(
service_id: u16,
instance_id: u16,
message: Message<P>,
) -> (
oneshot::Receiver<Result<(), Error>>,
oneshot::Receiver<Result<P, Error>>,
Self,
) {
let (send_complete_tx, send_complete_rx) = oneshot::channel();
let (response_tx, response_rx) = oneshot::channel();
(
send_complete_rx,
response_rx,
Self::SendToService {
service_id,
instance_id,
message,
send_complete: send_complete_tx,
response: response_tx,
},
)
}
pub fn subscribe(
service_id: u16,
instance_id: u16,
major_version: u8,
ttl: u32,
event_group_id: u16,
client_port: u16,
) -> (oneshot::Receiver<Result<(), Error>>, Self) {
let (sender, receiver) = oneshot::channel();
(
receiver,
Self::Subscribe {
service_id,
instance_id,
major_version,
ttl,
event_group_id,
client_port,
response: sender,
},
)
}
pub fn query_reboot_flag() -> (oneshot::Receiver<crate::protocol::sd::RebootFlag>, Self) {
let (sender, receiver) = oneshot::channel();
(receiver, Self::QueryRebootFlag(sender))
}
#[cfg(test)]
pub fn force_sd_session_wrapped_for_test(wrapped: bool) -> (oneshot::Receiver<()>, Self) {
let (sender, receiver) = oneshot::channel();
(
receiver,
Self::ForceSdSessionWrappedForTest(wrapped, sender),
)
}
}
pub(super) struct Inner<PayloadDefinitions: PayloadWireFormat> {
control_receiver: Receiver<ControlMessage<PayloadDefinitions>>,
request_queue: VecDeque<ControlMessage<PayloadDefinitions>>,
pending_responses: HashMap<u32, oneshot::Sender<Result<PayloadDefinitions, Error>>>,
update_sender: mpsc::UnboundedSender<ClientUpdate<PayloadDefinitions>>,
interface: Ipv4Addr,
discovery_socket: Option<SocketManager<PayloadDefinitions>>,
discovery_unicast_socket: Option<SocketManager<PayloadDefinitions>>,
unicast_sockets: HashMap<u16, SocketManager<PayloadDefinitions>>,
session_tracker: SessionTracker,
service_registry: ServiceRegistry,
run: bool,
client_id: u16,
session_counter: u16,
sd_session_id: u16,
sd_session_has_wrapped: bool,
e2e_registry: Arc<Mutex<E2ERegistry>>,
multicast_loopback: bool,
phantom: std::marker::PhantomData<PayloadDefinitions>,
}
impl<PayloadDefinitions: PayloadWireFormat> std::fmt::Debug for Inner<PayloadDefinitions> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Inner")
.field("interface", &self.interface)
.field("session_tracker", &self.session_tracker)
.field("run", &self.run)
.field("client_id", &self.client_id)
.field("session_counter", &self.session_counter)
.finish_non_exhaustive()
}
}
#[allow(clippy::too_many_arguments)]
fn process_discovery<P>(
source: SocketAddr,
transport: TransportKind,
someip_header: protocol::Header,
sd_header: <P as PayloadWireFormat>::SdHeader,
session_tracker: &mut SessionTracker,
service_registry: &mut ServiceRegistry,
e2e_registry: &Arc<Mutex<E2ERegistry>>,
update_sender: &mpsc::UnboundedSender<ClientUpdate<P>>,
) where
P: PayloadWireFormat + Clone + std::fmt::Debug + 'static,
{
let session_id = (someip_header.request_id() & 0xFFFF) as u16;
let sd_payload = P::new_sd_payload(&sd_header);
let reboot_flag = sd_payload.sd_flags().map_or(
crate::protocol::sd::RebootFlag::Continuous,
crate::protocol::sd::Flags::reboot,
);
let mut rebooted = false;
for (svc_id, inst_id) in sd_payload.service_instances() {
let verdict =
session_tracker.check(source, transport, svc_id, inst_id, session_id, reboot_flag);
if verdict == SessionVerdict::Reboot {
rebooted = true;
}
}
for ep in sd_payload.offered_endpoints() {
let id = ServiceInstanceId {
service_id: ep.service_id,
instance_id: ep.instance_id,
};
if ep.is_offer {
if let Some(addr) = ep.addr {
service_registry.insert(
id,
ServiceEndpointInfo {
addr,
local_port: 0,
major_version: ep.major_version,
minor_version: ep.minor_version,
},
);
trace!(
"Registry: added 0x{:04X}.0x{:04X} -> {}",
ep.service_id, ep.instance_id, addr,
);
}
} else {
service_registry.remove(id);
trace!(
"Registry: removed 0x{:04X}.0x{:04X}",
ep.service_id, ep.instance_id,
);
}
}
if rebooted {
e2e_registry
.lock()
.expect("e2e registry lock poisoned")
.reset_source(source.ip());
let _ = update_sender.send(ClientUpdate::SenderRebooted(source));
}
let discovery_msg = DiscoveryMessage {
source,
someip_header,
sd_header,
};
let _ = update_sender.send(ClientUpdate::DiscoveryUpdated(discovery_msg));
}
impl<PayloadDefinitions> Inner<PayloadDefinitions>
where
PayloadDefinitions: PayloadWireFormat + Clone + std::fmt::Debug + 'static,
{
pub fn spawn(
interface: Ipv4Addr,
e2e_registry: Arc<Mutex<E2ERegistry>>,
multicast_loopback: bool,
) -> (
Sender<ControlMessage<PayloadDefinitions>>,
mpsc::UnboundedReceiver<ClientUpdate<PayloadDefinitions>>,
) {
info!("Initializing SOME/IP Client");
let (control_sender, control_receiver) = mpsc::channel(4);
let (update_sender, update_receiver) = mpsc::unbounded_channel();
let inner = Self {
control_receiver,
request_queue: VecDeque::new(),
pending_responses: HashMap::new(),
update_sender,
interface,
discovery_socket: None,
discovery_unicast_socket: None,
unicast_sockets: HashMap::new(),
session_tracker: SessionTracker::default(),
service_registry: ServiceRegistry::default(),
run: true,
client_id: 0x1234,
session_counter: 1,
sd_session_id: 1,
sd_session_has_wrapped: false,
e2e_registry,
multicast_loopback,
phantom: std::marker::PhantomData,
};
inner.run();
(control_sender, update_receiver)
}
fn bind_discovery(&mut self) -> Result<(), Error> {
if self.discovery_socket.is_some() {
Ok(())
} else {
let socket = SocketManager::bind_discovery_seeded(
self.interface,
Arc::clone(&self.e2e_registry),
self.sd_session_id,
self.sd_session_has_wrapped,
self.multicast_loopback,
)?;
self.discovery_socket = Some(socket);
match SocketManager::bind_discovery_unicast(
self.interface,
Arc::clone(&self.e2e_registry),
) {
Ok(unicast) => self.discovery_unicast_socket = Some(unicast),
Err(e) => error!("Failed to bind unicast discovery socket: {e}"),
}
Ok(())
}
}
async fn unbind_discovery(&mut self) {
debug!("Unbinding Discovery socket.");
if let Some(socket) = self.discovery_socket.take() {
self.sd_session_id = socket.session_id();
self.sd_session_has_wrapped =
socket.reboot_flag() == crate::protocol::sd::RebootFlag::Continuous;
socket.shut_down().await;
}
if let Some(socket) = self.discovery_unicast_socket.take() {
socket.shut_down().await;
}
}
fn set_interface(&mut self, interface: Ipv4Addr) {
self.interface = interface;
}
fn bind_unicast(&mut self, port: u16) -> Result<u16, Error> {
if port != 0
&& let Some(socket) = self.unicast_sockets.get(&port)
{
return Ok(socket.port());
}
let unicast_socket = SocketManager::bind(port, Arc::clone(&self.e2e_registry))?;
let bound_port = unicast_socket.port();
self.unicast_sockets.insert(bound_port, unicast_socket);
debug!("Bound unicast socket on port {}", bound_port);
Ok(bound_port)
}
async fn receive_discovery(
socket_manager: &mut Option<SocketManager<PayloadDefinitions>>,
) -> Result<
(
SocketAddr,
protocol::Header,
<PayloadDefinitions as PayloadWireFormat>::SdHeader,
),
Error,
> {
if let Some(receiver) = socket_manager {
match receiver.receive().await {
Some(result) => match result {
Ok(received) => {
let someip_header = received.message.header().clone();
if let Some(sd_header) = received.message.sd_header() {
Ok((received.source, someip_header, sd_header.to_owned()))
} else {
Err(Error::UnexpectedDiscoveryMessage(someip_header))
}
}
Err(err) => Err(err),
},
None => Err(Error::SocketClosedUnexpectedly),
}
} else {
future::pending().await
}
}
async fn receive_any_unicast(
unicast_sockets: &mut HashMap<u16, SocketManager<PayloadDefinitions>>,
) -> Result<ReceivedMessage<PayloadDefinitions>, Error> {
if unicast_sockets.is_empty() {
return future::pending().await;
}
std::future::poll_fn(|cx| {
for socket in unicast_sockets.values_mut() {
if let Poll::Ready(result) = socket.poll_receive(cx) {
return Poll::Ready(match result {
Some(msg) => msg,
None => Err(Error::SocketClosedUnexpectedly),
});
}
}
Poll::Pending
})
.await
}
#[allow(clippy::too_many_lines)]
async fn handle_control_message(&mut self) {
if let Some(active_request) = self.request_queue.pop_front() {
match active_request {
ControlMessage::SetInterface(interface, response) => {
if self.discovery_socket.is_some() {
info!(
"Discovery socket currently bound to interface: {}, unbinding.",
self.interface
);
self.unbind_discovery().await;
self.request_queue
.push_front(ControlMessage::SetInterface(interface, response));
return;
}
if self.interface != interface {
self.set_interface(interface);
self.request_queue
.push_front(ControlMessage::SetInterface(interface, response));
return;
}
info!("Binding to interface: {}", interface);
let bind_result = self.bind_discovery();
match &bind_result {
Ok(()) => {
info!("Successfully Bound to interface: {}", interface);
}
Err(e) => {
warn!("Failed to bind to interface: {}. Error: {:?}", interface, e);
}
}
if response.send(bind_result).is_err() {
warn!("SetInterface response receiver dropped (caller canceled)");
}
}
ControlMessage::BindDiscovery(response) => {
let result = self.bind_discovery();
if response.send(result).is_err() {
warn!("BindDiscovery response receiver dropped (caller canceled)");
}
}
ControlMessage::UnbindDiscovery(response) => {
self.unbind_discovery().await;
if response.send(Ok(())).is_err() {
warn!("UnbindDiscovery response receiver dropped (caller canceled)");
}
}
ControlMessage::SendSD(target, header, response) => {
match &mut self.discovery_socket {
None => {
match self.bind_discovery() {
Ok(()) => {
self.request_queue.push_front(ControlMessage::SendSD(
target, header, response,
));
}
Err(e) => {
error!(
"Failed to bind discovery socket for sending SD message: {:?}",
e
);
if response.send(Err(e)).is_err() {
warn!(
"SendSD error response receiver dropped (caller canceled)"
);
}
}
}
}
Some(discovery_socket) => {
let message = Message::<PayloadDefinitions>::new_sd(
u32::from(discovery_socket.session_id()),
&header,
);
debug!("Sending {:?} to {}", &message, target);
let send_result = self
.discovery_socket
.as_mut()
.unwrap()
.send(target, message)
.await;
if response.send(send_result).is_err() {
warn!("SendSD response receiver dropped (caller canceled)");
}
}
}
}
ControlMessage::AddEndpoint(
service_id,
instance_id,
addr,
local_port,
response,
) => {
self.service_registry.insert(
ServiceInstanceId {
service_id,
instance_id,
},
ServiceEndpointInfo {
addr,
local_port,
major_version: 0xFF,
minor_version: 0xFFFF_FFFF,
},
);
debug!(
"Added endpoint for service 0x{:04X}.0x{:04X} -> {}",
service_id, instance_id, addr,
);
if response.send(Ok(())).is_err() {
warn!("AddEndpoint response receiver dropped (caller canceled)");
}
}
ControlMessage::RemoveEndpoint(service_id, instance_id, response) => {
self.service_registry.remove(ServiceInstanceId {
service_id,
instance_id,
});
debug!(
"Removed endpoint for service 0x{:04X}.0x{:04X}",
service_id, instance_id,
);
if response.send(Ok(())).is_err() {
warn!("RemoveEndpoint response receiver dropped (caller canceled)");
}
}
ControlMessage::SendToService {
service_id,
instance_id,
mut message,
send_complete,
response,
} => {
let id = ServiceInstanceId {
service_id,
instance_id,
};
let Some(endpoint) = self.service_registry.get(id) else {
let _ = send_complete.send(Err(Error::ServiceNotFound));
return;
};
let target = endpoint.addr;
let desired_port = endpoint.local_port;
let source_port = if desired_port == 0 {
if self.unicast_sockets.is_empty() {
match self.bind_unicast(0) {
Ok(port) => {
debug!("Auto-bound unicast on port {} for SendToService", port);
port
}
Err(e) => {
let _ = send_complete.send(Err(e));
return;
}
}
} else {
*self.unicast_sockets.keys().next().unwrap()
}
} else {
match self.bind_unicast(desired_port) {
Ok(port) => port,
Err(e) => {
let _ = send_complete.send(Err(e));
return;
}
}
};
let socket = self.unicast_sockets.get_mut(&source_port).unwrap();
let request_id =
(u32::from(self.client_id) << 16) | u32::from(self.session_counter);
message.set_request_id(request_id);
self.session_counter = self.session_counter.wrapping_add(1);
if self.session_counter == 0 {
self.session_counter = 1;
}
let send_result = socket.send(target, message).await;
match send_result {
Ok(()) => {
let _ = send_complete.send(Ok(()));
self.pending_responses.insert(request_id, response);
}
Err(e) => {
let _ = send_complete.send(Err(e));
}
}
}
#[cfg(test)]
ControlMessage::ForceSdSessionWrappedForTest(wrapped, response) => {
self.sd_session_has_wrapped = wrapped;
let _ = response.send(());
}
ControlMessage::QueryRebootFlag(response) => {
let flag = if let Some(socket) = self.discovery_socket.as_ref() {
socket.reboot_flag()
} else {
crate::protocol::sd::RebootFlag::from(!self.sd_session_has_wrapped)
};
if response.send(flag).is_err() {
warn!("QueryRebootFlag response receiver dropped (caller canceled)");
}
}
ControlMessage::Subscribe {
service_id,
instance_id,
major_version,
ttl,
event_group_id,
client_port,
response,
} => {
let id = ServiceInstanceId {
service_id,
instance_id,
};
if self.service_registry.get(id).is_none() {
let _ = response.send(Err(Error::ServiceNotFound));
return;
}
let unicast_port = match self.bind_unicast(client_port) {
Ok(port) => {
debug!("Bound unicast on port {} for Subscribe", port);
port
}
Err(e) => {
let _ = response.send(Err(e));
return;
}
};
match &mut self.discovery_socket {
None => match self.bind_discovery() {
Ok(()) => {
self.request_queue.push_front(ControlMessage::Subscribe {
service_id,
instance_id,
major_version,
ttl,
event_group_id,
client_port,
response,
});
}
Err(e) => {
let _ = response.send(Err(e));
}
},
Some(discovery_socket) => {
let sd_header = PayloadDefinitions::new_subscription_sd_header(
service_id,
instance_id,
major_version,
ttl,
event_group_id,
self.interface,
crate::protocol::sd::TransportProtocol::Udp,
unicast_port,
discovery_socket.reboot_flag(),
);
let session_id = u32::from(discovery_socket.session_id());
let message =
Message::<PayloadDefinitions>::new_sd(session_id, &sd_header);
let reg = self.service_registry.get(id).unwrap();
let target =
SocketAddrV4::new(*reg.addr.ip(), protocol::sd::MULTICAST_PORT);
debug!("Sending Subscribe {:?} to {}", &message, target);
let send_result = self
.discovery_socket
.as_mut()
.unwrap()
.send(target, message)
.await;
if response.send(send_result).is_err() {
warn!("Subscribe response receiver dropped (caller canceled)");
}
}
}
}
}
}
}
#[allow(clippy::too_many_lines)]
fn run(mut self) {
tokio::spawn(async move {
info!("SOME/IP Client processing loop started");
loop {
let Self {
control_receiver,
pending_responses,
discovery_socket,
discovery_unicast_socket,
unicast_sockets,
update_sender,
request_queue,
session_tracker,
service_registry,
e2e_registry,
run,
..
} = &mut self;
select! {
() = tokio::time::sleep(std::time::Duration::from_millis(125)) => {}
ctrl = control_receiver.recv() => {
if let Some(ctrl) = ctrl {
debug!("Received control message: {:?}", ctrl);
request_queue.push_back(ctrl);
} else {
*run = false;
}
}
discovery = Inner::receive_discovery(discovery_socket) => {
trace!("Received discovery message: {:?}", discovery);
match discovery {
Ok((source, someip_header, sd_header)) => {
process_discovery::<PayloadDefinitions>(
source,
TransportKind::Multicast,
someip_header,
sd_header,
session_tracker,
service_registry,
e2e_registry,
update_sender,
);
}
Err(err) => {
error!("Error receiving discovery message: {:?}", err);
let _ = update_sender.send(ClientUpdate::Error(err));
}
}
}
unicast_discovery = Inner::receive_discovery(discovery_unicast_socket) => {
trace!("Received unicast discovery message: {:?}", unicast_discovery);
match unicast_discovery {
Ok((source, someip_header, sd_header)) => {
process_discovery::<PayloadDefinitions>(
source,
TransportKind::Unicast,
someip_header,
sd_header,
session_tracker,
service_registry,
e2e_registry,
update_sender,
);
}
Err(err) => {
error!("Error receiving unicast discovery message: {:?}", err);
let _ = update_sender.send(ClientUpdate::Error(err));
}
}
}
unicast = Inner::receive_any_unicast(unicast_sockets) => {
trace!("Received unicast message: {:?}", unicast);
match unicast {
Ok(received) => {
let ReceivedMessage { message: received_message, e2e_status, source } = received;
let request_id = received_message.header().request_id();
if let Some(sender) = pending_responses.remove(&request_id) {
let _ = sender.send(Ok(received_message.payload().clone()));
continue;
}
let _ = update_sender.send(ClientUpdate::Unicast { message: received_message, e2e_status, source });
}
Err(err) => {
let _ = update_sender.send(ClientUpdate::Error(err));
}
}
}
}
if !*run {
info!("SOME/IP Client processing loop exiting");
break;
}
self.handle_control_message().await;
}
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::sd::test_support::{TestPayload, empty_sd_header};
use std::format;
type TestControl = ControlMessage<TestPayload>;
#[test]
fn test_control_message_constructors() {
let (_rx, msg) = TestControl::set_interface(Ipv4Addr::LOCALHOST);
assert!(matches!(msg, ControlMessage::SetInterface(..)));
let (_rx, msg) = TestControl::bind_discovery();
assert!(matches!(msg, ControlMessage::BindDiscovery(..)));
let (_rx, msg) = TestControl::unbind_discovery();
assert!(matches!(msg, ControlMessage::UnbindDiscovery(..)));
let target = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1234);
let sd_header = empty_sd_header();
let (_rx, msg) = TestControl::send_sd(target, sd_header);
assert!(matches!(msg, ControlMessage::SendSD(..)));
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (_rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
assert!(matches!(msg, ControlMessage::AddEndpoint(..)));
let (_rx, msg) = TestControl::remove_endpoint(0x1234, 0x0001);
assert!(matches!(msg, ControlMessage::RemoveEndpoint(..)));
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (_send_rx, _resp_rx, msg) = TestControl::send_to_service(0x1234, 0x0001, message);
assert!(matches!(msg, ControlMessage::SendToService { .. }));
let (_rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x01, 0);
assert!(matches!(msg, ControlMessage::Subscribe { .. }));
}
#[test]
fn test_control_message_debug() {
let (_rx, msg) = TestControl::set_interface(Ipv4Addr::LOCALHOST);
let s = format!("{msg:?}");
assert!(s.contains("SetInterface"));
let (_rx, msg) = TestControl::bind_discovery();
assert!(!format!("{msg:?}").is_empty());
let (_rx, msg) = TestControl::unbind_discovery();
assert!(format!("{msg:?}").contains("UnbindDiscovery"));
let target = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1234);
let sd_header = empty_sd_header();
let (_rx, msg) = TestControl::send_sd(target, sd_header);
assert!(format!("{msg:?}").contains("SendSD"));
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (_rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
let s = format!("{msg:?}");
assert!(s.contains("AddEndpoint"));
let (_rx, msg) = TestControl::remove_endpoint(0x1234, 0x0001);
let s = format!("{msg:?}");
assert!(s.contains("RemoveEndpoint"));
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (_send_rx, _resp_rx, msg) = TestControl::send_to_service(0x1234, 0x0001, message);
let s = format!("{msg:?}");
assert!(s.contains("SendToService"));
assert!(s.contains("service_id"));
assert!(s.contains("instance_id"));
let (_rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x01, 0);
let s = format!("{msg:?}");
assert!(s.contains("Subscribe"));
assert!(s.contains("service_id"));
assert!(s.contains("event_group_id"));
}
#[tokio::test]
async fn test_inner_spawn_and_shutdown() {
let (control_sender, mut update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
drop(control_sender);
let result =
tokio::time::timeout(std::time::Duration::from_secs(2), update_receiver.recv()).await;
assert!(result.is_ok());
assert!(result.unwrap().is_none());
}
async fn assert_inner_alive(control_sender: &Sender<ControlMessage<TestPayload>>) {
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
let (rx, msg) = TestControl::add_endpoint(0xFFFE, 0xFFFE, addr, 0);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out — inner loop appears dead")
.expect("Oneshot closed — inner loop appears dead");
assert!(result.is_ok());
}
#[tokio::test]
async fn test_dropped_receiver_bind_discovery_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_dropped_receiver_unbind_discovery_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::unbind_discovery();
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_dropped_receiver_set_interface_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::set_interface(Ipv4Addr::LOCALHOST);
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_dropped_receiver_send_sd_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let target = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 30490);
let sd_header = empty_sd_header();
let (rx, msg) = TestControl::send_sd(target, sd_header);
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_queued_messages_all_complete() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx_set, msg_set) = TestControl::set_interface(Ipv4Addr::LOCALHOST);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 9999);
let (rx_add, msg_add) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg_set).await.unwrap();
control_sender.send(msg_add).await.unwrap();
let set_result = tokio::time::timeout(std::time::Duration::from_secs(3), rx_set)
.await
.expect("Timed out waiting for SetInterface")
.expect("SetInterface oneshot closed");
assert!(set_result.is_ok());
let add_result = tokio::time::timeout(std::time::Duration::from_secs(3), rx_add)
.await
.expect("Timed out waiting for AddEndpoint")
.expect("AddEndpoint oneshot closed");
assert!(add_result.is_ok());
assert_inner_alive(&control_sender).await;
}
#[test]
fn test_send_to_service_constructor_returns_two_receivers() {
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (send_rx, resp_rx, _msg) = TestControl::send_to_service(0x1234, 0x0001, message);
if let ControlMessage::SendToService {
send_complete,
response,
..
} = _msg
{
send_complete.send(Ok(())).unwrap();
assert!(send_rx.blocking_recv().unwrap().is_ok());
let payload = TestPayload {
header: empty_sd_header(),
};
response.send(Ok(payload.clone())).unwrap();
assert_eq!(resp_rx.blocking_recv().unwrap().unwrap(), payload);
} else {
panic!("expected SendToService variant");
}
}
#[tokio::test]
async fn test_dropped_receiver_add_endpoint_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_dropped_receiver_remove_endpoint_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::remove_endpoint(0x1234, 0x0001);
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_dropped_receiver_send_to_service_send_complete_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (send_rx, _resp_rx, msg) = TestControl::send_to_service(0x1234, 0x0001, message);
drop(send_rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_bind_discovery_with_loopback() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
true,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
}
#[tokio::test]
async fn test_bind_discovery_idempotent() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
}
#[tokio::test]
async fn test_send_sd_auto_binds_discovery() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let target = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 30490);
let sd_header = empty_sd_header();
let (rx, msg) = TestControl::send_sd(target, sd_header);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out waiting for SendSD")
.expect("SendSD oneshot closed");
assert!(result.is_ok());
}
#[tokio::test]
async fn test_send_to_service_auto_binds_unicast() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (send_rx, _resp_rx, msg) = TestControl::send_to_service(0x1234, 0x0001, message);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), send_rx)
.await
.expect("Timed out waiting for SendToService")
.expect("SendToService oneshot closed");
assert!(result.is_ok(), "send should succeed: {result:?}");
}
#[tokio::test]
async fn test_subscribe_with_endpoint_sends_sd() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x01, 0);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out waiting for Subscribe")
.expect("Subscribe oneshot closed");
assert!(result.is_ok(), "subscribe should succeed: {result:?}");
}
#[tokio::test]
async fn test_subscribe_auto_binds_discovery() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x01, 0);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out waiting for Subscribe")
.expect("Subscribe oneshot closed");
assert!(result.is_ok(), "subscribe should auto-bind: {result:?}");
}
#[tokio::test]
async fn test_subscribe_unknown_service_returns_error() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::subscribe(0xFFFF, 0xFFFF, 1, 3, 0x01, 0);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out")
.expect("oneshot closed");
assert!(matches!(result, Err(Error::ServiceNotFound)));
}
#[tokio::test]
async fn test_send_to_service_reuses_existing_unicast_socket() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (send_rx, _resp_rx, msg) = TestControl::send_to_service(0x1234, 0x0001, message);
control_sender.send(msg).await.unwrap();
send_rx.await.unwrap().unwrap();
let message = Message::<TestPayload>::new_sd(1, &empty_sd_header());
let (send_rx, _resp_rx, msg) = TestControl::send_to_service(0x1234, 0x0001, message);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), send_rx)
.await
.expect("Timed out")
.expect("oneshot closed");
assert!(
result.is_ok(),
"second send should reuse socket: {result:?}"
);
}
#[tokio::test]
async fn test_dropped_receiver_subscribe_service_not_found_continues() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x01, 0);
drop(rx);
control_sender.send(msg).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_set_interface_changes_interface() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::set_interface(Ipv4Addr::new(127, 0, 0, 2));
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(3), rx)
.await
.expect("Timed out waiting for SetInterface")
.expect("SetInterface oneshot closed");
let _ = result;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_set_interface_with_discovery_bound_changes_interface() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::set_interface(Ipv4Addr::new(127, 0, 0, 2));
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(3), rx)
.await
.expect("Timed out waiting for SetInterface")
.expect("SetInterface oneshot closed");
let _ = result;
assert_inner_alive(&control_sender).await;
}
#[tokio::test]
async fn test_subscribe_specific_port_reuse() {
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000);
let (rx, msg) = TestControl::add_endpoint(0x1234, 0x0001, addr, 0);
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x01, 44444);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out")
.expect("oneshot closed");
assert!(result.is_ok(), "first subscribe should succeed: {result:?}");
let (rx, msg) = TestControl::subscribe(0x1234, 0x0001, 1, 3, 0x02, 44444);
control_sender.send(msg).await.unwrap();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), rx)
.await
.expect("Timed out")
.expect("oneshot closed");
assert!(
result.is_ok(),
"second subscribe should reuse port: {result:?}"
);
}
#[tokio::test]
async fn test_sd_session_id_persists_across_rebind() {
use crate::protocol::MessageView;
use std::vec;
use tokio::net::UdpSocket;
let (control_sender, _update_receiver) = Inner::<TestPayload>::spawn(
Ipv4Addr::LOCALHOST,
Arc::new(Mutex::new(E2ERegistry::new())),
false,
);
let raw = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let target = SocketAddrV4::new(Ipv4Addr::LOCALHOST, raw.local_addr().unwrap().port());
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::send_sd(target, empty_sd_header());
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let mut buf = vec![0u8; 1400];
let (len, _) =
tokio::time::timeout(std::time::Duration::from_secs(2), raw.recv_from(&mut buf))
.await
.expect("timed out waiting for first SD message")
.unwrap();
let first = MessageView::parse(&buf[..len]).unwrap();
let session_id_before = (first.header().request_id() & 0xFFFF) as u16;
let reboot_flag_before = first.sd_header().unwrap().flags().reboot();
assert!(session_id_before >= 1, "session_id must never be 0");
let (rx, msg) = TestControl::unbind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::bind_discovery();
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (rx, msg) = TestControl::send_sd(target, empty_sd_header());
control_sender.send(msg).await.unwrap();
rx.await.unwrap().unwrap();
let (len, _) =
tokio::time::timeout(std::time::Duration::from_secs(2), raw.recv_from(&mut buf))
.await
.expect("timed out waiting for second SD message")
.unwrap();
let second = MessageView::parse(&buf[..len]).unwrap();
let session_id_after = (second.header().request_id() & 0xFFFF) as u16;
let reboot_flag_after = second.sd_header().unwrap().flags().reboot();
assert!(
session_id_after > session_id_before,
"session_id should continue after rebind (before={session_id_before}, after={session_id_after})"
);
assert_eq!(
reboot_flag_after, reboot_flag_before,
"reboot_flag should be preserved across rebind"
);
}
}