use std::{io, sync::Arc, time::Duration};
use bytes::BytesMut;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream, ToSocketAddrs, tcp::OwnedReadHalf},
sync::Mutex,
task::block_in_place,
time::timeout,
};
use tokio_util::codec::Decoder;
use tokio_util::sync::CancellationToken;
use crate::XvcServer;
use xvc_protocol::{
Message, OwnedMessage, Version, XvcInfo, error::ReadError, tokio_codec::MessageDecoder,
};
#[derive(Debug, Clone)]
pub struct Config {
pub max_vector_size: u32,
pub read_write_timeout: Duration,
}
impl Default for Config {
fn default() -> Self {
Self {
max_vector_size: 10 * 1024 * 1024,
read_write_timeout: Duration::from_secs(30),
}
}
}
#[derive(Debug)]
pub struct Server<T: XvcServer> {
server: Arc<Mutex<T>>,
config: Config,
}
#[derive(Default)]
pub struct Builder {
config: Config,
}
impl Builder {
pub fn new() -> Builder {
Builder::default()
}
pub fn max_vector_size(mut self, size: u32) -> Self {
self.config.max_vector_size = size;
self
}
pub fn rw_timeout(mut self, timeout: Duration) -> Self {
self.config.read_write_timeout = timeout;
self
}
pub fn build<T: XvcServer>(self, server: T) -> Server<T> {
Server::new(server, self.config)
}
}
impl<T: XvcServer> Server<T> {
pub fn new(server: T, config: Config) -> Server<T> {
Server {
server: Arc::new(Mutex::new(server)),
config,
}
}
pub async fn listen(&self, addr: impl ToSocketAddrs) -> io::Result<()>
where
T: Send + 'static,
{
let listener = TcpListener::bind(addr).await?;
self.listen_on(listener, CancellationToken::new()).await
}
pub async fn listen_on(
&self,
listener: TcpListener,
shutdown: CancellationToken,
) -> io::Result<()>
where
T: Send + 'static,
{
log::info!("Server listening for connections");
loop {
tokio::select! {
_ = shutdown.cancelled() => {
log::info!("Shutdown signal received, stopping listener");
break;
}
result = listener.accept() => {
match result {
Ok((stream, addr)) => {
let guard = match Arc::clone(&self.server).try_lock_owned() {
Ok(guard) => guard,
Err(_) => {
log::warn!("Rejected concurrent client from {}: another client is already active", addr);
continue;
}
};
stream.set_nodelay(true)?;
log::info!("New client connection from {}", addr);
let config = self.config.clone();
tokio::spawn(async move {
if let Err(e) = handle_client(guard, config, stream).await {
log::error!("Client error: {}", e);
}
});
}
Err(e) => log::error!("Connection error: {}", e),
}
}
}
}
Ok(())
}
}
async fn handle_client<T>(
server: tokio::sync::OwnedMutexGuard<T>,
config: Config,
stream: TcpStream,
) -> Result<(), ReadError>
where
T: XvcServer + Send + 'static,
{
let (mut read_half, mut write_half) = stream.into_split();
let mut buf = BytesMut::new();
let mut decoder = MessageDecoder::new(config.max_vector_size as usize);
loop {
match read_message(
&mut read_half,
&mut buf,
&mut decoder,
config.read_write_timeout,
)
.await
{
Ok(Some(msg)) => {
let response = block_in_place(|| compute_response(&*server, &config, msg))?;
write_half.write_all(&response).await?;
}
Ok(None) => break,
Err(e) => return Err(e),
}
}
Ok(())
}
async fn read_message(
read: &mut OwnedReadHalf,
buf: &mut BytesMut,
decoder: &mut MessageDecoder,
rw_timeout: Duration,
) -> Result<Option<OwnedMessage>, ReadError> {
loop {
if let Some(msg) = decoder.decode(buf)? {
return Ok(Some(msg));
}
match timeout(rw_timeout, read.read_buf(buf)).await {
Ok(Ok(0)) => return Ok(None), Ok(Ok(_)) => {} Ok(Err(e)) => return Err(ReadError::from(e)),
Err(_elapsed) => {
log::warn!("Client read timeout, closing connection");
return Ok(None);
}
}
}
}
fn compute_response<T: XvcServer>(
server: &T,
config: &Config,
msg: OwnedMessage,
) -> Result<Vec<u8>, ReadError> {
let mut buf = Vec::new();
match msg {
Message::GetInfo => {
log::info!("Received GetInfo message");
let info = XvcInfo::new(Version::V1_0, config.max_vector_size);
info.write_to(&mut buf)?;
log::debug!("Sent XVC info response");
}
Message::SetTck { period_ns } => {
log::debug!("Received SetTck message: period_ns={}", period_ns);
match server.set_tck(period_ns) {
Ok(ret_period) => {
log::debug!("Set TCK returned: period_ns={}", ret_period);
buf.extend_from_slice(&ret_period.to_le_bytes());
}
Err(e) => {
log::error!("Set TCK error: {e}");
buf.extend_from_slice(&period_ns.to_le_bytes());
}
}
}
Message::Shift { num_bits, tms, tdi } => {
log::debug!(
"Received Shift message: num_bits={}, tms_len={}, tdi_len={}",
num_bits,
tms.len(),
tdi.len()
);
log::trace!("Shift TMS data: {:02x?}", &tms[..]);
log::trace!("Shift TDI data: {:02x?}", &tdi[..]);
buf = vec![0; tdi.len()];
match server.shift(num_bits, &tms, &tdi, &mut buf) {
Ok(()) => {
log::trace!("Shift result TDO data: {:02x?}", &buf[..]);
}
Err(e) => {
log::error!("Shift error: {e}");
}
}
}
}
Ok(buf)
}