use super::*;
use std::sync::Arc;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::json;
#[derive(Debug, Serialize, Deserialize)]
struct Greet {
name: String,
}
#[derive(Debug, Serialize, Deserialize, PartialEq, Eq)]
struct Greeting {
text: String,
}
fn greeting_router() -> RpcRouter {
RpcRouter::new()
.typed("greet", |req: Greet| async move {
Ok(Greeting {
text: format!("hello {}", req.name),
})
})
.typed("explode", |_req: ()| async move {
Err::<(), _>(RpcError::new(-32001, "handler said no"))
})
}
fn frame(id: u64, method: &str, params: serde_json::Value) -> Vec<u8> {
serde_json::to_vec(&json!({
"jsonrpc": "2.0",
"id": id,
"method": method,
"params": params,
}))
.expect("serialize test frame")
}
#[tokio::test]
async fn dispatch_routes_to_the_registered_handler() {
let response = greeting_router()
.dispatch(&frame(7, "greet", json!({ "name": "ada" })))
.await;
assert_eq!(response.id, json!(7), "the request id must be echoed back");
assert!(response.error.is_none(), "unexpected error: {response:?}");
let greeting: Greeting =
serde_json::from_value(response.result.expect("a result")).expect("decode result");
assert_eq!(
greeting,
Greeting {
text: "hello ada".to_string()
}
);
}
#[tokio::test]
async fn dispatch_reports_method_not_found_for_an_unregistered_method() {
let response = greeting_router()
.dispatch(&frame(1, "review", json!(null)))
.await;
let error = response.error.expect("an error");
assert_eq!(error.code, CODE_METHOD_NOT_FOUND);
assert!(
error.message.contains("review") && error.message.contains("greet"),
"the refusal must name both the bad method and the served ones: {}",
error.message
);
assert_eq!(response.id, json!(1));
}
#[tokio::test]
async fn dispatch_rejects_an_unparseable_frame() {
let response = greeting_router().dispatch(b"{not json").await;
assert_eq!(
response.error.expect("an error").code,
CODE_PARSE_ERROR,
"an unreadable frame is a parse error, not a method-not-found"
);
assert_eq!(
response.id,
serde_json::Value::Null,
"there was no readable id to echo"
);
}
#[tokio::test]
async fn dispatch_rejects_a_wrong_jsonrpc_version() {
let raw = serde_json::to_vec(&json!({
"jsonrpc": "1.0",
"id": 3,
"method": "greet",
"params": { "name": "ada" },
}))
.expect("serialize");
let response = greeting_router().dispatch(&raw).await;
assert_eq!(response.error.expect("an error").code, CODE_INVALID_REQUEST);
}
#[tokio::test]
async fn dispatch_reports_invalid_params_for_an_undecodable_payload() {
let response = greeting_router()
.dispatch(&frame(4, "greet", json!({ "name": 17 })))
.await;
let error = response.error.expect("an error");
assert_eq!(error.code, CODE_INVALID_PARAMS);
assert!(
error.message.contains("params do not decode"),
"the serde reason must survive: {}",
error.message
);
}
#[tokio::test]
async fn dispatch_propagates_a_handler_error_verbatim() {
let response = greeting_router()
.dispatch(&frame(5, "explode", json!(null)))
.await;
assert_eq!(
response.error.expect("an error"),
RpcError::new(-32001, "handler said no"),
"a handler's own code and message must not be rewritten"
);
}
struct EchoFallback {
refuse: &'static str,
}
#[async_trait::async_trait]
impl RpcFallback for EchoFallback {
async fn call(
&self,
method: &str,
params: serde_json::Value,
) -> Result<serde_json::Value, RpcError> {
if method == self.refuse {
return Err(RpcError::new(-32002, format!("fallback refused {method}")));
}
Ok(json!({ "method": method, "params": params }))
}
}
#[tokio::test]
async fn dispatch_routes_an_unregistered_method_to_the_fallback() {
let router = greeting_router().fallback(EchoFallback { refuse: "" });
let response = router
.dispatch(&frame(11, "review", json!({ "n": 1 })))
.await;
assert!(response.error.is_none(), "unexpected error: {response:?}");
assert_eq!(
response.result,
Some(json!({ "method": "review", "params": { "n": 1 } })),
"the fallback must receive both the method name and the params"
);
assert_eq!(response.id, json!(11), "the request id is still echoed");
}
#[tokio::test]
async fn dispatch_prefers_a_registered_method_over_the_fallback() {
let router = greeting_router().fallback(EchoFallback { refuse: "" });
let response = router
.dispatch(&frame(12, "greet", json!({ "name": "ada" })))
.await;
assert_eq!(
response.result,
Some(json!({ "text": "hello ada" })),
"the registered handler answered, not the fallback"
);
}
#[tokio::test]
async fn dispatch_maps_a_fallback_error_to_an_rpc_error_response() {
let router = greeting_router().fallback(EchoFallback { refuse: "review" });
let response = router.dispatch(&frame(13, "review", json!(null))).await;
assert_eq!(
response.error.expect("an error"),
RpcError::new(-32002, "fallback refused review"),
"the fallback's own code and message must survive verbatim"
);
assert!(response.result.is_none());
assert_eq!(response.id, json!(13));
}
#[test]
fn method_names_are_sorted_and_complete() {
let router = greeting_router();
assert_eq!(
router.method_names().collect::<Vec<_>>(),
vec!["explode", "greet"]
);
}
fn spawn_server(
dir: &std::path::Path,
router: RpcRouter,
options: RpcServeOptions,
) -> (
std::path::PathBuf,
tokio::sync::oneshot::Sender<()>,
tokio::task::JoinHandle<Result<(), RpcServerError>>,
) {
let socket = dir.join("rpc.sock");
let (tx, rx) = tokio::sync::oneshot::channel();
let server = RpcServer::new(socket.clone(), router).with_options(options);
let handle = tokio::spawn(async move {
server
.run(async move {
let _ = rx.await;
})
.await
});
(socket, tx, handle)
}
async fn call(
socket: &std::path::Path,
id: u64,
method: &str,
params: serde_json::Value,
) -> RpcResponse {
let request = json!({ "jsonrpc": "2.0", "id": id, "method": method, "params": params });
crate::uds::send_framed_request(socket, &request, Duration::from_secs(10))
.await
.expect("round trip")
}
async fn await_socket(socket: &std::path::Path) {
for _ in 0..200 {
if crate::uds::socket_is_serving(socket, Duration::from_millis(200)).await {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("server never began serving {}", socket.display());
}
#[tokio::test]
async fn serve_round_trips_a_request_over_a_real_socket() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) =
spawn_server(tmp.path(), greeting_router(), RpcServeOptions::default());
await_socket(&socket).await;
let response = call(&socket, 42, "greet", json!({ "name": "grace" })).await;
assert_eq!(response.id, json!(42));
let greeting: Greeting =
serde_json::from_value(response.result.expect("a result")).expect("decode");
assert_eq!(greeting.text, "hello grace");
}
#[tokio::test]
async fn serve_answers_an_unknown_method_rather_than_hanging_up() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) =
spawn_server(tmp.path(), greeting_router(), RpcServeOptions::default());
await_socket(&socket).await;
let response = call(&socket, 1, "nope", json!(null)).await;
assert_eq!(
response.error.expect("an error").code,
CODE_METHOD_NOT_FOUND
);
}
#[tokio::test]
async fn serve_handles_concurrent_connections_without_serialising() {
let peers = 4usize;
let barrier = Arc::new(tokio::sync::Barrier::new(peers));
let router = RpcRouter::new().typed("wait", move |_req: ()| {
let barrier = Arc::clone(&barrier);
async move {
barrier.wait().await;
Ok::<bool, RpcError>(true)
}
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
let mut calls = Vec::new();
for n in 0..peers {
let socket = socket.clone();
calls.push(tokio::spawn(async move {
call(&socket, n as u64, "wait", json!(null)).await
}));
}
for handle in calls {
let response = handle.await.expect("join");
assert_eq!(
response.result,
Some(json!(true)),
"a serialised server deadlocks here instead of answering"
);
}
}
#[tokio::test]
async fn serve_survives_a_panicking_handler_and_answers_the_next_connection() {
let router = RpcRouter::new()
.typed("boom", |_req: ()| async move {
panic!("a handler exploded");
#[allow(unreachable_code)]
Ok::<(), RpcError>(())
})
.typed("greet", |req: Greet| async move {
Ok(Greeting {
text: format!("hello {}", req.name),
})
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
let request = json!({ "jsonrpc": "2.0", "id": 1, "method": "boom", "params": null });
let panicked: Result<RpcResponse, _> =
crate::uds::send_framed_request(&socket, &request, Duration::from_secs(10)).await;
assert!(
panicked.is_err(),
"a panicking handler answers nothing; the client must see a transport \
failure rather than hang: {panicked:?}"
);
let response = call(&socket, 2, "greet", json!({ "name": "ada" })).await;
assert_eq!(
response.result,
Some(json!({ "text": "hello ada" })),
"one panicking connection must not stop the accept loop"
);
}
#[tokio::test]
async fn serve_stops_on_shutdown() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, stop, handle) =
spawn_server(tmp.path(), greeting_router(), RpcServeOptions::default());
await_socket(&socket).await;
stop.send(()).expect("signal shutdown");
tokio::time::timeout(Duration::from_secs(5), handle)
.await
.expect("the server must return once shutdown resolves")
.expect("join")
.expect("clean shutdown");
}
#[tokio::test]
async fn server_round_trips_and_removes_its_socket_on_shutdown() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, stop, handle) =
spawn_server(tmp.path(), greeting_router(), RpcServeOptions::default());
await_socket(&socket).await;
let response = call(&socket, 9, "greet", json!({ "name": "ada" })).await;
assert!(response.error.is_none());
stop.send(()).expect("signal shutdown");
handle.await.expect("join").expect("clean shutdown");
assert!(
!socket.exists(),
"the socket file must be gone after shutdown, or the next bind fails"
);
}
fn stream_frame(id: u64, method: &str, params: serde_json::Value) -> serde_json::Value {
json!({ "jsonrpc": "2.0", "id": id, "method": method, "params": params, "stream": true })
}
fn token_router(count: usize) -> RpcRouter {
greeting_router().typed_stream("tokens", move |_req: ()| async move {
let (tx, rx) = tokio::sync::mpsc::channel(4);
tokio::spawn(async move {
for n in 0..count {
if tx.send(Ok(json!(format!("t{n}")))).await.is_err() {
return;
}
}
});
Ok(rx)
})
}
#[test]
fn stream_frames_carry_the_phase_discriminant() {
let item = serde_json::to_value(RpcStreamFrame::item(json!(1), json!("tok"))).expect("encode");
assert_eq!(item["stream"], json!("item"));
assert_eq!(item["result"], json!("tok"));
let end = serde_json::to_value(RpcStreamFrame::end(json!(1))).expect("encode");
assert_eq!(end["stream"], json!("end"));
assert!(end.get("result").is_none() && end.get("error").is_none());
let failed =
serde_json::to_value(RpcStreamFrame::error(json!(1), RpcError::internal("no"))).expect("e");
assert_eq!(failed["stream"], json!("error"));
assert_eq!(failed["error"]["message"], json!("no"));
let unary = serde_json::to_value(RpcResponse::success(json!(1), json!("x"))).expect("encode");
assert!(
unary.get("stream").is_none(),
"a unary response must never carry the discriminant"
);
}
#[test]
fn stream_names_are_sorted_and_separate_from_unary_names() {
let router = token_router(1);
assert_eq!(router.stream_names().collect::<Vec<_>>(), vec!["tokens"]);
assert_eq!(
router.method_names().collect::<Vec<_>>(),
vec!["explode", "greet"],
"a streaming name must not appear in the unary table"
);
}
#[tokio::test]
async fn dispatch_streaming_answers_a_unary_request_unchanged() {
for (id, method, params) in [
(1u64, "greet", json!({ "name": "ada" })),
(2, "explode", json!(null)),
(3, "nope", json!(null)),
] {
let raw = frame(id, method, params.clone());
let unary = greeting_router().dispatch(&raw).await;
let wide = match greeting_router().dispatch_streaming(&raw).await {
RpcOutcome::Single(response) => response,
other => panic!("a request without the flag must not stream: {other:?}"),
};
assert_eq!(
serde_json::to_value(&unary).expect("encode"),
serde_json::to_value(&wide).expect("encode"),
"dispatch_streaming changed the answer for {method}"
);
}
}
#[tokio::test]
async fn stream_opt_in_is_read_from_the_request_frame() {
let router = token_router(1);
let without = frame(1, "tokens", json!(null));
match router.dispatch_streaming(&without).await {
RpcOutcome::Single(response) => {
assert_eq!(response.error.expect("an error").code, CODE_STREAM_REQUIRED)
}
other => panic!("no flag must not produce a stream: {other:?}"),
}
let with = serde_json::to_vec(&stream_frame(1, "tokens", json!(null))).expect("encode");
match router.dispatch_streaming(&with).await {
RpcOutcome::Stream { id, .. } => assert_eq!(id, json!(1)),
other => panic!("the flag must produce a stream: {other:?}"),
}
}
async fn collect_stream(
socket: &std::path::Path,
id: u64,
method: &str,
params: serde_json::Value,
) -> Result<Vec<String>, crate::uds::UdsRpcError> {
let request = stream_frame(id, method, params);
let mut stream: crate::uds::FramedStream<String> =
crate::uds::send_framed_stream_request(socket, &request, Duration::from_secs(10)).await?;
let mut items = Vec::new();
while let Some(item) = stream.next_frame().await {
items.push(item?);
}
Ok(items)
}
#[tokio::test]
async fn stream_round_trips_many_frames_over_a_real_socket() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) =
spawn_server(tmp.path(), token_router(3), RpcServeOptions::default());
await_socket(&socket).await;
let items = collect_stream(&socket, 1, "tokens", json!(null))
.await
.expect("the stream must complete");
assert_eq!(items, vec!["t0", "t1", "t2"]);
}
#[tokio::test]
async fn stream_of_zero_items_still_ends_on_a_terminal_frame() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) =
spawn_server(tmp.path(), token_router(0), RpcServeOptions::default());
await_socket(&socket).await;
let items = collect_stream(&socket, 1, "tokens", json!(null))
.await
.expect("an empty stream is still a complete one");
assert!(items.is_empty());
}
#[tokio::test]
async fn stream_reports_a_handler_error_as_a_terminal_frame() {
let router = RpcRouter::new().typed_stream("tokens", |_req: ()| async move {
let (tx, rx) = tokio::sync::mpsc::channel(4);
tokio::spawn(async move {
let _ = tx.send(Ok(json!("t0"))).await;
let _ = tx.send(Ok(json!("t1"))).await;
let _ = tx
.send(Err(RpcError::new(-32003, "the model gave up")))
.await;
});
Ok(rx)
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
let err = collect_stream(&socket, 1, "tokens", json!(null))
.await
.expect_err("a mid-stream failure must not read as a complete answer");
match err {
crate::uds::UdsRpcError::Stream { error, .. } => {
assert_eq!(error, RpcError::new(-32003, "the model gave up"),);
}
other => panic!("expected Stream, got {other:?}"),
}
}
#[tokio::test]
async fn stream_reports_an_open_failure_as_a_terminal_frame() {
let router = RpcRouter::new().typed_stream("tokens", |_req: ()| async move {
Err(RpcError::new(-32004, "no model configured"))
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
let err = collect_stream(&socket, 1, "tokens", json!(null))
.await
.expect_err("an open failure is still an error");
assert!(
matches!(&err, crate::uds::UdsRpcError::Stream { error, .. } if error.code == -32004),
"expected the handler's own code, got {err:?}"
);
}
#[tokio::test]
async fn stream_reports_invalid_params_before_opening_the_stream() {
let router = RpcRouter::new().typed_stream("tokens", |_req: Greet| async move {
let (_tx, rx) = tokio::sync::mpsc::channel(1);
Ok(rx)
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
let err = collect_stream(&socket, 1, "tokens", json!({ "name": 17 }))
.await
.expect_err("bad params must not open a stream");
assert!(
matches!(&err, crate::uds::UdsRpcError::Stream { error, .. }
if error.code == CODE_INVALID_PARAMS),
"expected invalid_params, got {err:?}"
);
}
#[tokio::test]
async fn stream_request_for_a_non_streaming_method_is_refused() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) =
spawn_server(tmp.path(), token_router(1), RpcServeOptions::default());
await_socket(&socket).await;
for method in ["greet", "no-such-method"] {
let err = tokio::time::timeout(
Duration::from_secs(5),
collect_stream(&socket, 1, method, json!({ "name": "ada" })),
)
.await
.unwrap_or_else(|_| panic!("{method} hung instead of failing"))
.expect_err("a method that does not stream must refuse");
match err {
crate::uds::UdsRpcError::Stream { error, .. } => {
assert_eq!(error.code, CODE_STREAM_UNSUPPORTED);
assert!(
error.message.contains("tokens"),
"the refusal must name what this listener does stream: {}",
error.message
);
}
other => panic!("expected a terminal error frame for {method}, got {other:?}"),
}
}
}
#[tokio::test]
async fn unary_request_for_a_streaming_method_is_refused_in_one_frame() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) =
spawn_server(tmp.path(), token_router(3), RpcServeOptions::default());
await_socket(&socket).await;
let response = tokio::time::timeout(
Duration::from_secs(5),
call(&socket, 1, "tokens", json!(null)),
)
.await
.expect("a unary call on a streaming method must not hang");
let error = response.error.expect("an error");
assert_eq!(error.code, CODE_STREAM_REQUIRED);
assert!(
error.message.contains("stream"),
"the refusal must say how to ask again: {}",
error.message
);
}
#[tokio::test]
async fn stream_refuses_an_item_larger_than_the_frame_budget() {
let router = RpcRouter::new().typed_stream("tokens", |_req: ()| async move {
let (tx, rx) = tokio::sync::mpsc::channel(2);
tokio::spawn(async move {
let _ = tx.send(Ok(json!("small"))).await;
let _ = tx.send(Ok(json!("x".repeat(4096)))).await;
});
Ok(rx)
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(
tmp.path(),
router,
RpcServeOptions {
max_frame_bytes: 512,
..RpcServeOptions::default()
},
);
await_socket(&socket).await;
let request = stream_frame(1, "tokens", json!(null));
let mut stream: crate::uds::FramedStream<String> =
crate::uds::send_framed_stream_request_capped(
&socket,
&request,
Duration::from_secs(10),
512,
)
.await
.expect("open");
assert_eq!(
stream.next_frame().await.expect("an item").expect("ok"),
"small",
"the frames before the oversized one still arrive"
);
let err = stream
.next_frame()
.await
.expect("a report")
.expect_err("an oversized item must not be written");
assert!(
matches!(&err, crate::uds::UdsRpcError::Stream { error, .. }
if error.message.contains("frame budget")),
"expected a terminal budget refusal, got {err:?}"
);
}
#[tokio::test]
async fn stream_serves_one_frame_requests_on_other_connections_while_running() {
let (release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let release = Arc::new(tokio::sync::Mutex::new(Some(release_rx)));
let router = greeting_router().typed_stream("tokens", move |_req: ()| {
let release = Arc::clone(&release);
async move {
let (tx, rx) = tokio::sync::mpsc::channel(4);
tokio::spawn(async move {
let _ = tx.send(Ok(json!("first"))).await;
if let Some(gate) = release.lock().await.take() {
let _ = gate.await;
}
let _ = tx.send(Ok(json!("last"))).await;
});
Ok(rx)
}
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
let request = stream_frame(1, "tokens", json!(null));
let mut stream: crate::uds::FramedStream<String> =
crate::uds::send_framed_stream_request(&socket, &request, Duration::from_secs(10))
.await
.expect("open");
assert_eq!(
stream.next_frame().await.expect("an item").expect("ok"),
"first",
"the stream is live before the interleaved calls"
);
for n in 0..3u64 {
let response = tokio::time::timeout(
Duration::from_secs(5),
call(&socket, 100 + n, "greet", json!({ "name": "ada" })),
)
.await
.expect("a one-frame call must not queue behind a live stream");
assert_eq!(response.result, Some(json!({ "text": "hello ada" })));
}
release_tx.send(()).expect("release the stream");
assert_eq!(
stream.next_frame().await.expect("an item").expect("ok"),
"last"
);
assert!(
stream.next_frame().await.is_none(),
"the terminal frame ends the stream"
);
}
#[tokio::test]
async fn stream_survives_a_client_that_disconnects_mid_stream() {
let router = greeting_router().typed_stream("tokens", |_req: ()| async move {
let (tx, rx) = tokio::sync::mpsc::channel(1);
tokio::spawn(async move {
for n in 0..10_000u32 {
if tx.send(Ok(json!(format!("t{n}")))).await.is_err() {
return;
}
}
});
Ok(rx)
});
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, _handle) = spawn_server(tmp.path(), router, RpcServeOptions::default());
await_socket(&socket).await;
{
let request = stream_frame(1, "tokens", json!(null));
let mut stream: crate::uds::FramedStream<String> =
crate::uds::send_framed_stream_request(&socket, &request, Duration::from_secs(10))
.await
.expect("open");
assert_eq!(
stream.next_frame().await.expect("an item").expect("ok"),
"t0"
);
}
let response = tokio::time::timeout(
Duration::from_secs(5),
call(&socket, 2, "greet", json!({ "name": "ada" })),
)
.await
.expect("an abandoned stream must not wedge the accept loop");
assert_eq!(response.result, Some(json!({ "text": "hello ada" })));
}
#[tokio::test]
async fn serve_rejects_an_oversized_frame() {
let options = RpcServeOptions {
max_frame_bytes: 64,
..RpcServeOptions::default()
};
let (mut client, server) = tokio::net::UnixStream::pair().expect("socketpair");
let writer = tokio::spawn(async move {
use tokio::io::AsyncWriteExt as _;
let huge = frame(1, "greet", json!({ "name": "x".repeat(4096) }));
let _ = client.write_all(&huge).await;
let _ = client.flush().await;
tokio::time::sleep(Duration::from_secs(2)).await;
});
let outcome = handle_connection(server, Arc::new(greeting_router()), options).await;
match outcome {
Err(RpcServerError::FrameTooLarge { limit }) => assert_eq!(limit, 64),
other => panic!("expected FrameTooLarge, got {other:?}"),
}
writer.abort();
}
async fn serve_one_frame(body: Vec<u8>, max_frame_bytes: u64) -> Result<Served, RpcServerError> {
let (mut client, server) = tokio::net::UnixStream::pair().expect("socketpair");
let writer = tokio::spawn(async move {
use tokio::io::AsyncWriteExt as _;
let _ = client.write_all(&body).await;
let _ = client.flush().await;
tokio::time::sleep(Duration::from_secs(2)).await;
});
let options = RpcServeOptions {
max_frame_bytes,
..RpcServeOptions::default()
};
let outcome = handle_connection(server, Arc::new(greeting_router()), options).await;
writer.abort();
outcome
}
#[tokio::test]
async fn frame_of_exactly_the_budget_including_its_newline_is_accepted() {
let empty = frame(1, "greet", json!({ "name": "" })).len();
let budget = (empty + 40) as u64;
let padding = "x".repeat(budget as usize - 1 - empty);
let mut body = frame(1, "greet", json!({ "name": padding }));
assert_eq!(
body.len() as u64,
budget - 1,
"the JSON body must be one byte short of the budget"
);
body.push(b'\n');
assert_eq!(
serve_one_frame(body.clone(), budget)
.await
.expect("a frame that exactly fills the budget is accepted"),
Served::Answered {
errored: false,
liveness: false
}
);
match serve_one_frame(body, budget - 1).await {
Err(RpcServerError::FrameTooLarge { limit }) => assert_eq!(limit, budget - 1),
other => panic!("one byte over the budget must be refused, got {other:?}"),
}
}
#[tokio::test]
async fn handle_connection_reports_a_liveness_probe_rather_than_a_failure() {
let (client, server) = tokio::net::UnixStream::pair().expect("socketpair");
drop(client);
let served = handle_connection(
server,
Arc::new(greeting_router()),
RpcServeOptions::default(),
)
.await
.expect("a closed probe is not an error");
assert_eq!(served, Served::LivenessProbe);
}
#[tokio::test]
async fn handle_connection_reports_an_error_response_as_answered() {
let (mut client, server) = tokio::net::UnixStream::pair().expect("socketpair");
let writer = tokio::spawn(async move {
use tokio::io::AsyncWriteExt as _;
let mut bytes = frame(1, "nope", json!(null));
bytes.push(b'\n');
let _ = client.write_all(&bytes).await;
let _ = client.flush().await;
tokio::time::sleep(Duration::from_secs(2)).await;
});
let served = handle_connection(
server,
Arc::new(greeting_router()),
RpcServeOptions::default(),
)
.await
.expect("a refusal is still an answer");
assert_eq!(
served,
Served::Answered {
errored: true,
liveness: false
}
);
writer.abort();
}
fn spawn_idle_server(
dir: &std::path::Path,
idle: Duration,
) -> (
std::path::PathBuf,
tokio::sync::oneshot::Sender<()>,
tokio::task::JoinHandle<ServeExit>,
) {
spawn_idle_server_with(dir, idle, greeting_router())
}
fn spawn_idle_server_with(
dir: &std::path::Path,
idle: Duration,
router: RpcRouter,
) -> (
std::path::PathBuf,
tokio::sync::oneshot::Sender<()>,
tokio::task::JoinHandle<ServeExit>,
) {
let socket = dir.join("idle.sock");
let (tx, rx) = tokio::sync::oneshot::channel();
let bound = socket.clone();
let handle = tokio::spawn(async move {
let listener = crate::uds::bind_hardened(&bound).expect("bind");
let exit = serve_until_idle(
&listener,
Arc::new(router),
RpcServeOptions::default(),
async move {
let _ = rx.await;
},
Some(IdleTracker::new(idle)),
)
.await;
let _ = std::fs::remove_file(&bound);
exit
});
(socket, tx, handle)
}
#[tokio::test]
async fn serve_until_idle_exits_when_the_window_elapses() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, handle) = spawn_idle_server(tmp.path(), Duration::from_millis(300));
await_socket(&socket).await;
let exit = tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("the loop must exit on its own once the window elapses")
.expect("join");
assert_eq!(exit, ServeExit::Idle);
assert!(
!socket.exists(),
"an idle exit must unlink the socket, or the next spawn cannot bind"
);
}
#[tokio::test]
async fn serve_until_idle_is_reset_by_an_answered_request() {
let tmp = tempfile::tempdir().expect("tempdir");
let window = Duration::from_millis(600);
let (socket, _stop, handle) = spawn_idle_server(tmp.path(), window);
await_socket(&socket).await;
tokio::time::sleep(window / 2).await;
let response = call(&socket, 1, "greet", json!({ "name": "ada" })).await;
assert!(response.error.is_none(), "unexpected error: {response:?}");
tokio::time::sleep(window).await;
assert!(
!handle.is_finished(),
"the answered request must have restarted the idle window"
);
let exit = tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("it must still exit once the restarted window elapses")
.expect("join");
assert_eq!(exit, ServeExit::Idle);
}
#[tokio::test]
async fn serve_until_idle_ignores_liveness_probes() {
let tmp = tempfile::tempdir().expect("tempdir");
let window = Duration::from_millis(400);
let (socket, _stop, handle) = spawn_idle_server(tmp.path(), window);
await_socket(&socket).await;
for _ in 0..4 {
assert!(crate::uds::socket_is_serving(&socket, Duration::from_millis(200)).await);
tokio::time::sleep(window / 8).await;
}
assert!(
!handle.is_finished(),
"liveness probes across half the window must not have exited it early"
);
let exit = tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("probes must not hold the window open")
.expect("join");
assert_eq!(exit, ServeExit::Idle);
}
fn liveness_router() -> RpcRouter {
greeting_router().typed_liveness(
"health",
|_req: ()| async move { Ok(json!({ "status": "ok" })) },
)
}
const POLL_CADENCE: Duration = Duration::from_millis(50);
const POLL_WINDOW: Duration = Duration::from_millis(200);
fn spawn_poller(socket: std::path::PathBuf, method: &'static str) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
loop {
let request = json!({ "jsonrpc": "2.0", "id": 1, "method": method, "params": {} });
let _ = crate::uds::send_framed_request::<_, RpcResponse>(
&socket,
&request,
Duration::from_secs(2),
)
.await;
tokio::time::sleep(POLL_CADENCE).await;
}
})
}
#[tokio::test]
async fn serve_until_idle_ignores_a_registered_liveness_method() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, _stop, handle) =
spawn_idle_server_with(tmp.path(), POLL_WINDOW, liveness_router());
await_socket(&socket).await;
let poller = spawn_poller(socket.clone(), "health");
let exit = tokio::time::timeout(Duration::from_secs(10), handle).await;
poller.abort();
let exit = exit
.expect("a liveness method must not re-arm the idle window")
.expect("join");
assert_eq!(exit, ServeExit::Idle);
}
#[tokio::test]
async fn serve_until_idle_is_held_open_by_a_non_liveness_call() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, stop, handle) = spawn_idle_server_with(tmp.path(), POLL_WINDOW, liveness_router());
await_socket(&socket).await;
let poller = spawn_poller(socket.clone(), "greet");
tokio::time::sleep(POLL_WINDOW * 5).await;
let still_running = !handle.is_finished();
poller.abort();
assert!(
still_running,
"an unmarked method's answer must still restart the idle window"
);
stop.send(()).expect("signal shutdown");
assert_eq!(
tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("the loop must stop on the signal")
.expect("join"),
ServeExit::Shutdown
);
}
#[test]
fn liveness_names_are_sorted_and_separate_from_the_method_table() {
let router = liveness_router();
assert_eq!(router.liveness_names().collect::<Vec<_>>(), vec!["health"]);
assert!(
router.method_names().any(|m| m == "health"),
"a liveness method is still a registered method: {router:?}"
);
}
#[test]
fn frame_is_liveness_reads_the_method_name_off_the_frame() {
let router = liveness_router();
assert!(router.frame_is_liveness(&frame(1, "health", json!({}))));
assert!(!router.frame_is_liveness(&frame(1, "greet", json!({ "name": "ada" }))));
assert!(
!router.frame_is_liveness(b"not json at all"),
"a frame about to be refused is not activity either way"
);
}
#[test]
fn frame_is_liveness_is_false_for_a_router_that_marks_nothing() {
assert!(!greeting_router().frame_is_liveness(&frame(1, "greet", json!({ "name": "ada" }))));
}
#[tokio::test]
async fn serve_until_without_a_policy_never_exits_on_its_own() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, stop, handle) =
spawn_server(tmp.path(), greeting_router(), RpcServeOptions::default());
await_socket(&socket).await;
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(
!handle.is_finished(),
"with no idle policy the loop runs until its shutdown future resolves"
);
stop.send(()).expect("signal shutdown");
handle.await.expect("join").expect("clean shutdown");
}
#[tokio::test]
async fn idle_tracker_counts_open_connections_and_restores_on_drop() {
let tracker = IdleTracker::new(Duration::from_secs(60));
assert_eq!(tracker.open_connections(), 0);
let mut answered = tracker.connection_opened();
let dropped = tracker.connection_opened();
assert_eq!(tracker.open_connections(), 2);
drop(dropped);
assert_eq!(
tracker.open_connections(),
1,
"a cancelled or panicking connection must still release its slot"
);
answered.answered();
answered.release().await;
assert_eq!(tracker.open_connections(), 0);
assert_eq!(tracker.timeout(), Duration::from_secs(60));
}
#[tokio::test]
async fn a_client_queued_when_the_idle_window_elapses_is_served_not_reset() {
let tmp = tempfile::tempdir().expect("tempdir");
let socket = tmp.path().join("drain.sock");
let listener = crate::uds::bind_hardened(&socket).expect("bind");
let tracker = IdleTracker::new(Duration::from_millis(1));
let mut client = tokio::net::UnixStream::connect(&socket)
.await
.expect("connect queues in the backlog; nothing is accepting yet");
let mut request = frame(7, "greet", json!({ "name": "queued" }));
request.push(b'\n');
client.write_all(&request).await.expect("write request");
tokio::time::sleep(Duration::from_millis(20)).await;
let serving = tokio::spawn(async move {
serve_until_idle(
&listener,
Arc::new(greeting_router()),
RpcServeOptions::default(),
std::future::pending::<()>(),
Some(tracker),
)
.await
});
let mut line = String::new();
let read = tokio::time::timeout(
Duration::from_secs(10),
BufReader::new(&mut client).read_line(&mut line),
)
.await
.expect("the queued client must not hang")
.expect("read response");
assert!(
read > 0,
"the queued connection was reset instead of served: the loop exited \
with it still in the backlog"
);
let response: RpcResponse = serde_json::from_str(&line).expect("parse response");
assert!(response.error.is_none(), "unexpected error: {response:?}");
let exit = tokio::time::timeout(Duration::from_secs(10), serving)
.await
.expect("the loop must still exit once the backlog really is empty")
.expect("join");
assert_eq!(exit, ServeExit::Idle);
}
#[derive(Clone, Default)]
struct EventLog(Arc<std::sync::Mutex<Vec<&'static str>>>);
impl EventLog {
fn record(&self, event: &'static str) {
if let Ok(mut events) = self.0.lock() {
events.push(event);
}
}
fn events(&self) -> Vec<&'static str> {
self.0.lock().map(|e| e.clone()).unwrap_or_default()
}
}
fn spawn_draining_server(
dir: &std::path::Path,
handler_for: Duration,
drain_budget: Duration,
) -> (
std::path::PathBuf,
EventLog,
tokio::sync::mpsc::UnboundedReceiver<()>,
tokio::sync::oneshot::Sender<()>,
tokio::task::JoinHandle<ServeExit>,
) {
let socket = dir.join("drain.sock");
let (stop_tx, stop_rx) = tokio::sync::oneshot::channel();
let (started_tx, started_rx) = tokio::sync::mpsc::unbounded_channel();
let log = EventLog::default();
let bound = socket.clone();
let handler_log = log.clone();
let loop_log = log.clone();
let handle = tokio::spawn(async move {
let router = RpcRouter::new().typed("slow", move |_req: ()| {
let log = handler_log.clone();
let started = started_tx.clone();
async move {
let _ = started.send(());
tokio::time::sleep(handler_for).await;
log.record("answered");
Ok::<_, RpcError>(json!({ "ok": true }))
}
});
let listener = crate::uds::bind_hardened(&bound).expect("bind");
let options = RpcServeOptions {
shutdown_drain: drain_budget,
..RpcServeOptions::default()
};
let exit = serve_until_idle(
&listener,
Arc::new(router),
options,
async move {
let _ = stop_rx.await;
},
None,
)
.await;
loop_log.record("returned");
let _ = std::fs::remove_file(&bound);
exit
});
(socket, log, started_rx, stop_tx, handle)
}
#[test]
fn default_serve_options_reserve_cleanup_time_inside_the_grace_window() {
let grace = crate::shutdown::termination_grace();
let drain = RpcServeOptions::default().shutdown_drain;
assert!(
drain + crate::shutdown::CLEANUP_RESERVE <= grace,
"the default drain ({drain:?}) plus the cleanup reserve ({:?}) must fit \
inside the termination grace window ({grace:?}), or the caller's \
post-serve work runs on borrowed time",
crate::shutdown::CLEANUP_RESERVE
);
assert!(
drain > Duration::ZERO,
"reserving cleanup time must not collapse the drain to nothing"
);
}
#[tokio::test]
async fn shutdown_drains_an_in_flight_connection_before_it_returns() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, log, mut started, stop, handle) = spawn_draining_server(
tmp.path(),
Duration::from_millis(400),
Duration::from_secs(5),
);
await_socket(&socket).await;
let dialled = socket.clone();
let client = tokio::spawn(async move { call(&dialled, 1, "slow", json!(null)).await });
started.recv().await.expect("the handler must start");
stop.send(()).expect("signal shutdown");
let response = tokio::time::timeout(Duration::from_secs(10), client)
.await
.expect("the in-flight request must complete during the drain")
.expect("join");
assert_eq!(
response.result,
Some(json!({ "ok": true })),
"a request already accepted must be answered, not dropped"
);
let exit = tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("the loop must return once the drain finishes")
.expect("join");
assert_eq!(exit, ServeExit::Shutdown);
assert_eq!(
log.events(),
vec!["answered", "returned"],
"the loop must not return — and so let its caller unlink — until the \
in-flight handler is done"
);
}
#[tokio::test]
async fn shutdown_refuses_a_connection_dialled_after_the_signal() {
let tmp = tempfile::tempdir().expect("tempdir");
let (socket, log, mut started, stop, handle) = spawn_draining_server(
tmp.path(),
Duration::from_millis(600),
Duration::from_secs(5),
);
await_socket(&socket).await;
let dialled = socket.clone();
let client = tokio::spawn(async move { call(&dialled, 1, "slow", json!(null)).await });
started.recv().await.expect("the handler must start");
stop.send(()).expect("signal shutdown");
let request = json!({ "jsonrpc": "2.0", "id": 2, "method": "slow", "params": null });
let refused = tokio::time::timeout(
Duration::from_secs(3),
crate::uds::send_framed_request::<_, RpcResponse>(
&socket,
&request,
Duration::from_secs(3),
),
)
.await
.expect("a post-signal dial must be refused, not left hanging in the backlog");
let during_refusal = log.events();
assert!(
matches!(refused, Err(crate::uds::UdsRpcError::NoResponse { .. })),
"a dial during the drain must be accepted and closed — not reset by a \
dropped listener, and not refused at connect by an unlinked path: \
{refused:?}"
);
assert!(
!during_refusal.contains(&"returned"),
"the refusal must land while the drain still owns the listener; the \
loop had already returned: {during_refusal:?}"
);
let _ = client.await;
let exit = tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("the loop must still return")
.expect("join");
assert_eq!(exit, ServeExit::Shutdown);
}
#[tokio::test]
async fn shutdown_returns_when_the_drain_budget_expires() {
let tmp = tempfile::tempdir().expect("tempdir");
let budget = Duration::from_millis(150);
let (socket, log, mut started, stop, handle) =
spawn_draining_server(tmp.path(), Duration::from_secs(30), budget);
await_socket(&socket).await;
let dialled = socket.clone();
let client = tokio::spawn(async move { call(&dialled, 1, "slow", json!(null)).await });
started.recv().await.expect("the handler must start");
let signalled = std::time::Instant::now();
stop.send(()).expect("signal shutdown");
let exit = tokio::time::timeout(Duration::from_secs(10), handle)
.await
.expect("the drain must be bounded")
.expect("join");
let elapsed = signalled.elapsed();
assert_eq!(exit, ServeExit::Shutdown);
assert!(
elapsed < Duration::from_secs(5),
"the loop must return on the budget, not on the handler: {elapsed:?}"
);
assert_eq!(
log.events(),
vec!["returned"],
"the handler had not finished, so only the loop's own event is recorded"
);
client.abort();
}