mod router;
mod wire;
#[cfg(test)]
#[path = "tests.rs"]
mod tests;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncReadExt as _, AsyncWriteExt, BufReader};
use tokio::net::{UnixListener, UnixStream};
use crate::uds::{MAX_FRAME_BYTES, UdsSecurityError, bind_hardened, ensure_peer_is_self};
pub use router::{RpcMethod, RpcRouter, typed_method};
pub use wire::{
CODE_INTERNAL_ERROR, CODE_INVALID_PARAMS, CODE_INVALID_REQUEST, CODE_METHOD_NOT_FOUND,
CODE_PARSE_ERROR, JSONRPC_VERSION, RpcError, RpcRequest, RpcResponse,
};
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum RpcServerError {
#[error("bind the rpc socket at {path}: {source}")]
Bind {
path: PathBuf,
#[source]
source: UdsSecurityError,
},
#[error("refuse an accepted connection: {source}")]
Peer {
#[source]
source: UdsSecurityError,
},
#[error("read request frame: {source}")]
Read {
#[source]
source: std::io::Error,
},
#[error("request frame exceeded {limit} bytes without a newline")]
FrameTooLarge {
limit: u64,
},
#[error("serialize response frame: {source}")]
Encode {
#[source]
source: serde_json::Error,
},
#[error("write response frame: {source}")]
Write {
#[source]
source: std::io::Error,
},
}
#[derive(Debug, Clone, Copy)]
pub struct RpcServeOptions {
pub read_timeout: Duration,
pub max_frame_bytes: u64,
}
impl Default for RpcServeOptions {
fn default() -> Self {
Self {
read_timeout: Duration::from_secs(30),
max_frame_bytes: MAX_FRAME_BYTES,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Served {
Answered {
errored: bool,
},
LivenessProbe,
}
pub async fn handle_connection(
mut stream: UnixStream,
router: Arc<RpcRouter>,
options: RpcServeOptions,
) -> Result<Served, RpcServerError> {
ensure_peer_is_self(&stream).map_err(|source| RpcServerError::Peer { source })?;
let mut frame: Vec<u8> = Vec::new();
{
let mut reader = BufReader::new((&mut stream).take(options.max_frame_bytes));
let read = tokio::time::timeout(options.read_timeout, reader.read_until(b'\n', &mut frame))
.await
.map_err(|_| RpcServerError::Read {
source: std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("no complete frame within {:?}", options.read_timeout),
),
})?
.map_err(|source| RpcServerError::Read { source })?;
if read == 0 && frame.is_empty() {
return Ok(Served::LivenessProbe);
}
if !frame.ends_with(b"\n") && frame.len() as u64 >= options.max_frame_bytes {
return Err(RpcServerError::FrameTooLarge {
limit: options.max_frame_bytes,
});
}
}
let response = router.dispatch(&frame).await;
let errored = response.is_error();
let bytes =
crate::uds::encode_frame(&response).map_err(|source| RpcServerError::Encode { source })?;
stream
.write_all(&bytes)
.await
.map_err(|source| RpcServerError::Write { source })?;
stream
.flush()
.await
.map_err(|source| RpcServerError::Write { source })?;
Ok(Served::Answered { errored })
}
pub async fn serve_until(
listener: &UnixListener,
router: Arc<RpcRouter>,
options: RpcServeOptions,
shutdown: impl std::future::Future<Output = ()> + Send,
) {
tokio::pin!(shutdown);
loop {
let accepted = tokio::select! {
biased;
() = &mut shutdown => return,
accepted = listener.accept() => accepted,
};
let stream = match accepted {
Ok((stream, _)) => stream,
Err(e) => {
tracing::warn!(error = %e, "uds rpc listener accept failed");
continue;
}
};
let router = Arc::clone(&router);
tokio::spawn(async move {
let inner = tokio::spawn(handle_connection(stream, router, options));
match inner.await {
Ok(Ok(Served::Answered { .. })) => {}
Ok(Ok(Served::LivenessProbe)) => {
tracing::debug!("liveness probe connected and closed without a frame");
}
Ok(Err(e)) => {
tracing::warn!(error = %e, "uds rpc connection failed; nothing was answered");
}
Err(join) if join.is_panic() => {
tracing::error!(
error = %join,
"uds rpc handler panicked; the connection was dropped unanswered"
);
}
Err(join) => {
tracing::warn!(error = %join, "uds rpc connection task was cancelled");
}
}
});
}
}
#[derive(Debug)]
pub struct RpcServer {
socket: PathBuf,
router: Arc<RpcRouter>,
options: RpcServeOptions,
}
impl RpcServer {
pub fn new(socket: impl Into<PathBuf>, router: RpcRouter) -> Self {
Self {
socket: socket.into(),
router: Arc::new(router),
options: RpcServeOptions::default(),
}
}
pub fn with_options(mut self, options: RpcServeOptions) -> Self {
self.options = options;
self
}
pub fn socket(&self) -> &Path {
&self.socket
}
pub async fn run(
self,
shutdown: impl std::future::Future<Output = ()> + Send,
) -> Result<(), RpcServerError> {
let listener = bind_hardened(&self.socket).map_err(|source| RpcServerError::Bind {
path: self.socket.clone(),
source,
})?;
tracing::info!(
socket = %self.socket.display(),
methods = ?self.router.method_names().collect::<Vec<_>>(),
"uds rpc server bound"
);
serve_until(&listener, Arc::clone(&self.router), self.options, shutdown).await;
if let Err(e) = std::fs::remove_file(&self.socket) {
tracing::debug!(socket = %self.socket.display(), error = %e, "socket already gone");
}
drop(listener);
Ok(())
}
}