use std::{
fs,
io::ErrorKind,
net::SocketAddr,
path::{Path, PathBuf},
sync::Arc,
};
use async_trait::async_trait;
#[cfg(feature = "serde")]
use serde::{Serialize, de::DeserializeOwned};
use tokio::{
io,
io::{AsyncRead, AsyncWrite, DuplexStream},
net::{
TcpListener as TokioTcpListener, TcpStream, UnixListener as TokioUnixListener, UnixStream,
},
sync::watch,
task::{JoinHandle, JoinSet},
};
use tracing::{trace, warn};
use crate::{
Connection, ConnectionMaker, RequestHandle, RpcSender, Value,
connection::{ConnectionMakerFn, ConnectionRuntime},
error::*,
};
pub fn duplex(buffer_size: usize) -> (DuplexStream, DuplexStream) {
io::duplex(buffer_size)
}
#[async_trait]
pub trait Listener: Send + Sync + 'static {
type Stream: AsyncRead + AsyncWrite + Unpin + Send + 'static;
async fn accept(&self) -> Result<Self::Stream>;
}
trait AsyncStream: AsyncRead + AsyncWrite {}
impl<T> AsyncStream for T where T: AsyncRead + AsyncWrite {}
type BoxedStream = Box<dyn AsyncStream + Unpin + Send>;
struct ErasedListener<L> {
inner: L,
}
#[async_trait]
impl<L> Listener for ErasedListener<L>
where
L: Listener,
{
type Stream = BoxedStream;
async fn accept(&self) -> Result<Self::Stream> {
Ok(Box::new(self.inner.accept().await?))
}
}
struct ConfiguredListener {
listener: Box<dyn Listener<Stream = BoxedStream>>,
local_addr: Option<SocketAddr>,
}
impl ConfiguredListener {
fn local_addr(&self) -> Result<SocketAddr> {
self.local_addr
.ok_or_else(|| RpcError::Protocol(ProtocolError::MissingSocketAddr))
}
fn into_parts(self) -> (Box<dyn Listener<Stream = BoxedStream>>, Option<SocketAddr>) {
(self.listener, self.local_addr)
}
}
struct TcpListener {
inner: TokioTcpListener,
}
impl TcpListener {
async fn bind(addr: &str) -> Result<Self> {
trace!("Binding TCP listener to address: {}", addr);
let listener = TokioTcpListener::bind(addr).await?;
Ok(Self { inner: listener })
}
}
#[async_trait]
impl Listener for TcpListener {
type Stream = TcpStream;
async fn accept(&self) -> Result<Self::Stream> {
let (stream, addr) = self.inner.accept().await?;
trace!("Accepted TCP connection from: {}", addr);
Ok(stream)
}
}
struct UnixListener {
inner: TokioUnixListener,
path: PathBuf,
}
impl UnixListener {
async fn bind<P: AsRef<Path>>(path: P) -> Result<Self> {
let path_str = path.as_ref().to_string_lossy();
trace!("Binding Unix listener to path: {}", path_str);
let listener = TokioUnixListener::bind(&path)?;
Ok(Self {
inner: listener,
path: path.as_ref().to_path_buf(),
})
}
}
#[async_trait]
impl Listener for UnixListener {
type Stream = UnixStream;
async fn accept(&self) -> Result<Self::Stream> {
let (stream, _) = self.inner.accept().await?;
trace!("Accepted Unix connection");
Ok(stream)
}
}
impl Drop for UnixListener {
fn drop(&mut self) {
match fs::remove_file(&self.path) {
Ok(()) => {}
Err(e) if e.kind() == ErrorKind::NotFound => {}
Err(e) => {
warn!("Failed to remove unix socket at {:?}: {}", self.path, e);
}
}
}
}
pub struct Server<T>
where
T: Connection,
{
connection_maker: Arc<dyn ConnectionMaker<T> + Send + Sync>,
listener: Option<ConfiguredListener>,
}
impl<T> Server<T>
where
T: Connection,
{
pub fn from_maker<M>(maker: M) -> Self
where
M: ConnectionMaker<T> + Send + Sync + 'static,
{
Self {
connection_maker: Arc::new(maker),
listener: None,
}
}
pub fn from_fn<F>(f: F) -> Self
where
F: Fn() -> T + Send + Sync + 'static,
{
Self::from_maker(ConnectionMakerFn::new(f))
}
pub fn local_addr(&self) -> Result<SocketAddr> {
self.listener
.as_ref()
.ok_or_else(|| RpcError::Protocol(ProtocolError::ListenerNotConfigured))?
.local_addr()
}
pub async fn tcp(mut self, addr: &str) -> Result<Self> {
let tcp_listener = TcpListener::bind(addr).await?;
let local_addr = Some(tcp_listener.inner.local_addr()?);
self.configure_listener(
Box::new(ErasedListener {
inner: tcp_listener,
}),
local_addr,
);
Ok(self)
}
pub async fn unix<P: AsRef<Path>>(mut self, path: P) -> Result<Self> {
self.configure_listener(
Box::new(ErasedListener {
inner: UnixListener::bind(path).await?,
}),
None,
);
Ok(self)
}
pub fn with_listener<L>(mut self, listener: L) -> Result<Self>
where
L: Listener,
{
self.configure_listener(Box::new(ErasedListener { inner: listener }), None);
Ok(self)
}
pub async fn spawn(self) -> Result<ServerHandle> {
let Self {
connection_maker,
listener,
} = self;
let listener =
listener.ok_or_else(|| RpcError::Protocol(ProtocolError::ListenerNotConfigured))?;
let (listener, local_addr) = listener.into_parts();
let (shutdown_tx, shutdown_rx) = watch::channel(false);
let task =
tokio::spawn(
async move { run_listener(listener, connection_maker, shutdown_rx).await },
);
Ok(ServerHandle {
shutdown_tx,
task,
local_addr,
})
}
pub async fn run(self) -> Result<()> {
self.spawn().await?.join().await
}
fn configure_listener(
&mut self,
listener: Box<dyn Listener<Stream = BoxedStream>>,
local_addr: Option<SocketAddr>,
) {
self.listener = Some(ConfiguredListener {
listener,
local_addr,
});
}
}
impl<T> Server<T>
where
T: Connection + Default,
{
pub fn from_listener<L>(listener: L) -> Result<Self>
where
L: Listener,
{
Self::from_maker(T::default()).with_listener(listener)
}
}
#[derive(Debug)]
pub struct ServerHandle {
shutdown_tx: watch::Sender<bool>,
task: JoinHandle<Result<()>>,
local_addr: Option<SocketAddr>,
}
impl ServerHandle {
pub fn shutdown(&self) {
let _send_result = self.shutdown_tx.send(true);
}
pub async fn join(self) -> Result<()> {
self.task
.await
.map_err(|source| RpcError::task_failed("server accept loop", source))?
}
pub fn local_addr(&self) -> Result<SocketAddr> {
self.local_addr
.ok_or_else(|| RpcError::Protocol(ProtocolError::MissingSocketAddr))
}
}
async fn run_listener<T>(
listener: Box<dyn Listener<Stream = BoxedStream>>,
connection_maker: Arc<dyn ConnectionMaker<T> + Send + Sync>,
mut shutdown_rx: watch::Receiver<bool>,
) -> Result<()>
where
T: Connection,
{
let mut connections = JoinSet::new();
loop {
tokio::select! {
_ = shutdown_rx.changed() => {
if *shutdown_rx.borrow() {
break;
}
}
accepted = listener.accept() => {
let stream = accepted?;
let connection = connection_maker.make_connection();
connections.spawn(async move {
serve_connection(stream, connection).await;
});
}
Some(joined) = connections.join_next(), if !connections.is_empty() => {
if let Err(e) = joined
&& !e.is_cancelled()
{
warn!("Error joining connection task: {}", e);
}
}
}
}
connections.abort_all();
while let Some(joined) = connections.join_next().await {
if let Err(e) = joined
&& !e.is_cancelled()
{
warn!("Error joining connection task: {}", e);
}
}
Ok(())
}
async fn serve_connection<S, T>(stream: S, connection: T)
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
T: Connection,
{
let runtime = ConnectionRuntime::new(stream, connection);
match runtime.run().await {
Ok(()) => {
trace!("Connection handler finished successfully");
}
Err(RpcError::Disconnect { .. }) => {
trace!("Client disconnected");
}
Err(e) => {
warn!("Connection error: {}", e);
}
}
}
#[derive(Debug)]
pub struct Client {
sender: RpcSender,
handle: Option<JoinHandle<()>>,
shutdown_tx: watch::Sender<bool>,
}
impl Client {
pub fn sender(&self) -> RpcSender {
self.sender.clone()
}
pub async fn from_stream<S, T>(stream: S, service: T) -> Result<Self>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
T: Connection,
{
Self::new(stream, service).await
}
pub async fn connect_unix<P, T>(path: P, service: T) -> Result<Self>
where
P: AsRef<Path>,
T: Connection,
{
let path_str = path.as_ref().to_string_lossy().to_string();
let stream = UnixStream::connect(path)
.await
.map_err(|source| RpcError::Connect { source })?;
trace!("Unix connection established to: {:?}", path_str);
Self::new(stream, service).await
}
pub async fn connect_tcp<T>(addr: &str, service: T) -> Result<Self>
where
T: Connection,
{
let stream = TcpStream::connect(addr)
.await
.map_err(|source| RpcError::Connect { source })?;
trace!("TCP connection established to: {}", addr);
Self::new(stream, service).await
}
async fn new<S, T>(stream: S, service: T) -> Result<Self>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
T: Connection,
{
let runtime = ConnectionRuntime::new(stream, service);
let shutdown_tx = runtime.shutdown_sender();
let rpc_sender = runtime.sender();
let handler_task = tokio::spawn(async move {
if let Err(e) = runtime.run().await {
match e {
RpcError::Disconnect { .. } => {
tracing::trace!("Client disconnected");
}
e => {
tracing::warn!("Handler error: {}", e);
}
}
}
});
Ok(Self {
sender: rpc_sender,
handle: Some(handler_task),
shutdown_tx,
})
}
pub async fn send_request(&self, method: &str, params: &[Value]) -> Result<Value> {
self.sender.send_request(method, params).await
}
pub async fn start_request(&self, method: &str, params: &[Value]) -> Result<RequestHandle> {
self.sender.start_request(method, params).await
}
pub async fn send_notification(&self, method: &str, params: &[Value]) -> Result<()> {
self.sender.send_notification(method, params).await
}
#[cfg(feature = "serde")]
pub async fn call<Req, Resp>(&self, method: &str, req: &Req) -> Result<Resp>
where
Req: Serialize,
Resp: DeserializeOwned,
{
self.sender.call(method, req).await
}
#[cfg(feature = "serde")]
pub async fn notify<Req>(&self, method: &str, req: &Req) -> Result<()>
where
Req: Serialize,
{
self.sender.notify(method, req).await
}
pub async fn join(mut self) -> Result<()> {
let handle = self
.handle
.take()
.ok_or_else(|| RpcError::resource_already_taken("client handler"))?;
handle
.await
.map_err(|source| RpcError::task_failed("client handler", source))?;
Ok(())
}
pub fn shutdown(&self) {
let _send_result = self.shutdown_tx.send(true);
}
pub async fn close(self) -> Result<()> {
self.shutdown();
self.join().await
}
}
impl Drop for Client {
fn drop(&mut self) {
let _send_result = self.shutdown_tx.send(true);
if let Some(handle) = &self.handle {
handle.abort();
}
}
}