mod error;
mod inner;
mod service_registry;
mod session;
mod socket_manager;
pub use error::Error;
use crate::e2e::{E2ECheckStatus, E2EKey, E2EProfile, E2ERegistry};
use crate::{protocol, protocol::Message, traits::PayloadWireFormat};
use inner::{ControlMessage, Inner};
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::{Arc, Mutex, RwLock};
use tokio::sync::{mpsc, oneshot};
use tracing::info;
pub struct PendingResponse<P> {
receiver: oneshot::Receiver<Result<P, Error>>,
}
impl<P> std::fmt::Debug for PendingResponse<P> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PendingResponse").finish_non_exhaustive()
}
}
impl<P> PendingResponse<P> {
pub async fn response(self) -> Result<P, Error> {
self.receiver
.await
.expect("inner loop dropped response channel")
}
}
pub struct DiscoveryMessage<P: PayloadWireFormat> {
pub source: SocketAddr,
pub someip_header: protocol::Header,
pub sd_header: P::SdHeader,
}
impl<P: PayloadWireFormat> std::fmt::Debug for DiscoveryMessage<P> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DiscoveryMessage")
.field("source", &self.source)
.field("someip_header", &self.someip_header)
.field("sd_header", &self.sd_header)
.finish()
}
}
pub enum ClientUpdate<P: PayloadWireFormat> {
DiscoveryUpdated(DiscoveryMessage<P>),
SenderRebooted(SocketAddr),
Unicast {
message: Message<P>,
e2e_status: Option<E2ECheckStatus>,
},
Error(Error),
}
impl<P: PayloadWireFormat> std::fmt::Debug for ClientUpdate<P> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::DiscoveryUpdated(msg) => f.debug_tuple("DiscoveryUpdated").field(msg).finish(),
Self::SenderRebooted(addr) => f.debug_tuple("SenderRebooted").field(addr).finish(),
Self::Unicast {
message,
e2e_status,
} => f
.debug_struct("Unicast")
.field("message", message)
.field("e2e_status", e2e_status)
.finish(),
Self::Error(err) => f.debug_tuple("Error").field(err).finish(),
}
}
}
pub struct ClientUpdates<MessageDefinitions: PayloadWireFormat> {
update_receiver: mpsc::UnboundedReceiver<ClientUpdate<MessageDefinitions>>,
}
impl<MessageDefinitions: PayloadWireFormat> std::fmt::Debug for ClientUpdates<MessageDefinitions> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClientUpdates").finish_non_exhaustive()
}
}
impl<MessageDefinitions: PayloadWireFormat> ClientUpdates<MessageDefinitions> {
pub async fn recv(&mut self) -> Option<ClientUpdate<MessageDefinitions>> {
self.update_receiver.recv().await
}
}
#[derive(Clone)]
pub struct Client<MessageDefinitions: PayloadWireFormat> {
interface: Arc<RwLock<Ipv4Addr>>,
control_sender: mpsc::Sender<inner::ControlMessage<MessageDefinitions>>,
e2e_registry: Arc<Mutex<E2ERegistry>>,
}
impl<MessageDefinitions: PayloadWireFormat> std::fmt::Debug for Client<MessageDefinitions> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Client")
.field(
"interface",
&*self.interface.read().expect("interface lock poisoned"),
)
.finish_non_exhaustive()
}
}
impl<MessageDefinitions> Client<MessageDefinitions>
where
MessageDefinitions: PayloadWireFormat + Clone + std::fmt::Debug + 'static,
{
#[must_use]
pub fn new(interface: Ipv4Addr) -> (Self, ClientUpdates<MessageDefinitions>) {
let e2e_registry = Arc::new(Mutex::new(E2ERegistry::new()));
let (control_sender, update_receiver) = Inner::spawn(interface, Arc::clone(&e2e_registry));
let client = Self {
interface: Arc::new(RwLock::new(interface)),
control_sender,
e2e_registry,
};
let updates = ClientUpdates { update_receiver };
(client, updates)
}
#[must_use]
pub fn interface(&self) -> Ipv4Addr {
*self.interface.read().expect("interface lock poisoned")
}
pub async fn set_interface(&self, interface: Ipv4Addr) -> Result<(), Error> {
let (response, message) = ControlMessage::set_interface(interface);
self.control_sender.send(message).await.unwrap();
response.await.unwrap()?;
*self.interface.write().expect("interface lock poisoned") = interface;
Ok(())
}
pub async fn bind_discovery(&self) -> Result<(), Error> {
let (response, message) = ControlMessage::bind_discovery();
self.control_sender.send(message).await.unwrap();
response.await.unwrap()
}
pub async fn unbind_discovery(&self) -> Result<(), Error> {
let (response, message) = ControlMessage::unbind_discovery();
self.control_sender.send(message).await.unwrap();
response.await.unwrap()
}
pub async fn subscribe(
&self,
service_id: u16,
instance_id: u16,
major_version: u8,
ttl: u32,
event_group_id: u16,
client_port: u16,
) -> Result<(), Error> {
let (response, message) = ControlMessage::subscribe(
service_id,
instance_id,
major_version,
ttl,
event_group_id,
client_port,
);
self.control_sender.send(message).await.unwrap();
response.await.unwrap()
}
pub async fn subscribe_no_wait(
&self,
service_id: u16,
instance_id: u16,
major_version: u8,
ttl: u32,
event_group_id: u16,
client_port: u16,
) {
let (response, message) = ControlMessage::subscribe(
service_id,
instance_id,
major_version,
ttl,
event_group_id,
client_port,
);
let _ = self.control_sender.send(message).await;
tokio::spawn(async move {
let _ = response.await;
});
}
pub async fn send_sd_message(
&self,
target: SocketAddrV4,
sd_header: <MessageDefinitions as PayloadWireFormat>::SdHeader,
) -> Result<(), Error> {
let (response, message) = ControlMessage::send_sd(target, sd_header);
self.control_sender.send(message).await.unwrap();
response.await.unwrap()
}
pub fn start_sd_announcements(
&self,
sd_header: <MessageDefinitions as PayloadWireFormat>::SdHeader,
interval: std::time::Duration,
) -> tokio::task::JoinHandle<()>
where
<MessageDefinitions as PayloadWireFormat>::SdHeader: Send + 'static,
{
use crate::protocol::sd;
let weak_sender = self.control_sender.downgrade();
let target = SocketAddrV4::new(sd::MULTICAST_IP, sd::MULTICAST_PORT);
let interval = interval.max(std::time::Duration::from_millis(100));
tokio::spawn(async move {
let mut tick = tokio::time::interval(interval);
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
tick.tick().await;
let mut count = 0u64;
loop {
tick.tick().await;
let Some(sender) = weak_sender.upgrade() else {
tracing::info!("Client shut down, stopping SD announcements");
break;
};
let (response, message) = ControlMessage::send_sd(target, sd_header.clone());
let send_ok = sender.send(message).await.is_ok();
drop(sender);
if !send_ok {
tracing::warn!("SD announcement channel closed, stopping");
break;
}
match response.await {
Ok(Ok(())) => {
count += 1;
if count == 1 {
tracing::info!("Sent first client SD announcement");
} else {
tracing::trace!("Sent {count} client SD announcements");
}
}
Ok(Err(e)) => {
tracing::error!("Failed to send SD announcement: {e:?}");
}
Err(_) => {
tracing::warn!("SD announcement response dropped, stopping");
break;
}
}
}
})
}
pub async fn add_endpoint(
&self,
service_id: u16,
instance_id: u16,
addr: SocketAddrV4,
local_port: u16,
) -> Result<(), Error> {
let (response, message) =
ControlMessage::add_endpoint(service_id, instance_id, addr, local_port);
self.control_sender.send(message).await.unwrap();
response.await.unwrap()
}
pub async fn remove_endpoint(&self, service_id: u16, instance_id: u16) -> Result<(), Error> {
let (response, message) = ControlMessage::remove_endpoint(service_id, instance_id);
self.control_sender.send(message).await.unwrap();
response.await.unwrap()
}
pub async fn send_to_service(
&self,
service_id: u16,
instance_id: u16,
message: crate::protocol::Message<MessageDefinitions>,
) -> Result<PendingResponse<MessageDefinitions>, Error> {
let (send_rx, response_rx, ctrl_msg) =
ControlMessage::send_to_service(service_id, instance_id, message);
self.control_sender.send(ctrl_msg).await.unwrap();
send_rx.await.unwrap()?;
Ok(PendingResponse {
receiver: response_rx,
})
}
pub async fn request(
&self,
service_id: u16,
instance_id: u16,
message: crate::protocol::Message<MessageDefinitions>,
) -> Result<MessageDefinitions, Error> {
let (send_rx, response_rx, ctrl_msg) =
ControlMessage::send_to_service(service_id, instance_id, message);
self.control_sender.send(ctrl_msg).await.unwrap();
send_rx.await.unwrap()?;
response_rx
.await
.expect("inner loop dropped response channel")
}
pub fn register_e2e(&self, key: E2EKey, profile: E2EProfile) {
self.e2e_registry
.lock()
.expect("e2e registry lock poisoned")
.register(key, profile);
}
pub fn unregister_e2e(&self, key: &E2EKey) {
self.e2e_registry
.lock()
.expect("e2e registry lock poisoned")
.unregister(key);
}
pub fn shut_down(self) {
drop(self.control_sender);
info!("Shutting Down SOME/IP client");
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::sd::test_support::{TestPayload, empty_sd_header};
use crate::traits::WireFormat;
use std::format;
type TestClient = Client<TestPayload>;
#[tokio::test]
async fn test_client_new_and_interface() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
assert_eq!(client.interface(), Ipv4Addr::LOCALHOST);
client.shut_down();
}
#[tokio::test]
async fn test_client_debug() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let debug_str = format!("{client:?}");
assert!(debug_str.contains("Client"));
assert!(debug_str.contains("127.0.0.1"));
client.shut_down();
}
#[tokio::test]
async fn test_client_update_debug() {
use std::net::SocketAddr;
let sd_header = empty_sd_header();
let someip_header = crate::protocol::Header::new_sd(1, sd_header.required_size());
let discovery_msg = DiscoveryMessage {
source: SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 30490),
someip_header,
sd_header,
};
let update: ClientUpdate<TestPayload> = ClientUpdate::DiscoveryUpdated(discovery_msg);
let debug_str = format!("{update:?}");
assert!(debug_str.contains("DiscoveryUpdated"));
let update: ClientUpdate<TestPayload> =
ClientUpdate::SenderRebooted(SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 30490));
let debug_str = format!("{update:?}");
assert!(debug_str.contains("SenderRebooted"));
let msg = crate::protocol::Message::new_sd(1, &empty_sd_header());
let update: ClientUpdate<TestPayload> = ClientUpdate::Unicast {
message: msg,
e2e_status: None,
};
let debug_str = format!("{update:?}");
assert!(debug_str.contains("Unicast"));
let update: ClientUpdate<TestPayload> = ClientUpdate::Error(Error::ServiceNotFound);
let debug_str = format!("{update:?}");
assert!(debug_str.contains("Error"));
}
#[tokio::test]
async fn test_subscribe_unknown_service_returns_error() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let result = client.subscribe(0xFFFF, 0xFFFF, 1, 3, 0x01, 0).await;
assert!(
matches!(result, Err(Error::ServiceNotFound)),
"expected ServiceNotFound, got {result:?}"
);
client.shut_down();
}
#[tokio::test]
async fn test_subscribe_no_wait_unknown_service_does_not_panic() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client
.subscribe_no_wait(0xFFFF, 0xFFFF, 1, 3, 0x01, 0)
.await;
client.shut_down();
}
#[tokio::test]
async fn test_bind_discovery_and_unbind() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client.bind_discovery().await.unwrap();
client.unbind_discovery().await.unwrap();
client.shut_down();
}
#[tokio::test]
async fn test_set_interface() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let new_addr = Ipv4Addr::LOCALHOST;
client.set_interface(new_addr).await.unwrap();
assert_eq!(client.interface(), new_addr);
client.shut_down();
}
#[tokio::test]
async fn test_add_endpoint_succeeds() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let addr = SocketAddrV4::new(Ipv4Addr::new(192, 168, 1, 1), 30000);
client.add_endpoint(0x1234, 0x0001, addr, 0).await.unwrap();
client.shut_down();
}
#[tokio::test]
async fn test_send_to_service_unknown_returns_error() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let msg = crate::protocol::Message::new_sd(1, &empty_sd_header());
let result = client.send_to_service(0xFFFF, 0xFFFF, msg).await;
assert!(
matches!(result, Err(Error::ServiceNotFound)),
"expected ServiceNotFound, got {result:?}"
);
client.shut_down();
}
#[tokio::test]
async fn test_remove_endpoint_succeeds() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let addr = SocketAddrV4::new(Ipv4Addr::new(192, 168, 1, 1), 30000);
client.add_endpoint(0x1234, 0x0001, addr, 0).await.unwrap();
client.remove_endpoint(0x1234, 0x0001).await.unwrap();
client.shut_down();
}
#[test]
fn test_pending_response_debug() {
let (_tx, rx) = oneshot::channel::<Result<TestPayload, Error>>();
let pending = PendingResponse { receiver: rx };
let s = format!("{pending:?}");
assert!(s.contains("PendingResponse"));
}
#[tokio::test]
async fn test_pending_response_resolves_ok() {
let (tx, rx) = oneshot::channel::<Result<TestPayload, Error>>();
let pending = PendingResponse { receiver: rx };
let payload = TestPayload {
header: empty_sd_header(),
};
tx.send(Ok(payload.clone())).unwrap();
let result = pending.response().await;
assert_eq!(result.unwrap(), payload);
}
#[tokio::test]
async fn test_pending_response_resolves_err() {
let (tx, rx) = oneshot::channel::<Result<TestPayload, Error>>();
let pending = PendingResponse { receiver: rx };
tx.send(Err(Error::ServiceNotFound)).unwrap();
let result = pending.response().await;
assert!(
matches!(result, Err(Error::ServiceNotFound)),
"expected ServiceNotFound, got {result:?}"
);
}
#[tokio::test]
async fn test_send_sd_message() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client.bind_discovery().await.unwrap();
let target = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 30490);
let sd_header = empty_sd_header();
client.send_sd_message(target, sd_header).await.unwrap();
client.shut_down();
}
#[tokio::test]
async fn test_send_to_service_success_returns_pending_response() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let addr = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 30000);
client.add_endpoint(0x1234, 0x0001, addr, 0).await.unwrap();
let msg = crate::protocol::Message::new_sd(1, &empty_sd_header());
let pending = client.send_to_service(0x1234, 0x0001, msg).await;
assert!(pending.is_ok());
client.shut_down();
}
#[tokio::test]
async fn test_recv_returns_none_after_shutdown() {
let (client, mut updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client.shut_down();
let result = tokio::time::timeout(std::time::Duration::from_secs(2), updates.recv()).await;
assert!(result.is_ok());
assert!(result.unwrap().is_none());
}
#[tokio::test]
async fn test_register_and_unregister_e2e() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let key = E2EKey {
service_id: 0x1234,
method_or_event_id: 0x0001,
};
let profile = E2EProfile::Profile4(crate::e2e::Profile4Config::new(42, 10));
client.register_e2e(key, profile);
client.unregister_e2e(&key);
client.shut_down();
}
#[tokio::test]
async fn test_client_is_clone() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let client2 = client.clone();
assert_eq!(client.interface(), client2.interface());
client.shut_down();
}
#[tokio::test]
async fn test_client_updates_debug() {
let (_client, updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let debug_str = format!("{updates:?}");
assert!(debug_str.contains("ClientUpdates"));
}
#[tokio::test]
async fn test_request_unknown_service_returns_error() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let msg = crate::protocol::Message::new_sd(1, &empty_sd_header());
let result = client.request(0xFFFF, 0xFFFF, msg).await;
assert!(
matches!(result, Err(Error::ServiceNotFound)),
"expected ServiceNotFound, got {result:?}"
);
client.shut_down();
}
#[tokio::test]
async fn test_start_sd_announcements_does_not_panic() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client.bind_discovery().await.unwrap();
let sd_header = empty_sd_header();
let handle =
client.start_sd_announcements(sd_header, std::time::Duration::from_millis(100));
tokio::time::sleep(std::time::Duration::from_millis(250)).await;
handle.abort();
let result = handle.await;
let err = result.unwrap_err();
assert!(
err.is_cancelled(),
"task should have been cancelled, not panicked"
);
client.shut_down();
}
#[tokio::test]
async fn test_start_sd_announcements_without_discovery_bound() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
let sd_header = empty_sd_header();
let handle =
client.start_sd_announcements(sd_header, std::time::Duration::from_millis(100));
tokio::time::sleep(std::time::Duration::from_millis(250)).await;
handle.abort();
let result = handle.await;
let err = result.unwrap_err();
assert!(
err.is_cancelled(),
"task should have been cancelled, not panicked"
);
client.shut_down();
}
#[tokio::test]
async fn test_start_sd_announcements_abort_stops_task() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client.bind_discovery().await.unwrap();
let sd_header = empty_sd_header();
let handle =
client.start_sd_announcements(sd_header, std::time::Duration::from_millis(100));
handle.abort();
let result = handle.await;
let err = result.unwrap_err();
assert!(
err.is_cancelled(),
"task should have been cancelled, not panicked"
);
client.shut_down();
}
#[tokio::test]
async fn test_start_sd_announcements_stops_on_shutdown() {
let (client, _updates) = TestClient::new(Ipv4Addr::LOCALHOST);
client.bind_discovery().await.unwrap();
let sd_header = empty_sd_header();
let handle =
client.start_sd_announcements(sd_header, std::time::Duration::from_millis(100));
client.shut_down();
let join_result = tokio::time::timeout(std::time::Duration::from_secs(2), handle)
.await
.expect("task should have exited within timeout");
assert!(
join_result.is_ok() || join_result.as_ref().unwrap_err().is_cancelled(),
"task should have exited cleanly, not panicked"
);
}
}