use crate::{
Error, TCP_PORT,
logical_address::LogicalAddress,
message_codec::MessageCodec,
messages::{
Decode, DiagnosticMessage, DiagnosticPowerModeCode, Encode, FurtherActionRequired, Message,
OwnedMessage, OwnedPayload, Payload, PayloadType, ProtocolVersion,
RoutingActivationRequest, RoutingActivationResponseCode, VehicleIdentificationResponse,
VinGidSyncStatus,
},
};
use async_trait::async_trait;
use futures::{SinkExt, StreamExt};
use std::{
boxed::Box,
net::{IpAddr, SocketAddr},
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use tokio::{
net::{TcpListener, TcpStream, UdpSocket},
time::sleep,
};
use tokio_util::codec::{FramedRead, FramedWrite};
use tracing::{debug, error, warn};
const ACCEPT_ERROR_BACKOFF: Duration = Duration::from_millis(100);
#[derive(Debug)]
pub struct ClientConnectionInfo {
pub ip_address: IpAddr,
pub logical_address: LogicalAddress,
}
struct ActiveConnectionGuard<'a> {
active_connections: &'a AtomicUsize,
}
impl<'a> ActiveConnectionGuard<'a> {
fn new(active_connections: &'a AtomicUsize) -> Self {
active_connections.fetch_add(1, Ordering::Relaxed);
Self { active_connections }
}
}
impl Drop for ActiveConnectionGuard<'_> {
fn drop(&mut self) {
self.active_connections.fetch_sub(1, Ordering::Relaxed);
}
}
#[async_trait]
pub trait ResponseWriter: Send {
async fn send(&mut self, message: OwnedMessage) -> Result<(), Error>;
}
struct FramedResponseWriter<'a, W> {
sink: &'a mut FramedWrite<W, MessageCodec>,
}
#[async_trait]
impl<W> ResponseWriter for FramedResponseWriter<'_, W>
where
W: tokio::io::AsyncWrite + Unpin + Send,
{
async fn send(&mut self, message: OwnedMessage) -> Result<(), Error> {
self.sink.send(&message).await?;
Ok(())
}
}
#[async_trait]
pub trait ServerConnectionHandler {
fn get_vin(&self) -> [u8; 17];
fn get_logical_address(&self) -> LogicalAddress;
fn get_entity_id(&self) -> [u8; 6];
fn get_group_id(&self) -> Option<[u8; 6]>;
async fn routing_activation(
&self,
request: &RoutingActivationRequest,
) -> Result<OwnedMessage, Error>;
async fn diagnostic_message(
&self,
message: &DiagnosticMessage<'_>,
responses: &mut dyn ResponseWriter,
) -> Result<(), Error>;
fn received_vehicle_identification_request(
&self,
_client_info: &ClientConnectionInfo,
) -> Result<VehicleIdentificationResponse, Error> {
Ok(VehicleIdentificationResponse {
entity_id: self.get_entity_id(),
logical_address: self.get_logical_address(),
vin: self.get_vin(),
group_id: self.get_group_id(),
further_action: FurtherActionRequired::NoFurtherActionRequired,
vin_gid_sync_status: VinGidSyncStatus::Synchronized,
})
}
fn vehicle_identification_with_eid(
&self,
_client_info: &ClientConnectionInfo,
eid: &[u8; 6],
) -> Result<Option<VehicleIdentificationResponse>, Error> {
if self.get_entity_id() == *eid {
Ok(Some(VehicleIdentificationResponse {
entity_id: self.get_entity_id(),
logical_address: self.get_logical_address(),
vin: self.get_vin(),
group_id: self.get_group_id(),
further_action: FurtherActionRequired::NoFurtherActionRequired,
vin_gid_sync_status: VinGidSyncStatus::Synchronized,
}))
} else {
Ok(None)
}
}
fn vehicle_identification_with_vin(
&self,
_client_info: &ClientConnectionInfo,
vin: &[u8; 17],
) -> Result<Option<VehicleIdentificationResponse>, Error> {
if self.get_vin() == *vin {
Ok(Some(VehicleIdentificationResponse {
entity_id: self.get_entity_id(),
logical_address: self.get_logical_address(),
vin: self.get_vin(),
group_id: self.get_group_id(),
further_action: FurtherActionRequired::NoFurtherActionRequired,
vin_gid_sync_status: VinGidSyncStatus::Synchronized,
}))
} else {
Ok(None)
}
}
async fn alive_check(&self, client_info: &ClientConnectionInfo) -> Result<OwnedMessage, Error> {
Ok(OwnedMessage::alive_check_response(
self.protocol_version(),
client_info.logical_address,
))
}
async fn diagnostic_power_mode_information(
&self,
_client_info: &ClientConnectionInfo,
) -> Result<DiagnosticPowerModeCode, Error> {
Ok(DiagnosticPowerModeCode::NotSupported)
}
fn protocol_version(&self) -> ProtocolVersion {
ProtocolVersion::V2012
}
}
#[derive(Debug)]
pub struct Server<T> {
connection_handler: Arc<T>,
active_connections: AtomicUsize,
}
impl<T> Server<T>
where
T: ServerConnectionHandler + Sync,
{
pub fn new(connection_handler: T) -> Result<Self, Error> {
Ok(Server {
connection_handler: Arc::new(connection_handler),
active_connections: AtomicUsize::new(0),
})
}
pub async fn run_server(&self) -> Result<(), Error> {
let tcp_listener = TcpListener::bind(("0.0.0.0", TCP_PORT)).await?;
self.run_server_with_listener(tcp_listener).await
}
pub async fn run_server_with_listener(&self, tcp_listener: TcpListener) -> Result<(), Error> {
loop {
match tcp_listener.accept().await {
Ok((tcp_stream, client_socket_addr)) => {
if let Err(client_error) = self
.handle_client_connection(client_socket_addr, tcp_stream)
.await
{
error!("Client error: {client_error}");
}
}
Err(accept_error) => {
error!("Failed to accept TCP client, continuing: {accept_error}");
sleep(ACCEPT_ERROR_BACKOFF).await;
}
}
}
}
pub async fn run_udp_responder(&self, socket: UdpSocket) -> Result<(), Error> {
let mut buf = [0u8; 1024];
loop {
let (len, peer) = match socket.recv_from(&mut buf).await {
Ok(received) => received,
Err(recv_error) => {
warn!("UDP receive failed, continuing: {recv_error}");
sleep(ACCEPT_ERROR_BACKOFF).await;
continue;
}
};
let (message, _rest) = match Message::decode(&buf[..len]) {
Ok(decoded) => decoded,
Err(decode_error) => {
warn!("Undecodable UDP datagram from {peer}, ignoring: {decode_error}");
continue;
}
};
if !matches!(message.payload, Payload::VehicleIdentificationRequest) {
debug!(
"Unsupported UDP payload type {:?} from {peer}, ignoring",
message.header.payload_type
);
continue;
}
if !matches!(
message.header.payload_type,
PayloadType::VehicleIdentificationRequest
) {
debug!(
"Ignoring directed identification request {:?} from {peer}: this crate \
cannot match the EID/VIN it names",
message.header.payload_type
);
continue;
}
let client_info = ClientConnectionInfo {
ip_address: peer.ip(),
logical_address: LogicalAddress(0x0000),
};
let response = match self
.connection_handler
.received_vehicle_identification_request(&client_info)
{
Ok(response) => response,
Err(handler_error) => {
warn!("Identification handler failed for {peer}, skipping: {handler_error}");
continue;
}
};
let reply = OwnedMessage::vehicle_identification_response(
self.connection_handler.protocol_version(),
response,
);
let mut encoded = match reply.encoded_size() {
Ok(size) => std::vec::Vec::with_capacity(size),
Err(size_error) => {
warn!("Failed to size identification response for {peer}: {size_error}");
continue;
}
};
if let Err(encode_error) = reply.encode(&mut encoded) {
warn!("Failed to encode identification response for {peer}: {encode_error}");
continue;
}
if let Err(send_error) = socket.send_to(&encoded, peer).await {
warn!("Failed to answer identification probe from {peer}: {send_error}");
}
}
}
pub async fn handle_client_connection(
&self,
client_socket_addr: SocketAddr,
tcp_stream: TcpStream,
) -> Result<(), Error> {
let _active_connection_guard = ActiveConnectionGuard::new(&self.active_connections);
if let Err(nodelay_error) = tcp_stream.set_nodelay(true) {
warn!("Failed to set TCP_NODELAY for {client_socket_addr}: {nodelay_error}");
}
let (rx, tx) = tcp_stream.into_split();
let mut read_stream = FramedRead::new(rx, MessageCodec::new());
let mut write_sink = FramedWrite::new(tx, MessageCodec::new());
let mut tester_logical_address: Option<LogicalAddress> = None;
loop {
match read_stream.next().await {
Some(Ok(message)) => {
if let Some(response) = self
.handle_client_message(
client_socket_addr,
message,
&mut write_sink,
&mut tester_logical_address,
)
.await?
{
write_sink.send(&response).await?;
}
}
Some(Err(codec_error)) => {
error!(
"Client decoding error, closing connection. source: {client_socket_addr}, {codec_error}"
);
return Ok(());
}
None => {
warn!("Client stream closed, client addr: {client_socket_addr}");
return Ok(());
}
}
}
}
async fn handle_client_message<W>(
&self,
client_socket_addr: SocketAddr,
request_message: OwnedMessage,
write_sink: &mut FramedWrite<W, MessageCodec>,
tester_logical_address: &mut Option<LogicalAddress>,
) -> Result<Option<OwnedMessage>, Error>
where
W: tokio::io::AsyncWrite + Unpin + Send,
{
let connection_info = ClientConnectionInfo {
ip_address: client_socket_addr.ip(),
logical_address: tester_logical_address.unwrap_or(LogicalAddress(0x0000)),
};
match request_message.payload {
OwnedPayload::AliveCheckRequest => self
.connection_handler
.alive_check(&connection_info)
.await
.map(Some),
OwnedPayload::DiagnosticMessage(diagnostic_message) => {
let mut responses = FramedResponseWriter { sink: write_sink };
self.connection_handler
.diagnostic_message(&diagnostic_message.as_ref(), &mut responses)
.await?;
Ok(None)
}
OwnedPayload::EntityStatusRequest => {
warn!(
"Entity Status Request is not yet supported, ignoring. source: {client_socket_addr}"
);
Ok(None)
}
OwnedPayload::RoutingActivationRequest(request) => {
let response = self.connection_handler.routing_activation(&request).await?;
if let OwnedPayload::RoutingActivationResponse(activation_response) =
&response.payload
&& matches!(
activation_response.routing_activation_response_code,
RoutingActivationResponseCode::RoutingSuccessfullyActivated
| RoutingActivationResponseCode::RoutingSuccessfullyActivatedConfirmationRequired
)
{
*tester_logical_address = Some(request.source_address);
}
Ok(Some(response))
}
OwnedPayload::RoutingActivationResponse(_routing_activation_response) => {
warn!(
"Client sent a server-role RoutingActivationResponse message, source: {client_socket_addr}"
);
Err(Error::UnexpectedMessageType(
request_message.header.payload_type,
))
}
OwnedPayload::VehicleIdentificationRequest => {
warn!(
"Vehicle Identification Request is not yet supported, ignoring. source: {client_socket_addr}"
);
Ok(None)
}
OwnedPayload::VehicleIdentificationResponse(_vehicle_identification_response) => {
warn!(
"Client sent a server-role VehicleIdentificationResponse message, source: {client_socket_addr}"
);
Err(Error::UnexpectedMessageType(
request_message.header.payload_type,
))
}
_ => Err(Error::UnexpectedMessageType(
request_message.header.payload_type,
)),
}
}
}
#[cfg(test)]
mod tests {
use super::ActiveConnectionGuard;
use std::sync::atomic::{AtomicUsize, Ordering};
#[test]
fn guard_increments_on_creation_and_decrements_on_drop() {
let active_connections = AtomicUsize::new(0);
{
let _guard = ActiveConnectionGuard::new(&active_connections);
assert_eq!(active_connections.load(Ordering::Relaxed), 1);
}
assert_eq!(active_connections.load(Ordering::Relaxed), 0);
}
#[test]
fn guard_decrements_even_on_early_return_via_question_mark() {
fn inner(active_connections: &AtomicUsize) -> Result<(), ()> {
let _guard = ActiveConnectionGuard::new(active_connections);
Err(())?;
Ok(())
}
let active_connections = AtomicUsize::new(0);
let result = inner(&active_connections);
assert!(result.is_err());
assert_eq!(active_connections.load(Ordering::Relaxed), 0);
}
#[test]
fn guard_decrements_on_panic_unwind() {
let active_connections = AtomicUsize::new(0);
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _guard = ActiveConnectionGuard::new(&active_connections);
panic!("simulated failure while holding the guard");
}));
assert!(result.is_err());
assert_eq!(active_connections.load(Ordering::Relaxed), 0);
}
}