use std::path::{Path, PathBuf};
use std::time::Duration;
use serde::Serialize;
use serde::de::DeserializeOwned;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use super::{UdsSecurityError, connect_hardened};
pub const MAX_FRAME_BYTES: u64 = 8 * 1024 * 1024;
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum UdsRpcError {
#[error("dial {path}: {source}")]
Dial {
path: PathBuf,
#[source]
source: UdsSecurityError,
},
#[error("serialize request frame for {path}: {source}")]
Encode {
path: PathBuf,
#[source]
source: serde_json::Error,
},
#[error("write request frame to {path}: {source}")]
Write {
path: PathBuf,
#[source]
source: std::io::Error,
},
#[error("read response frame from {path}: {source}")]
Read {
path: PathBuf,
#[source]
source: std::io::Error,
},
#[error("{path} closed the connection without sending a response frame")]
NoResponse {
path: PathBuf,
},
#[error("response frame from {path} exceeded {limit} bytes without a newline")]
FrameTooLarge {
path: PathBuf,
limit: u64,
},
#[error("decode response frame from {path}: {source}")]
Decode {
path: PathBuf,
#[source]
source: serde_json::Error,
},
#[error("{path} did not complete the exchange within {timeout:?}")]
Timeout {
path: PathBuf,
timeout: Duration,
},
}
pub async fn send_framed_request<Req, Resp>(
path: &Path,
request: &Req,
timeout: Duration,
) -> Result<Resp, UdsRpcError>
where
Req: Serialize + ?Sized,
Resp: DeserializeOwned,
{
match tokio::time::timeout(timeout, exchange::<Req, Resp>(path, request)).await {
Ok(result) => result,
Err(_) => Err(UdsRpcError::Timeout {
path: path.to_path_buf(),
timeout,
}),
}
}
async fn exchange<Req, Resp>(path: &Path, request: &Req) -> Result<Resp, UdsRpcError>
where
Req: Serialize + ?Sized,
Resp: DeserializeOwned,
{
let mut frame = serde_json::to_vec(request).map_err(|source| UdsRpcError::Encode {
path: path.to_path_buf(),
source,
})?;
frame.push(b'\n');
let mut stream = connect_hardened(path)
.await
.map_err(|source| UdsRpcError::Dial {
path: path.to_path_buf(),
source,
})?;
let write = async {
stream.write_all(&frame).await?;
stream.flush().await?;
stream.shutdown().await
};
write.await.map_err(|source| UdsRpcError::Write {
path: path.to_path_buf(),
source,
})?;
let mut reader = BufReader::new(stream.take(MAX_FRAME_BYTES));
let mut line: Vec<u8> = Vec::new();
let read = match reader.read_until(b'\n', &mut line).await {
Ok(read) => read,
Err(source) => return Err(classify_read_failure(path, source, line.is_empty())),
};
if read == 0 && line.is_empty() {
return Err(UdsRpcError::NoResponse {
path: path.to_path_buf(),
});
}
if !line.ends_with(b"\n") && line.len() as u64 >= MAX_FRAME_BYTES {
return Err(UdsRpcError::FrameTooLarge {
path: path.to_path_buf(),
limit: MAX_FRAME_BYTES,
});
}
serde_json::from_slice(&line).map_err(|source| UdsRpcError::Decode {
path: path.to_path_buf(),
source,
})
}
fn classify_read_failure(
path: &Path,
source: std::io::Error,
nothing_buffered: bool,
) -> UdsRpcError {
let hung_up = matches!(
source.kind(),
std::io::ErrorKind::ConnectionReset | std::io::ErrorKind::ConnectionAborted
);
if hung_up && nothing_buffered {
return UdsRpcError::NoResponse {
path: path.to_path_buf(),
};
}
UdsRpcError::Read {
path: path.to_path_buf(),
source,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::uds::bind_hardened;
use serde::Deserialize;
use std::path::PathBuf;
use tokio::net::UnixListener;
#[derive(Debug, Serialize)]
struct Ping {
method: &'static str,
n: u32,
}
#[derive(Debug, Deserialize, PartialEq, Eq)]
struct Pong {
echoed: u32,
}
enum StubReply {
Bytes(Vec<u8>),
HangUp,
Silence,
}
fn spawn_stub(dir: &Path, replies: Vec<StubReply>) -> PathBuf {
let sock = dir.join("sockets").join("stub.sock");
let listener: UnixListener = bind_hardened(&sock).expect("bind stub socket");
tokio::spawn(async move {
for reply in replies {
let Ok((mut conn, _)) = listener.accept().await else {
return;
};
let mut sink = Vec::new();
let _ = conn.read_to_end(&mut sink).await;
match reply {
StubReply::Bytes(bytes) => {
let _ = conn.write_all(&bytes).await;
let _ = conn.flush().await;
}
StubReply::HangUp => {}
StubReply::Silence => {
tokio::time::sleep(Duration::from_secs(300)).await;
}
}
}
});
sock
}
#[tokio::test]
async fn send_framed_request_round_trips_a_typed_value() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = spawn_stub(
tmp.path(),
vec![StubReply::Bytes(b"{\"echoed\":41}\n".to_vec())],
);
let got: Pong = send_framed_request(
&sock,
&Ping {
method: "ping",
n: 41,
},
Duration::from_secs(5),
)
.await
.expect("round trip");
assert_eq!(got, Pong { echoed: 41 });
}
#[tokio::test]
async fn send_framed_request_accepts_a_frame_without_a_trailing_newline() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = spawn_stub(
tmp.path(),
vec![StubReply::Bytes(b"{\"echoed\":7}".to_vec())],
);
let got: Pong = send_framed_request(
&sock,
&Ping {
method: "ping",
n: 7,
},
Duration::from_secs(5),
)
.await
.expect("round trip");
assert_eq!(got, Pong { echoed: 7 });
}
#[tokio::test]
async fn send_framed_request_reports_no_response_when_peer_hangs_up() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = spawn_stub(tmp.path(), vec![StubReply::HangUp]);
let err = send_framed_request::<_, Pong>(
&sock,
&Ping {
method: "ping",
n: 1,
},
Duration::from_secs(5),
)
.await
.expect_err("a silent hang-up is not a response");
assert!(
matches!(err, UdsRpcError::NoResponse { .. }),
"expected NoResponse, got {err:?}"
);
}
#[tokio::test]
async fn send_framed_request_rejects_an_over_long_frame() {
let tmp = tempfile::tempdir().expect("tempdir");
let flood = vec![b'x'; (MAX_FRAME_BYTES + 1) as usize];
let sock = spawn_stub(tmp.path(), vec![StubReply::Bytes(flood)]);
let err = send_framed_request::<_, Pong>(
&sock,
&Ping {
method: "ping",
n: 1,
},
Duration::from_secs(30),
)
.await
.expect_err("an unterminated flood must not be buffered without bound");
assert!(
matches!(err, UdsRpcError::FrameTooLarge { limit, .. } if limit == MAX_FRAME_BYTES),
"expected FrameTooLarge, got {err:?}"
);
}
#[tokio::test]
async fn send_framed_request_reports_a_decode_failure() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = spawn_stub(tmp.path(), vec![StubReply::Bytes(b"not json\n".to_vec())]);
let err = send_framed_request::<_, Pong>(
&sock,
&Ping {
method: "ping",
n: 1,
},
Duration::from_secs(5),
)
.await
.expect_err("garbage is not a response");
assert!(
matches!(err, UdsRpcError::Decode { .. }),
"expected Decode, got {err:?}"
);
}
#[tokio::test]
async fn send_framed_request_reports_dial_failure_for_a_missing_socket() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = tmp.path().join("sockets").join("absent.sock");
let err = send_framed_request::<_, Pong>(
&sock,
&Ping {
method: "ping",
n: 1,
},
Duration::from_secs(5),
)
.await
.expect_err("no listener means no delivery");
assert!(
matches!(err, UdsRpcError::Dial { .. }),
"expected Dial, got {err:?}"
);
}
#[tokio::test]
async fn send_framed_request_times_out_on_a_silent_peer() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = spawn_stub(tmp.path(), vec![StubReply::Silence]);
let err = send_framed_request::<_, Pong>(
&sock,
&Ping {
method: "ping",
n: 1,
},
Duration::from_millis(150),
)
.await
.expect_err("a peer that never answers must not hold the caller open");
assert!(
matches!(err, UdsRpcError::Timeout { .. }),
"expected Timeout, got {err:?}"
);
}
#[test]
fn read_failure_from_an_abortive_close_reads_as_a_hang_up() {
for kind in [
std::io::ErrorKind::ConnectionReset,
std::io::ErrorKind::ConnectionAborted,
] {
let err = classify_read_failure(
Path::new("/tmp/relay.sock"),
std::io::Error::new(kind, "peer went away"),
true,
);
assert!(
matches!(err, UdsRpcError::NoResponse { .. }),
"expected NoResponse for {kind:?}, got {err:?}"
);
}
}
#[test]
fn read_failure_after_partial_bytes_stays_a_read_error() {
let err = classify_read_failure(
Path::new("/tmp/relay.sock"),
std::io::Error::new(std::io::ErrorKind::ConnectionReset, "peer went away"),
false,
);
assert!(
matches!(err, UdsRpcError::Read { .. }),
"expected Read, got {err:?}"
);
}
#[test]
fn read_failure_from_an_unrelated_errno_stays_a_read_error() {
let err = classify_read_failure(
Path::new("/tmp/relay.sock"),
std::io::Error::from(std::io::ErrorKind::PermissionDenied),
true,
);
assert!(
matches!(err, UdsRpcError::Read { .. }),
"expected Read, got {err:?}"
);
}
}