use super::{
Context, Result, ServerError,
aio::{RWriter, TReader},
message_handler,
};
use crate::{
raw::{R, T, Version},
server::{FileError, FileHandles, Filesystem, Requests},
};
use std::{collections::HashMap, net::SocketAddr, sync::Arc};
use tokio::sync::Mutex;
struct ConnectionParams {
msize: u32,
version: Version,
}
async fn handshake(
msize: u32,
version: &Version,
rw: &mut RWriter,
tr: &mut TReader,
) -> Result<ConnectionParams> {
loop {
let t = tr.next().await?;
let tag = t.tag();
match t {
T::Version(tag, client_msize, client_version) => {
tracing::debug!("client version {client_msize} {client_version}");
let conn_msize = msize.min(client_msize);
match version.try_negotiate(&client_version) {
Ok(conn_version) => {
rw.set_msize(conn_msize);
tr.set_msize(conn_msize);
rw.send(R::Version(tag, conn_msize, conn_version.clone()))
.await?;
return Ok(ConnectionParams {
version: conn_version,
msize: conn_msize,
});
}
Err(e) => {
rw.send(R::Error(tag, format!("{e:?}"), 0xFFFFFFFF)).await?;
return Err(ServerError::FailedToNegotiate);
}
};
}
_ => {
tracing::warn!("dropping unexpected message during handshake (tag={tag})");
}
}
}
}
pub struct MessageContext<'a, FilesystemT>
where
FilesystemT: Filesystem,
FilesystemT: Send,
FilesystemT: 'static,
{
pub(super) peer: SocketAddr,
pub(super) requests: &'a mut Requests,
pub(super) handles: &'a mut FileHandles<FilesystemT::File>,
pub(super) filesystems: Arc<Mutex<HashMap<String, FilesystemT>>>,
pub(super) msize: u32,
}
pub async fn connection_handler<FilesystemT>(
ctx: Context<FilesystemT>,
mut rw: RWriter,
mut tr: TReader,
) -> Result<()>
where
FilesystemT: Filesystem,
FilesystemT: Send,
FilesystemT: 'static,
{
let Context {
peer,
msize,
version,
mut handles,
mut requests,
filesystems,
} = ctx;
let ConnectionParams { msize, version } = handshake(msize, &version, &mut rw, &mut tr).await?;
tracing::info!("connection established with {peer}; version {version}, msize {msize}");
loop {
let t = tr.next().await?;
let tag = t.tag();
{
match requests.insert(tag, t.clone()) {
Ok(_) => {}
Err(_) => {
continue;
}
};
let mctx = MessageContext::<FilesystemT> {
peer,
requests: &mut requests,
handles: &mut handles,
filesystems: filesystems.clone(),
msize,
};
let reply = match message_handler(mctx, t).await {
Ok(r) => r,
Err(err) => match err {
ServerError::FileError(FileError(errno, desc)) => R::Error(tag, desc, errno),
_ => R::Error(tag, format!("{err:?}"), 0xFFFFFFFF),
},
};
tracing::debug!("reply tag={tag}: {:?}", reply);
match requests.remove(tag) {
Ok(_request) => {
rw.send(reply).await?;
}
Err(_) => {
tracing::trace!("reply tag={tag} not sent; was it flushed?");
}
}
}
}
}