From 806233cab7a89f0796a34e85473b2c517eed84e5 Mon Sep 17 00:00:00 2001 From: Hunter Sadler Date: Thu, 1 Oct 2026 16:27:29 -0600 Subject: [PATCH 1/2] feat(rust): bind weak ping handles before client initialization Share normal request correlation with transport-owned keepalive without retaining the client driver. Keep managed-host request IDs across retries and clean up cancelled requests deterministically. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- clients/rust/crates/ahp/README.md | 31 + clients/rust/crates/ahp/src/client.rs | 307 ++++++--- clients/rust/crates/ahp/src/client_tests.rs | 245 ++++++++ clients/rust/crates/ahp/src/hosts/runtime.rs | 10 +- clients/rust/crates/ahp/src/lib.rs | 9 +- clients/rust/crates/ahp/src/transport.rs | 21 + clients/rust/crates/ahp/tests/weak_ping.rs | 589 ++++++++++++++++++ .../20261001-rust-weak-ping-binding.json | 5 + 8 files changed, 1124 insertions(+), 93 deletions(-) create mode 100644 clients/rust/crates/ahp/src/client_tests.rs create mode 100644 clients/rust/crates/ahp/tests/weak_ping.rs create mode 100644 docs/.changes/20261001-rust-weak-ping-binding.json diff --git a/clients/rust/crates/ahp/README.md b/clients/rust/crates/ahp/README.md index 343027901..0b998db20 100644 --- a/clients/rust/crates/ahp/README.md +++ b/clients/rust/crates/ahp/README.md @@ -65,6 +65,37 @@ impl Transport for MyTransport { See `tests/client_roundtrip.rs` for a complete in-memory example. +### Transport-owned keepalive + +Override the optional synchronous +`Transport::bind_client(&mut self, ping: ahp::WeakPingHandle)` callback to pass +the handle to a transport-owned keepalive task. `Client::connect` binds it once, +before transport I/O starts or any `initialize` / `reconnect` request can be +sent. `BoxedTransport` forwards the callback, including for managed hosts. +Subsequent handshakes on the same client do not bind it again. + +`WeakPingHandle::ping().await` uses the same root-channel request, ID allocator, +response map, and configured timeout as `Client::ping`. It can therefore keep an +initialized connection alive while session discovery is still pending. The +transport owns the keepalive timing and any negotiation or activity gating; +binding does not start keepalive. + +Neither a retained handle nor an in-flight weak ping owns the client driver. +Explicit shutdown, transport closure, or dropping the last `Client` resolves +weak pings with `ClientError::Shutdown`. Request timeouts remain +`ClientError::Cancelled`, and server errors remain `ClientError::Rpc`. Cancelling +a request future removes its pending response entry without retracting an +already-sent request. Existing transports need no changes: binding defaults to +a no-op on both `Transport` and the object-safe `DynTransport` adapter. + +Managed hosts retain their request ID allocator across connection attempts for +the lifetime of one host supervisor, including failed handshakes. Late replies +from an earlier transport cannot match a newly allocated request on that +logical host. Independent hosts and standalone `Client::connect` calls still +start independent ID sequences. Exhausting the `u64` request ID space fails +explicitly with `ClientError::Transport(TransportError::Protocol(_))` rather +than reusing an ID. + ## See also - [`ahp-types`](https://crates.io/crates/ahp-types) — wire types only (no I/O) diff --git a/clients/rust/crates/ahp/src/client.rs b/clients/rust/crates/ahp/src/client.rs index 4b9e5f05e..a0d108862 100644 --- a/clients/rust/crates/ahp/src/client.rs +++ b/clients/rust/crates/ahp/src/client.rs @@ -26,7 +26,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::{ atomic::{AtomicU64, Ordering}, - Arc, + Arc, Weak, }; use std::time::Duration; @@ -53,10 +53,10 @@ use ahp_types::notifications::{ }; use serde::{de::DeserializeOwned, Serialize}; use serde_json::Value; -use tokio::sync::{broadcast, mpsc, oneshot, Mutex}; +use tokio::sync::{broadcast, mpsc, oneshot, watch, Mutex}; use tokio::task::JoinHandle; -use crate::error::ClientError; +use crate::error::{ClientError, TransportError}; use crate::transport::{Transport, TransportMessage}; /// Default size of a per-subscription broadcast channel. Consumers that @@ -189,8 +189,30 @@ pub struct DispatchHandle { type PendingMap = HashMap>>; +pub(crate) struct RequestIds { + next: std::sync::Mutex>, +} + +impl RequestIds { + pub(crate) fn new() -> Self { + Self { + next: std::sync::Mutex::new(Some(1)), + } + } + + fn allocate(&self) -> Result { + let mut next = self.next.lock().expect("request IDs mutex poisoned"); + let id = + next.ok_or_else(|| TransportError::Protocol("request ID space exhausted".into()))?; + *next = id.checked_add(1); + Ok(id) + } +} + struct Shared { - pending: Mutex, + // A synchronous lock lets request-future Drop remove entries without an await. + pending: std::sync::Mutex, + closed: watch::Sender, subscriptions: Mutex>>, /// Top-level all-events broadcast. /// @@ -201,7 +223,7 @@ struct Shared { /// alive inside the still-`Arc`-held `Shared`). all_events: std::sync::Mutex>>, outbound: mpsc::Sender, - next_id: AtomicU64, + request_ids: Arc, next_client_seq: AtomicU64, config: ClientConfig, /// Handler for inbound server-initiated requests (the symmetrical @@ -214,6 +236,138 @@ enum Outbound { Shutdown, } +impl Shared { + fn stop_requests(&self, message: Option<&str>) { + let mut pending = self.pending.lock().expect("pending mutex poisoned"); + self.closed.send_replace(true); + for (_, tx) in pending.drain() { + if let Some(message) = message { + let _ = tx.send(Err(JsonRpcError { + code: -32000, + message: message.into(), + data: None, + })); + } + } + } + + async fn request(&self, method: &str, params: P, weak: bool) -> Result + where + P: Serialize, + R: DeserializeOwned, + { + let mut closed = self.closed.subscribe(); + if *closed.borrow() { + return Err(ClientError::Shutdown); + } + let id = self.request_ids.allocate()?; + let params_val = serde_json::to_value(¶ms)?; + let params_any = if params_val.is_null() { + None + } else { + Some(ahp_types::common::AnyValue::from(params_val)) + }; + let req = JsonRpcMessage::Request(JsonRpcRequest { + jsonrpc: JsonRpcVersion::V2, + id, + method: method.into(), + params: params_any, + }); + + let (tx, rx) = oneshot::channel(); + { + let mut pending = self.pending.lock().expect("pending mutex poisoned"); + if *closed.borrow() { + return Err(ClientError::Shutdown); + } + pending.insert(id, tx); + } + let _pending = PendingRequest { shared: self, id }; + + let response = async { + self.outbound + .send(Outbound::Message(req)) + .await + .map_err(|_| ClientError::Shutdown)?; + + let result = match self.config.default_request_timeout { + Some(dur) => tokio::time::timeout(dur, rx) + .await + .map_err(|_| ClientError::Cancelled)?, + None => rx.await, + }; + match result { + Ok(Ok(value)) => Ok(serde_json::from_value(value)?), + Ok(Err(e)) => Err(ClientError::Rpc(e)), + Err(_) => Err(ClientError::Shutdown), + } + }; + tokio::select! { + biased; + _ = closed.wait_for(|closed| *closed), if weak => Err(ClientError::Shutdown), + result = response => result, + } + } + + async fn ping(&self, weak: bool) -> Result<(), ClientError> { + #[derive(Serialize)] + struct PingParams { + channel: &'static str, + } + self.request( + "ping", + PingParams { + channel: ROOT_RESOURCE_URI, + }, + weak, + ) + .await + } +} + +struct PendingRequest<'a> { + shared: &'a Shared, + id: u64, +} + +impl Drop for PendingRequest<'_> { + fn drop(&mut self) { + self.shared + .pending + .lock() + .expect("pending mutex poisoned") + .remove(&self.id); + } +} + +/// Non-owning handle for client-correlated keepalive requests. +/// +/// Supplied to [`Transport::bind_client`] before transport I/O starts. +/// Cloning the handle or awaiting [`Self::ping`] never keeps the client's +/// background driver alive. Pings share the client's normal request IDs, +/// response correlation, and configured request timeout. +#[derive(Clone)] +pub struct WeakPingHandle { + shared: Weak, +} + +impl WeakPingHandle { + /// Send the same root-channel request as [`Client::ping`]. + /// + /// Returns [`ClientError::Shutdown`] after explicit shutdown, transport + /// closure, or dropping the last [`Client`], including for an unanswered + /// in-flight ping. Timeouts return [`ClientError::Cancelled`]; server errors + /// return [`ClientError::Rpc`]. Dropping this future removes its pending + /// response entry but does not retract a request already sent. + /// + /// The caller decides when keepalive is appropriate for the connection; + /// binding the handle does not start a ping or negotiate keepalive. + pub async fn ping(&self) -> Result<(), ClientError> { + let shared = self.shared.upgrade().ok_or(ClientError::Shutdown)?; + shared.ping(true).await + } +} + // ─── Server-initiated request handling ─────────────────────────────────────── /// Future returned by a [`ServerRequestHandler`]. @@ -354,10 +508,14 @@ pub struct Client { struct DriveHandle { handle: Mutex>>, + shared: Weak, } impl Drop for DriveHandle { fn drop(&mut self) { + if let Some(shared) = self.shared.upgrade() { + shared.stop_requests(None); + } if let Ok(mut guard) = self.handle.try_lock() { if let Some(h) = guard.take() { h.abort(); @@ -372,42 +530,51 @@ impl Client { pub async fn connect( transport: T, config: ClientConfig, + ) -> Result { + Self::connect_with_request_ids(transport, config, Arc::new(RequestIds::new())).await + } + + pub(crate) async fn connect_with_request_ids( + mut transport: T, + config: ClientConfig, + request_ids: Arc, ) -> Result { let (outbound_tx, outbound_rx) = mpsc::channel::(64); let (all_events_tx, _) = broadcast::channel::(config.subscription_buffer); + let (closed, _) = watch::channel(false); let shared = Arc::new(Shared { - pending: Mutex::new(HashMap::new()), + pending: std::sync::Mutex::new(HashMap::new()), + closed, subscriptions: Mutex::new(HashMap::new()), all_events: std::sync::Mutex::new(Some(all_events_tx)), outbound: outbound_tx, - next_id: AtomicU64::new(1), + request_ids, next_client_seq: AtomicU64::new(1), config, server_request_handler: std::sync::Mutex::new(None), }); + transport.bind_client(WeakPingHandle { + shared: Arc::downgrade(&shared), + }); let handle = tokio::spawn(drive_transport(transport, shared.clone(), outbound_rx)); + let reader = Arc::new(DriveHandle { + handle: Mutex::new(Some(handle)), + shared: Arc::downgrade(&shared), + }); Ok(Self { shared, - _reader: Arc::new(DriveHandle { - handle: Mutex::new(Some(handle)), - }), + _reader: reader, }) } - /// Gracefully shut down the client, aborting any in-flight requests - /// with [`ClientError::Shutdown`]. + /// Gracefully shut down the client. + /// + /// In-flight normal requests retain the `-32000` [`ClientError::Rpc`] + /// shutdown error; weak pings resolve with [`ClientError::Shutdown`]. pub async fn shutdown(&self) { + self.shared.stop_requests(Some("client shut down")); let _ = self.shared.outbound.send(Outbound::Shutdown).await; - // Fail any pending in-flight requests. - let mut pending = self.shared.pending.lock().await; - for (_, tx) in pending.drain() { - let _ = tx.send(Err(JsonRpcError { - code: -32000, - message: "client shut down".into(), - data: None, - })); - } } /// Send a JSON-RPC request and await its result. @@ -416,53 +583,7 @@ impl Client { P: Serialize, R: DeserializeOwned, { - let id = self.shared.next_id.fetch_add(1, Ordering::Relaxed); - let params_val = serde_json::to_value(¶ms)?; - let params_any = if params_val.is_null() { - None - } else { - Some(ahp_types::common::AnyValue::from(params_val)) - }; - let req = JsonRpcMessage::Request(JsonRpcRequest { - jsonrpc: JsonRpcVersion::V2, - id, - method: method.into(), - params: params_any, - }); - - let (tx, rx) = oneshot::channel(); - { - let mut pending = self.shared.pending.lock().await; - pending.insert(id, tx); - } - - if self - .shared - .outbound - .send(Outbound::Message(req)) - .await - .is_err() - { - self.shared.pending.lock().await.remove(&id); - return Err(ClientError::Shutdown); - } - - let result = match self.shared.config.default_request_timeout { - Some(dur) => match tokio::time::timeout(dur, rx).await { - Ok(r) => r, - Err(_) => { - self.shared.pending.lock().await.remove(&id); - return Err(ClientError::Cancelled); - } - }, - None => rx.await, - }; - - match result { - Ok(Ok(value)) => Ok(serde_json::from_value(value)?), - Ok(Err(e)) => Err(ClientError::Rpc(e)), - Err(_) => Err(ClientError::Shutdown), - } + self.shared.request(method, params, false).await } /// Send a JSON-RPC notification (fire-and-forget). @@ -539,17 +660,7 @@ impl Client { /// server responds regardless of whether `initialize` has completed or any /// subscriptions are held. pub async fn ping(&self) -> Result<(), ClientError> { - #[derive(Serialize)] - struct PingParams { - channel: &'static str, - } - self.request( - "ping", - PingParams { - channel: ROOT_RESOURCE_URI, - }, - ) - .await + self.shared.ping(false).await } /// Subscribe to a URI and obtain a handle that streams @@ -842,6 +953,7 @@ async fn drive_transport( shared: Arc, mut outbound: mpsc::Receiver, ) { + let _requests = DriverRequests(shared.clone()); loop { tokio::select! { outbound_msg = outbound.recv() => { @@ -878,15 +990,8 @@ async fn drive_transport( } } - // Teardown: close everything so waiters see Shutdown. - let mut pending = shared.pending.lock().await; - for (_, tx) in pending.drain() { - let _ = tx.send(Err(JsonRpcError { - code: -32000, - message: "transport closed".into(), - data: None, - })); - } + // Teardown: close everything and fail outstanding requests. + shared.stop_requests(Some("transport closed")); let mut subs = shared.subscriptions.lock().await; subs.clear(); // Drop the top-level fan-out sender so any active @@ -898,15 +1003,33 @@ async fn drive_transport( } } +struct DriverRequests(Arc); + +impl Drop for DriverRequests { + fn drop(&mut self) { + self.0.stop_requests(None); + } +} + async fn dispatch_inbound(shared: &Arc, msg: JsonRpcMessage) { match msg { JsonRpcMessage::SuccessResponse(r) => { - if let Some(tx) = shared.pending.lock().await.remove(&r.id) { + if let Some(tx) = shared + .pending + .lock() + .expect("pending mutex poisoned") + .remove(&r.id) + { let _ = tx.send(Ok(r.result)); } } JsonRpcMessage::ErrorResponse(r) => { - if let Some(tx) = shared.pending.lock().await.remove(&r.id) { + if let Some(tx) = shared + .pending + .lock() + .expect("pending mutex poisoned") + .remove(&r.id) + { let _ = tx.send(Err(r.error)); } } @@ -1020,3 +1143,7 @@ async fn fan_out(shared: &Shared, channel: &Uri, event: SubscriptionEvent) { } } } + +#[cfg(test)] +#[path = "client_tests.rs"] +mod tests; diff --git a/clients/rust/crates/ahp/src/client_tests.rs b/clients/rust/crates/ahp/src/client_tests.rs new file mode 100644 index 000000000..32d73fa7a --- /dev/null +++ b/clients/rust/crates/ahp/src/client_tests.rs @@ -0,0 +1,245 @@ +#![allow(clippy::panic, clippy::unwrap_used)] + +use super::*; +use crate::BoxedTransport; + +struct TestTransport { + sent: mpsc::Sender, + received: mpsc::Receiver, + bound: Option, +} + +impl Transport for TestTransport { + fn bind_client(&mut self, ping: WeakPingHandle) { + self.bound = Some(ping); + } + + async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { + self.sent + .send(message) + .await + .map_err(|_| TransportError::Closed) + } + + async fn recv(&mut self) -> Result, TransportError> { + Ok(self.received.recv().await) + } +} + +async fn client( + config: ClientConfig, +) -> ( + Client, + mpsc::Receiver, + mpsc::Sender, +) { + let (sent, rx) = mpsc::channel(1); + let (tx, received) = mpsc::channel(1); + let transport = TestTransport { + sent, + received, + bound: None, + }; + ( + Client::connect(BoxedTransport::new(transport), config) + .await + .unwrap(), + rx, + tx, + ) +} + +#[tokio::test] +async fn cancelled_and_timed_out_requests_remove_pending_entries() { + let (client, mut sent, _received) = client(ClientConfig { + default_request_timeout: None, + ..ClientConfig::default() + }) + .await; + let weak = WeakPingHandle { + shared: Arc::downgrade(&client.shared), + }; + for is_weak in [false, true] { + let mut request = Box::pin(async { + if is_weak { + weak.ping().await + } else { + client.ping().await + } + }); + tokio::select! { + result = &mut request => panic!("premature result: {result:?}"), + _ = sent.recv() => {}, + } + assert_eq!(client.shared.pending.lock().unwrap().len(), 1); + drop(request); + assert!(client.shared.pending.lock().unwrap().is_empty()); + } + drop(client); + assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) + .await + .unwrap() + .is_none()); + assert!( + weak.shared.upgrade().is_none(), + "Shared leaked after driver drop" + ); + + let (client, mut sent, _received) = self::client(ClientConfig { + default_request_timeout: Some(Duration::ZERO), + ..ClientConfig::default() + }) + .await; + let weak = WeakPingHandle { + shared: Arc::downgrade(&client.shared), + }; + assert!(matches!(weak.ping().await, Err(ClientError::Cancelled))); + assert!(client.shared.pending.lock().unwrap().is_empty()); + assert!(sent.recv().await.is_some()); + assert!(matches!(client.ping().await, Err(ClientError::Cancelled))); + assert!(client.shared.pending.lock().unwrap().is_empty()); + assert!(sent.recv().await.is_some()); + drop(client); + assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) + .await + .unwrap() + .is_none()); +} + +#[tokio::test] +async fn owner_drop_clears_unanswered_weak_request_and_releases_shared_allocator() { + let (client, mut sent, _received) = client(ClientConfig { + default_request_timeout: None, + ..ClientConfig::default() + }) + .await; + let weak = WeakPingHandle { + shared: Arc::downgrade(&client.shared), + }; + let ids = Arc::downgrade(&client.shared.request_ids); + let mut request = Box::pin(weak.ping()); + tokio::select! { + result = &mut request => panic!("premature result: {result:?}"), + _ = sent.recv() => {}, + } + drop(client); + assert!(matches!(request.await, Err(ClientError::Shutdown))); + assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) + .await + .unwrap() + .is_none()); + assert!(weak.shared.upgrade().is_none()); + assert!(ids.upgrade().is_none()); +} + +#[tokio::test] +async fn request_id_exhaustion_never_wraps_or_enqueues_another_request() { + let (client, mut sent, _received) = client(ClientConfig { + default_request_timeout: Some(Duration::ZERO), + ..ClientConfig::default() + }) + .await; + *client.shared.request_ids.next.lock().unwrap() = Some(u64::MAX); + let weak = WeakPingHandle { + shared: Arc::downgrade(&client.shared), + }; + assert!(matches!(weak.ping().await, Err(ClientError::Cancelled))); + let JsonRpcMessage::Request(request) = sent.recv().await.unwrap().into_parsed().unwrap() else { + panic!("expected request"); + }; + assert_eq!(request.id, u64::MAX); + for is_weak in [false, true] { + let result = if is_weak { + weak.ping().await + } else { + client.ping().await + }; + assert!( + matches!(result, Err(ClientError::Transport(TransportError::Protocol(message))) if message == "request ID space exhausted") + ); + assert!(client.shared.pending.lock().unwrap().is_empty()); + assert!(sent.try_recv().is_err()); + } + assert!(client.shared.request_ids.next.lock().unwrap().is_none()); + client.shutdown().await; + assert!(matches!(weak.ping().await, Err(ClientError::Shutdown))); + drop(client); + assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) + .await + .unwrap() + .is_none()); +} + +struct BlockedTransport { + entered: Option>, + dropped: Option>, +} + +impl Transport for BlockedTransport { + async fn send(&mut self, _: TransportMessage) -> Result<(), TransportError> { + self.entered.take().unwrap().send(()).unwrap(); + std::future::pending().await + } + + async fn recv(&mut self) -> Result, TransportError> { + std::future::pending().await + } +} + +impl Drop for BlockedTransport { + fn drop(&mut self) { + self.dropped.take().unwrap().send(()).unwrap(); + } +} + +#[tokio::test] +async fn owner_drop_cancels_weak_ping_blocked_on_full_outbound_queue() { + let (entered, send_started) = oneshot::channel(); + let (dropped, transport_dropped) = oneshot::channel(); + let client = Client::connect( + BoxedTransport::new(BlockedTransport { + entered: Some(entered), + dropped: Some(dropped), + }), + ClientConfig { + default_request_timeout: None, + ..ClientConfig::default() + }, + ) + .await + .unwrap(); + let weak = WeakPingHandle { + shared: Arc::downgrade(&client.shared), + }; + client.notify("block", ()).await.unwrap(); + tokio::time::timeout(Duration::from_secs(2), send_started) + .await + .unwrap() + .unwrap(); + for _ in 0..64 { + client.notify("queued", ()).await.unwrap(); + } + assert_eq!(client.shared.outbound.capacity(), 0); + let mut request = Box::pin(weak.ping()); + std::future::poll_fn(|cx| { + assert!(request.as_mut().poll(cx).is_pending()); + std::task::Poll::Ready(()) + }) + .await; + assert_eq!(client.shared.pending.lock().unwrap().len(), 1); + drop(client); + assert!(weak + .shared + .upgrade() + .unwrap() + .pending + .lock() + .unwrap() + .is_empty()); + assert!(matches!(request.await, Err(ClientError::Shutdown))); + tokio::time::timeout(Duration::from_secs(2), transport_dropped) + .await + .unwrap() + .unwrap(); + assert!(weak.shared.upgrade().is_none()); +} diff --git a/clients/rust/crates/ahp/src/hosts/runtime.rs b/clients/rust/crates/ahp/src/hosts/runtime.rs index 6bfa33d7d..9d8e48af2 100644 --- a/clients/rust/crates/ahp/src/hosts/runtime.rs +++ b/clients/rust/crates/ahp/src/hosts/runtime.rs @@ -18,6 +18,7 @@ use ahp_types::state::{RootState, SessionSummary, SnapshotState}; use tokio::sync::{broadcast, mpsc, oneshot, Notify}; use tokio::task::JoinHandle; +use crate::client::RequestIds; use crate::reducers::{apply_action_to_root, ReduceOutcome}; use crate::{Client, ClientError, ClientEvent, DispatchHandle, SubscriptionEvent}; @@ -113,6 +114,7 @@ pub(super) fn spawn( let (cmd_tx, cmd_rx) = mpsc::channel(32); let runtime = HostRuntime { client_id: resolved_client_id, + request_ids: Arc::new(RequestIds::new()), config, cmd_rx, shared: shared.clone(), @@ -133,6 +135,7 @@ pub(super) fn spawn( struct HostRuntime { config: HostConfig, client_id: String, + request_ids: Arc, cmd_rx: mpsc::Receiver, shared: Arc, fan_out: broadcast::Sender, @@ -260,7 +263,12 @@ impl HostRuntime { .open_transport(self.config.id.clone()) .await?; - let client = Client::connect(transport, self.config.client_config.clone()).await?; + let client = Client::connect_with_request_ids( + transport, + self.config.client_config.clone(), + self.request_ids.clone(), + ) + .await?; // Attach the events receiver BEFORE the initialize/reconnect // handshake so any notifications the server pushes between the diff --git a/clients/rust/crates/ahp/src/lib.rs b/clients/rust/crates/ahp/src/lib.rs index 013229405..9ee30d67e 100644 --- a/clients/rust/crates/ahp/src/lib.rs +++ b/clients/rust/crates/ahp/src/lib.rs @@ -138,8 +138,12 @@ //! //! All async client APIs are cancel-safe at await points. The background //! driver is owned by the [`Client`] and aborted when the last clone is -//! dropped, or when [`Client::shutdown`] is called. In-flight requests -//! resolve with [`ClientError::Shutdown`] in either case. +//! dropped; [`Client::shutdown`] requests graceful transport closure. +//! Dropping the last client resolves in-flight requests with +//! [`ClientError::Shutdown`]. Explicit shutdown and transport closure retain +//! the normal `-32000` [`ClientError::Rpc`] teardown error. Non-owning +//! [`WeakPingHandle`] requests resolve with [`ClientError::Shutdown`] for +//! each of these lifecycle endings. #![forbid(unsafe_code)] #![warn(missing_docs)] @@ -157,6 +161,7 @@ pub use ahp_types; pub use client::{ Client, ClientConfig, ClientEvent, ClientEventStream, DispatchHandle, ResourceRequestHandlers, ServerRequestFuture, ServerRequestHandler, SessionSubscription, SubscriptionEvent, + WeakPingHandle, }; pub use error::{ClientError, TransportError}; pub use multi_host_state_mirror::{HostedResourceKey, MultiHostStateMirror}; diff --git a/clients/rust/crates/ahp/src/transport.rs b/clients/rust/crates/ahp/src/transport.rs index 3bf3a552e..3a1e4d790 100644 --- a/clients/rust/crates/ahp/src/transport.rs +++ b/clients/rust/crates/ahp/src/transport.rs @@ -64,6 +64,7 @@ use std::future::Future; use std::pin::Pin; +use crate::client::WeakPingHandle; use crate::error::TransportError; use ahp_types::messages::JsonRpcMessage; @@ -112,6 +113,15 @@ impl TransportMessage { /// client sends indefinitely until the underlying connection closes, /// and `recv` signals closure by returning `None`. pub trait Transport: Send + 'static { + /// Bind a non-owning ping handle before the client starts transport I/O. + /// + /// Called once by [`crate::Client::connect`], not on subsequent + /// initialize/reconnect requests on the same client. A transport may pass + /// the handle to its keepalive task without keeping the client alive. + /// The transport remains responsible for keepalive timing and negotiation. + /// The default implementation is a no-op. + fn bind_client(&mut self, _ping: WeakPingHandle) {} + /// Send a single message. /// /// Errors returned here are typically fatal for the transport @@ -154,6 +164,9 @@ pub trait Transport: Send + 'static { /// allocation cost (typically: registries that hold one transport per /// host). pub trait DynTransport: Send + 'static { + /// Object-safe analogue of [`Transport::bind_client`]. + fn bind_client(&mut self, _ping: WeakPingHandle) {} + /// Object-safe analogue of [`Transport::send`]. fn send<'a>( &'a mut self, @@ -172,6 +185,10 @@ pub trait DynTransport: Send + 'static { } impl DynTransport for T { + fn bind_client(&mut self, ping: WeakPingHandle) { + ::bind_client(self, ping); + } + fn send<'a>( &'a mut self, msg: TransportMessage, @@ -247,6 +264,10 @@ impl std::fmt::Debug for BoxedTransport { } impl Transport for BoxedTransport { + fn bind_client(&mut self, ping: WeakPingHandle) { + self.inner.bind_client(ping); + } + fn send( &mut self, msg: TransportMessage, diff --git a/clients/rust/crates/ahp/tests/weak_ping.rs b/clients/rust/crates/ahp/tests/weak_ping.rs new file mode 100644 index 000000000..7c2131120 --- /dev/null +++ b/clients/rust/crates/ahp/tests/weak_ping.rs @@ -0,0 +1,589 @@ +#![allow(clippy::panic, clippy::unwrap_used)] + +use std::collections::VecDeque; +use std::future::Future; +use std::pin::Pin; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; +use std::time::Duration; + +use ahp::hosts::{HostConfig, HostEvent, HostId, MultiHostClient, ReconnectPolicy}; +use ahp::{ + BoxedTransport, Client, ClientConfig, ClientError, DynTransport, Transport, TransportError, + TransportMessage, WeakPingHandle, +}; +use ahp_types::messages::{ + JsonRpcError, JsonRpcErrorResponse, JsonRpcMessage, JsonRpcRequest, JsonRpcSuccessResponse, + JsonRpcVersion, +}; +use serde_json::{json, Value}; +use tokio::sync::{mpsc, Mutex}; + +struct BoundTransport { + tx: mpsc::Sender, + rx: mpsc::Receiver, + bindings: mpsc::UnboundedSender, + bind_count: Arc, +} + +struct Peer { + tx: mpsc::Sender, + rx: mpsc::Receiver, + bindings: mpsc::UnboundedReceiver, + bind_count: Arc, +} + +fn pair() -> (BoundTransport, Peer) { + let (to_peer, rx) = mpsc::channel(16); + let (tx, from_peer) = mpsc::channel(16); + let (bindings, bound) = mpsc::unbounded_channel(); + let bind_count = Arc::new(AtomicUsize::new(0)); + ( + BoundTransport { + tx: to_peer, + rx: from_peer, + bindings, + bind_count: bind_count.clone(), + }, + Peer { + tx, + rx, + bindings: bound, + bind_count, + }, + ) +} + +impl Transport for BoundTransport { + fn bind_client(&mut self, ping: WeakPingHandle) { + self.bind_count.fetch_add(1, Ordering::SeqCst); + self.bindings.send(ping).unwrap(); + } + + async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { + assert_eq!(self.bind_count.load(Ordering::SeqCst), 1); + self.tx + .send(message) + .await + .map_err(|_| TransportError::Closed) + } + + async fn recv(&mut self) -> Result, TransportError> { + Ok(self.rx.recv().await) + } +} + +struct LegacyDynTransport { + tx: mpsc::Sender, + rx: mpsc::Receiver, +} + +impl DynTransport for LegacyDynTransport { + fn send<'a>( + &'a mut self, + message: TransportMessage, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + self.tx + .send(message) + .await + .map_err(|_| TransportError::Closed) + }) + } + + fn recv<'a>( + &'a mut self, + ) -> Pin, TransportError>> + Send + 'a>> + { + Box::pin(async move { Ok(self.rx.recv().await) }) + } + + fn close<'a>( + &'a mut self, + ) -> Pin> + Send + 'a>> { + Box::pin(async { Ok(()) }) + } +} + +#[tokio::test] +async fn legacy_object_safe_transport_needs_no_binding_implementation() { + let (transport, mut peer) = pair(); + let legacy = LegacyDynTransport { + tx: transport.tx, + rx: transport.rx, + }; + let client = Client::connect(BoxedTransport::from_dyn(Box::new(legacy)), no_timeout()) + .await + .unwrap(); + let mut request = Box::pin(client.ping()); + let sent = tokio::select! { + result = &mut request => panic!("premature result: {result:?}"), + sent = peer.request() => sent, + }; + peer.reply(sent.id, Value::Null).await; + bounded(request).await.unwrap(); + drop(client); + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +async fn bounded(future: impl Future) -> T { + tokio::time::timeout(Duration::from_secs(2), future) + .await + .expect("fixture timed out") +} + +impl Peer { + async fn request(&mut self) -> JsonRpcRequest { + let wire = bounded(self.rx.recv()).await.expect("driver closed"); + let JsonRpcMessage::Request(request) = wire.into_parsed().unwrap() else { + panic!("expected request") + }; + request + } + + async fn reply(&self, id: u64, result: Value) { + bounded( + self.tx + .send(TransportMessage::Parsed(JsonRpcMessage::SuccessResponse( + JsonRpcSuccessResponse { + jsonrpc: JsonRpcVersion::V2, + id, + result, + }, + ))), + ) + .await + .unwrap(); + } + + async fn bound(&mut self) -> WeakPingHandle { + bounded(self.bindings.recv()).await.unwrap() + } +} + +fn no_timeout() -> ClientConfig { + ClientConfig { + default_request_timeout: None, + ..ClientConfig::default() + } +} + +#[tokio::test] +async fn boxed_binding_precedes_handshakes_and_is_not_repeated() { + for from_dyn in [false, true] { + let (transport, mut peer) = pair(); + let transport = if from_dyn { + let dynamic: Box = Box::new(transport); + BoxedTransport::from_dyn(dynamic) + } else { + BoxedTransport::new(transport) + }; + let client = Client::connect(transport, no_timeout()).await.unwrap(); + let ping = peer + .bindings + .try_recv() + .expect("bind before connect returns"); + assert!(peer.rx.try_recv().is_err(), "binding must not send a ping"); + + for method in ["initialize", "reconnect", "initialize"] { + let request = client.request::<_, Value>(method, json!({})); + tokio::pin!(request); + let sent = tokio::select! { + request = &mut request => panic!("premature result: {request:?}"), + sent = peer.request() => sent, + }; + assert_eq!(sent.method, method); + peer.reply(sent.id, json!({"handshake": method})).await; + assert_eq!(bounded(request).await.unwrap()["handshake"], method); + } + assert_eq!(peer.bind_count.load(Ordering::SeqCst), 1); + assert!(peer.bindings.try_recv().is_err()); + drop(client); + assert!(bounded(peer.rx.recv()).await.is_none()); + assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); + } +} + +#[tokio::test] +async fn weak_and_normal_pings_share_ids_and_out_of_order_correlation() { + let (transport, mut peer) = pair(); + let client = Client::connect(BoxedTransport::new(transport), no_timeout()) + .await + .unwrap(); + let ping = peer.bound().await; + let ordinary = client.request::<_, Value>("listSessions", json!({"channel": "ahp-root://"})); + let heartbeat = ping.ping(); + let normal_ping = client.ping(); + tokio::pin!(ordinary, heartbeat, normal_ping); + + let requests = async { + let mut requests = Vec::new(); + for _ in 0..3 { + requests.push(peer.request().await); + } + requests + }; + let requests = tokio::select! { + result = &mut ordinary => panic!("premature result: {result:?}"), + result = &mut heartbeat => panic!("premature result: {result:?}"), + result = &mut normal_ping => panic!("premature result: {result:?}"), + requests = requests => requests, + }; + let mut ids: Vec<_> = requests.iter().map(|request| request.id).collect(); + ids.sort_unstable(); + assert_eq!(ids, [1, 2, 3]); + let discovery = requests + .iter() + .find(|r| r.method == "listSessions") + .unwrap(); + for request in requests.iter().rev().filter(|r| r.method == "ping") { + assert_eq!(request.params.as_ref().unwrap()["channel"], "ahp-root://"); + peer.reply(request.id, Value::Null).await; + } + bounded(heartbeat).await.unwrap(); + bounded(normal_ping).await.unwrap(); + // Discovery is still unanswered while both heartbeat responses resolve. + assert!(tokio::time::timeout(Duration::ZERO, &mut ordinary) + .await + .is_err()); + peer.reply(discovery.id, json!({"items": []})).await; + assert_eq!(bounded(ordinary).await.unwrap(), json!({"items": []})); + client.shutdown().await; + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn weak_ping_works_while_managed_host_discovery_blocks_client_access() { + let (transport, mut peer) = pair(); + let transport = Arc::new(Mutex::new(Some(transport))); + let config = HostConfig::new("discovering", "Discovering host", move |_| { + let transport = transport.clone(); + async move { Ok(BoxedTransport::new(transport.lock().await.take().unwrap())) } + }) + .with_client_config(no_timeout()) + .with_reconnect_policy(ReconnectPolicy::disabled()); + let multi = MultiHostClient::new(); + multi.add_host(config).await.unwrap(); + let ping = peer.bound().await; + let initialize = peer.request().await; + assert_eq!(initialize.method, "initialize"); + peer.reply( + initialize.id, + json!({"protocolVersion": ahp_types::PROTOCOL_VERSION, "serverSeq": 0, "snapshots": []}), + ) + .await; + let discovery = peer.request().await; + assert_eq!(discovery.method, "listSessions"); + assert!(multi.client(&HostId::from("discovering")).await.is_none()); + + let heartbeat = ping.ping(); + tokio::pin!(heartbeat); + let sent = tokio::select! { + result = &mut heartbeat => panic!("premature ping: {result:?}"), + sent = peer.request() => sent, + }; + assert_eq!(sent.method, "ping"); + assert_ne!(sent.id, initialize.id); + assert_ne!(sent.id, discovery.id); + peer.reply(sent.id, Value::Null).await; + bounded(heartbeat).await.unwrap(); + assert!(multi.client(&HostId::from("discovering")).await.is_none()); + bounded(multi.remove_host(&HostId::from("discovering"))) + .await + .unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); + assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); +} + +#[tokio::test] +async fn unanswered_weak_ping_does_not_retain_last_client_or_driver() { + let (transport, mut peer) = pair(); + let client = Client::connect(BoxedTransport::new(transport), no_timeout()) + .await + .unwrap(); + let clone = client.clone(); + let ping = peer.bound().await; + let heartbeat = ping.ping(); + tokio::pin!(heartbeat); + let sent = tokio::select! { + result = &mut heartbeat => panic!("premature ping: {result:?}"), + sent = peer.request() => sent, + }; + assert_eq!(sent.method, "ping"); + drop(client); + assert!(peer.rx.try_recv().is_err()); + drop(clone); + assert!(matches!( + bounded(heartbeat).await, + Err(ClientError::Shutdown) + )); + assert!(bounded(peer.rx.recv()).await.is_none(), "driver leaked"); + assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); +} + +#[tokio::test] +async fn shutdown_and_transport_close_preserve_normal_errors_but_cancel_weak_ping() { + for explicit in [false, true] { + let (transport, mut peer) = pair(); + let client = Client::connect(transport, no_timeout()).await.unwrap(); + let ping = peer.bound().await; + let normal = client.ping(); + let weak = ping.ping(); + tokio::pin!(normal, weak); + let requests = async { + peer.request().await; + peer.request().await; + }; + tokio::select! { + result = &mut normal => panic!("premature ping: {result:?}"), + result = &mut weak => panic!("premature ping: {result:?}"), + _ = requests => {}, + } + if explicit { + client.shutdown().await; + } else { + drop(peer.tx); + } + let Err(ClientError::Rpc(error)) = bounded(normal).await else { + panic!("normal ping must retain its existing RPC teardown error"); + }; + assert_eq!(error.code, -32000); + assert_eq!( + error.message, + if explicit { + "client shut down" + } else { + "transport closed" + } + ); + assert!(matches!(bounded(weak).await, Err(ClientError::Shutdown))); + assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); + assert!(bounded(peer.rx.recv()).await.is_none()); + } +} + +#[tokio::test] +async fn weak_ping_preserves_server_error_and_configured_timeout() { + let (transport, mut peer) = pair(); + let client = Client::connect(transport, no_timeout()).await.unwrap(); + let ping = peer.bound().await; + let heartbeat = ping.ping(); + tokio::pin!(heartbeat); + let request = tokio::select! { + result = &mut heartbeat => panic!("premature ping: {result:?}"), + request = peer.request() => request, + }; + peer.tx + .send(TransportMessage::Parsed(JsonRpcMessage::ErrorResponse( + JsonRpcErrorResponse { + jsonrpc: JsonRpcVersion::V2, + id: request.id, + error: JsonRpcError { + code: -32601, + message: "unsupported ping".into(), + data: None, + }, + }, + ))) + .await + .unwrap(); + let Err(ClientError::Rpc(error)) = bounded(heartbeat).await else { + panic!("expected server RPC error"); + }; + assert_eq!(error.code, -32601); + drop(client); + assert!(bounded(peer.rx.recv()).await.is_none()); + + let (transport, mut peer) = pair(); + let client = Client::connect( + transport, + ClientConfig { + default_request_timeout: Some(Duration::ZERO), + ..ClientConfig::default() + }, + ) + .await + .unwrap(); + let ping = peer.bound().await; + assert!(matches!( + bounded(ping.ping()).await, + Err(ClientError::Cancelled) + )); + assert_eq!(peer.request().await.method, "ping"); + drop(client); + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn managed_host_ids_survive_factory_failure_handshake_failure_and_reconnect() { + let (first, mut first_peer) = pair(); + let (second, mut second_peer) = pair(); + let (third, mut third_peer) = pair(); + let transports = Arc::new(Mutex::new(VecDeque::from([ + Err(TransportError::Closed), + Ok(first), + Ok(second), + Ok(third), + ]))); + let attempts = Arc::new(AtomicUsize::new(0)); + let factory_attempts = attempts.clone(); + let config = HostConfig::new("retained", "Retained delivery host", move |_| { + let transports = transports.clone(); + factory_attempts.fetch_add(1, Ordering::SeqCst); + async move { + transports + .lock() + .await + .pop_front() + .unwrap() + .map(BoxedTransport::new) + } + }) + .with_client_config(no_timeout()) + .with_reconnect_policy(ReconnectPolicy::immediate_forever()); + let multi = MultiHostClient::new(); + let mut events = multi.host_events(); + multi.add_host(config).await.unwrap(); + + let first_ping = first_peer.bound().await; + let abandoned = first_peer.request().await; + assert_eq!(abandoned.method, "initialize"); + assert_eq!(abandoned.id, 1); + drop(first_peer.tx); + assert!(bounded(first_peer.rx.recv()).await.is_none()); + assert!(matches!( + first_ping.ping().await, + Err(ClientError::Shutdown) + )); + + let second_ping = second_peer.bound().await; + let retry = second_peer.request().await; + assert_eq!(retry.method, "initialize"); + assert_eq!(retry.id, 2); + assert_eq!( + retry.params.as_ref().unwrap()["clientId"], + abandoned.params.as_ref().unwrap()["clientId"] + ); + let init = + json!({"protocolVersion": ahp_types::PROTOCOL_VERSION, "serverSeq": 10, "snapshots": []}); + second_peer.reply(abandoned.id, init.clone()).await; + let heartbeat = second_ping.ping(); + tokio::pin!(heartbeat); + let barrier = tokio::select! { + result = &mut heartbeat => panic!("premature result: {result:?}"), + sent = second_peer.request() => sent, + }; + assert_eq!( + barrier.method, "ping", + "stale handshake must not start discovery" + ); + assert_eq!(barrier.id, 3); + second_peer.reply(barrier.id, Value::Null).await; + bounded(heartbeat).await.unwrap(); + assert!(second_peer.rx.try_recv().is_err()); + assert!(multi.client(&HostId::from("retained")).await.is_none()); + + second_peer.reply(retry.id, init).await; + let discovery = second_peer.request().await; + assert_eq!(discovery.method, "listSessions"); + assert_eq!(discovery.id, 4); + second_peer.reply(discovery.id, json!({"items": []})).await; + loop { + if matches!( + bounded(events.recv()).await, + Some(HostEvent::Connected { .. }) + ) { + break; + } + } + let unanswered = second_ping.ping(); + tokio::pin!(unanswered); + let old_ping = tokio::select! { + result = &mut unanswered => panic!("premature result: {result:?}"), + sent = second_peer.request() => sent, + }; + assert_eq!(old_ping.id, 5); + drop(second_peer.tx); + assert!(matches!( + bounded(unanswered).await, + Err(ClientError::Shutdown) + )); + assert!(bounded(second_peer.rx.recv()).await.is_none()); + + let third_ping = third_peer.bound().await; + let reconnect = third_peer.request().await; + assert_eq!(reconnect.method, "reconnect"); + assert_eq!(reconnect.id, 6); + assert_eq!( + reconnect.params.as_ref().unwrap()["clientId"], + retry.params.as_ref().unwrap()["clientId"] + ); + third_peer.reply(old_ping.id, Value::Null).await; + let heartbeat = third_ping.ping(); + tokio::pin!(heartbeat); + let barrier = tokio::select! { + result = &mut heartbeat => panic!("premature result: {result:?}"), + sent = third_peer.request() => sent, + }; + assert_eq!( + barrier.method, "ping", + "stale ping must not satisfy reconnect" + ); + assert_eq!(barrier.id, 7); + third_peer.reply(barrier.id, Value::Null).await; + bounded(heartbeat).await.unwrap(); + assert!(third_peer.rx.try_recv().is_err()); + third_peer + .reply( + reconnect.id, + json!({"type": "replay", "actions": [], "missing": []}), + ) + .await; + let discovery = third_peer.request().await; + assert_eq!(discovery.method, "listSessions"); + assert_eq!(discovery.id, 8); + third_peer.reply(discovery.id, json!({"items": []})).await; + loop { + if matches!( + bounded(events.recv()).await, + Some(HostEvent::Connected { .. }) + ) { + break; + } + } + assert_eq!(attempts.load(Ordering::SeqCst), 4); + for count in [ + &first_peer.bind_count, + &second_peer.bind_count, + &third_peer.bind_count, + ] { + assert_eq!(count.load(Ordering::SeqCst), 1); + } + bounded(multi.remove_host(&HostId::from("retained"))) + .await + .unwrap(); + assert!(bounded(third_peer.rx.recv()).await.is_none()); + assert!(matches!( + third_ping.ping().await, + Err(ClientError::Shutdown) + )); +} + +#[tokio::test] +async fn independent_managed_hosts_start_independent_id_sequences() { + let multi = MultiHostClient::new(); + for id in ["first", "second"] { + let (transport, mut peer) = pair(); + let transport = Arc::new(Mutex::new(Some(transport))); + let config = HostConfig::new(id, id, move |_| { + let transport = transport.clone(); + async move { Ok(BoxedTransport::new(transport.lock().await.take().unwrap())) } + }) + .with_client_config(no_timeout()) + .with_reconnect_policy(ReconnectPolicy::disabled()); + multi.add_host(config).await.unwrap(); + assert_eq!(peer.request().await.id, 1); + bounded(multi.remove_host(&HostId::from(id))).await.unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); + } +} diff --git a/docs/.changes/20261001-rust-weak-ping-binding.json b/docs/.changes/20261001-rust-weak-ping-binding.json new file mode 100644 index 000000000..156373912 --- /dev/null +++ b/docs/.changes/20261001-rust-weak-ping-binding.json @@ -0,0 +1,5 @@ +{ + "type": "added", + "message": "Rust transport client binding supplies a non-owning `WeakPingHandle` for client-correlated keepalive during discovery; managed hosts retain request IDs across reconnect attempts, with deterministic weak-ping shutdown and cancelled-request cleanup.", + "targets": ["rust"] +} From 30a210fee4e8389f233f28c370a9d1be2cd46102 Mon Sep 17 00:00:00 2001 From: Hunter Sadler Date: Fri, 2 Oct 2026 14:31:32 -0600 Subject: [PATCH 2/2] fix(rust): separate host readiness from discovery and own keepalive Replace the transport ping binding with configurable driver-owned inbound-idle keepalive. Publish managed clients after handshake and reconcile concurrent session discovery without starving events or commands. Preserve logical-host request ID continuity and bound driver teardown. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- clients/rust/MULTI_HOST.md | 15 +- clients/rust/crates/ahp/Cargo.toml | 2 +- clients/rust/crates/ahp/README.md | 73 +- clients/rust/crates/ahp/src/client.rs | 312 +++++--- clients/rust/crates/ahp/src/client_tests.rs | 326 +++++---- clients/rust/crates/ahp/src/hosts/runtime.rs | 146 +++- clients/rust/crates/ahp/src/hosts/types.rs | 5 +- clients/rust/crates/ahp/src/lib.rs | 11 +- clients/rust/crates/ahp/src/transport.rs | 21 - .../crates/ahp/tests/connection_lifecycle.rs | 689 ++++++++++++++++++ clients/rust/crates/ahp/tests/weak_ping.rs | 589 --------------- .../20261001-rust-connection-lifecycle.json | 5 + .../20261001-rust-weak-ping-binding.json | 5 - 13 files changed, 1285 insertions(+), 914 deletions(-) create mode 100644 clients/rust/crates/ahp/tests/connection_lifecycle.rs delete mode 100644 clients/rust/crates/ahp/tests/weak_ping.rs create mode 100644 docs/.changes/20261001-rust-connection-lifecycle.json delete mode 100644 docs/.changes/20261001-rust-weak-ping-binding.json diff --git a/clients/rust/MULTI_HOST.md b/clients/rust/MULTI_HOST.md index bbe136179..b4e9a83a2 100644 --- a/clients/rust/MULTI_HOST.md +++ b/clients/rust/MULTI_HOST.md @@ -41,6 +41,19 @@ Snapshots are immutable. To observe changes, listen to the connection-event stre Each host runs in its own internal task, a `HostRuntime`, that owns the current `Client`, retries the configured `ReconnectPolicy`, and re-subscribes to known URIs across reconnects. +Connection readiness depends only on a successful `initialize` / `reconnect`. +The event receiver is installed before that handshake, and the client is +published before the session-cache refresh starts. `listSessions` is an ordinary +concurrent RPC: slow or failed discovery does not block other client requests. +Until it completes, `session_summaries` may be empty or retain the previous +connection's cache. Notifications received during the refresh are merged over +its result; cancelled or superseded refreshes cannot update a newer connection. + +The underlying client's automatic keepalive checks inbound wire silence, +independently of discovery. Configure it with `HostConfig::with_client_config` +and `ClientConfig::keepalive`; a liveness timeout closes the connection and +enters the normal reconnect policy. Set `keepalive: None` to disable it. + Every successful reconnect bumps a per-host **generation** counter. Any `HostClientHandle` you obtained from a previous connection refuses to dispatch on the new one and returns `HostError::HostReconnected`; request a fresh handle in that case. This prevents subtle bugs where a handle held across a reconnect silently writes to a different connection. ## Stable `clientId` per host @@ -122,7 +135,7 @@ handle.check_alive().await?; # Ok(()) } ``` -Configuration knobs live on `HostConfig` (`with_client_id`, `with_initial_subscriptions`, `with_client_config`, `with_reconnect_policy`) and on `ReconnectPolicy::{disabled, immediate_forever, exponential}`. For persistent identity across launches, plug in a persistent `ClientIdStore` via `MultiHostClient::with_client_id_store(...)` (see below) or load the `clientId` yourself and pass it through `HostConfig::with_client_id`. +Configuration knobs live on `HostConfig` (`with_client_id`, `with_initial_subscriptions`, `with_client_config`, `with_reconnect_policy`), `ClientConfig::keepalive`, and `ReconnectPolicy::{disabled, immediate_forever, exponential}`. For persistent identity across launches, plug in a persistent `ClientIdStore` via `MultiHostClient::with_client_id_store(...)` (see below) or load the `clientId` yourself and pass it through `HostConfig::with_client_id`. ## Persistent `clientId`s — `ClientIdStore` diff --git a/clients/rust/crates/ahp/Cargo.toml b/clients/rust/crates/ahp/Cargo.toml index c4f310848..5f87d70e0 100644 --- a/clients/rust/crates/ahp/Cargo.toml +++ b/clients/rust/crates/ahp/Cargo.toml @@ -29,7 +29,7 @@ tracing = { workspace = true } jiff = { workspace = true } [dev-dependencies] -tokio = { workspace = true, features = ["full"] } +tokio = { workspace = true, features = ["full", "test-util"] } serde_json = { workspace = true } ahp-ws = { path = "../ahp-ws" } tracing-subscriber = "0.3" diff --git a/clients/rust/crates/ahp/README.md b/clients/rust/crates/ahp/README.md index 0b998db20..41c54c84f 100644 --- a/clients/rust/crates/ahp/README.md +++ b/clients/rust/crates/ahp/README.md @@ -65,28 +65,57 @@ impl Transport for MyTransport { See `tests/client_roundtrip.rs` for a complete in-memory example. -### Transport-owned keepalive - -Override the optional synchronous -`Transport::bind_client(&mut self, ping: ahp::WeakPingHandle)` callback to pass -the handle to a transport-owned keepalive task. `Client::connect` binds it once, -before transport I/O starts or any `initialize` / `reconnect` request can be -sent. `BoxedTransport` forwards the callback, including for managed hosts. -Subsequent handshakes on the same client do not bind it again. - -`WeakPingHandle::ping().await` uses the same root-channel request, ID allocator, -response map, and configured timeout as `Client::ping`. It can therefore keep an -initialized connection alive while session discovery is still pending. The -transport owns the keepalive timing and any negotiation or activity gating; -binding does not start keepalive. - -Neither a retained handle nor an in-flight weak ping owns the client driver. -Explicit shutdown, transport closure, or dropping the last `Client` resolves -weak pings with `ClientError::Shutdown`. Request timeouts remain -`ClientError::Cancelled`, and server errors remain `ClientError::Rpc`. Cancelling -a request future removes its pending response entry without retracting an -already-sent request. Existing transports need no changes: binding defaults to -a no-op on both `Transport` and the object-safe `DynTransport` adapter. +### Automatic idle keepalive + +The client driver sends a root-channel `ping` after 30 seconds without received +wire traffic and closes the connection after 90 seconds of continuous inbound +silence. Any inbound frame proves liveness, not just a ping response. Outbound +requests and notifications do not reset these deadlines. There is at most one +automatic ping pending at a time, using the same ID allocator and response map +as ordinary requests. Automatic keepalive is independent of the request timeout +and requires no transport extension or separate heartbeat task. + +Configure these deadlines, or disable keepalive, through `ClientConfig`: + +```rust +use ahp::{ClientConfig, KeepaliveConfig}; +use std::time::Duration; + +let config = ClientConfig { + keepalive: Some(KeepaliveConfig { + idle_interval: Duration::from_secs(15), + liveness_timeout: Duration::from_secs(45), + }), + ..ClientConfig::default() +}; +let disabled = ClientConfig { keepalive: None, ..ClientConfig::default() }; +``` + +The idle interval must be nonzero, and the liveness timeout must be greater +than it. Invalid policy is rejected by `Client::connect` before I/O. Consumers +constructing `ClientConfig` with all fields explicitly must add `keepalive`; +struct updates using `..ClientConfig::default()` remain source-compatible. +The wire protocol is unchanged. These are local SDK settings, not automatic +interpretation of host metadata or negotiated deadlines. + +Because `Transport::send` and `recv` borrow the same transport, receive +observation pauses during a send. With keepalive enabled, each send is bounded +separately by `liveness_timeout` measured from the start of that write; a stalled +write is reported as a transport write timeout, not inbound silence. Transport +cleanup is bounded to five seconds, even if `close` cannot complete. + +Keepalive ends with shutdown, transport closure, or dropping the last client. +Liveness failure closes event streams so managed hosts follow their normal +reconnect policy. In-flight normal requests retain their `-32000` +`ClientError::Rpc` teardown errors, and request timeouts remain +`ClientError::Cancelled`. Cancelling a request future removes its pending entry +without retracting an already-sent request. + +Managed hosts become connected and expose their generation-checked client as +soon as `initialize` / `reconnect` completes, before `listSessions` resolves. +The concurrent session-cache refresh preserves intervening notifications and +cannot apply results after connection replacement. Refresh failures are logged +and do not change readiness. Managed hosts retain their request ID allocator across connection attempts for the lifetime of one host supervisor, including failed handshakes. Late replies diff --git a/clients/rust/crates/ahp/src/client.rs b/clients/rust/crates/ahp/src/client.rs index a0d108862..ace1eaf1f 100644 --- a/clients/rust/crates/ahp/src/client.rs +++ b/clients/rust/crates/ahp/src/client.rs @@ -55,6 +55,7 @@ use serde::{de::DeserializeOwned, Serialize}; use serde_json::Value; use tokio::sync::{broadcast, mpsc, oneshot, watch, Mutex}; use tokio::task::JoinHandle; +use tokio::time::Instant; use crate::error::{ClientError, TransportError}; use crate::transport::{Transport, TransportMessage}; @@ -71,6 +72,32 @@ pub struct ClientConfig { pub default_request_timeout: Option, /// Size of each subscription's broadcast ring buffer. pub subscription_buffer: usize, + /// Automatic connection liveness checks. `None` disables keepalive. + /// + /// Defaults to a ping after 30 seconds without inbound wire traffic and + /// connection failure after 90 seconds of continuous inbound silence. + pub keepalive: Option, +} + +/// Idle-based keepalive policy owned by the client's transport driver. +#[derive(Debug, Clone, Copy)] +pub struct KeepaliveConfig { + /// Inbound silence before sending one root-channel `ping`. + pub idle_interval: Duration, + /// Total inbound silence before declaring the connection failed. + /// + /// Must exceed [`Self::idle_interval`]. Any received wire message resets + /// both deadlines, whether or not it answers the outstanding ping. + pub liveness_timeout: Duration, +} + +impl Default for KeepaliveConfig { + fn default() -> Self { + Self { + idle_interval: Duration::from_secs(30), + liveness_timeout: Duration::from_secs(90), + } + } } impl Default for ClientConfig { @@ -78,6 +105,7 @@ impl Default for ClientConfig { Self { default_request_timeout: Some(Duration::from_secs(30)), subscription_buffer: DEFAULT_SUBSCRIPTION_BUFFER, + keepalive: Some(KeepaliveConfig::default()), } } } @@ -131,6 +159,19 @@ pub struct ClientEventStream { } impl ClientEventStream { + pub(crate) fn drain_buffered(&mut self) -> Vec { + let mut events = Vec::new(); + // Snapshot the backlog so a continuous producer cannot trap a refresh. + for _ in 0..self.rx.len() { + match self.rx.try_recv() { + Ok(event) => events.push(event), + Err(broadcast::error::TryRecvError::Lagged(_)) => continue, + Err(_) => break, + } + } + events + } + /// Await the next event. Returns `None` when the client has shut /// down (the underlying broadcast channel has closed). pub async fn recv(&mut self) -> Option { @@ -188,6 +229,11 @@ pub struct DispatchHandle { // ─── Internal plumbing ─────────────────────────────────────────────────────── type PendingMap = HashMap>>; +type PreparedRequest<'a> = ( + JsonRpcMessage, + PendingRequest<'a>, + oneshot::Receiver>, +); pub(crate) struct RequestIds { next: std::sync::Mutex>, @@ -233,7 +279,6 @@ struct Shared { enum Outbound { Message(JsonRpcMessage), - Shutdown, } impl Shared { @@ -251,75 +296,67 @@ impl Shared { } } - async fn request(&self, method: &str, params: P, weak: bool) -> Result - where - P: Serialize, - R: DeserializeOwned, - { - let mut closed = self.closed.subscribe(); - if *closed.borrow() { + fn prepare_request( + &self, + method: &str, + params: P, + ) -> Result, ClientError> { + if *self.closed.borrow() { return Err(ClientError::Shutdown); } let id = self.request_ids.allocate()?; let params_val = serde_json::to_value(¶ms)?; - let params_any = if params_val.is_null() { - None - } else { - Some(ahp_types::common::AnyValue::from(params_val)) - }; let req = JsonRpcMessage::Request(JsonRpcRequest { jsonrpc: JsonRpcVersion::V2, id, method: method.into(), - params: params_any, + params: if params_val.is_null() { + None + } else { + Some(params_val) + }, }); - let (tx, rx) = oneshot::channel(); { let mut pending = self.pending.lock().expect("pending mutex poisoned"); - if *closed.borrow() { + if *self.closed.borrow() { return Err(ClientError::Shutdown); } pending.insert(id, tx); } - let _pending = PendingRequest { shared: self, id }; + Ok((req, PendingRequest { shared: self, id }, rx)) + } - let response = async { - self.outbound - .send(Outbound::Message(req)) - .await - .map_err(|_| ClientError::Shutdown)?; + async fn request(&self, method: &str, params: P) -> Result + where + P: Serialize, + R: DeserializeOwned, + { + let (req, _pending, rx) = self.prepare_request(method, params)?; + self.outbound + .send(Outbound::Message(req)) + .await + .map_err(|_| ClientError::Shutdown)?; - let result = match self.config.default_request_timeout { - Some(dur) => tokio::time::timeout(dur, rx) - .await - .map_err(|_| ClientError::Cancelled)?, - None => rx.await, - }; - match result { - Ok(Ok(value)) => Ok(serde_json::from_value(value)?), - Ok(Err(e)) => Err(ClientError::Rpc(e)), - Err(_) => Err(ClientError::Shutdown), - } + let result = match self.config.default_request_timeout { + Some(dur) => tokio::time::timeout(dur, rx) + .await + .map_err(|_| ClientError::Cancelled)?, + None => rx.await, }; - tokio::select! { - biased; - _ = closed.wait_for(|closed| *closed), if weak => Err(ClientError::Shutdown), - result = response => result, + match result { + Ok(Ok(value)) => Ok(serde_json::from_value(value)?), + Ok(Err(e)) => Err(ClientError::Rpc(e)), + Err(_) => Err(ClientError::Shutdown), } } - async fn ping(&self, weak: bool) -> Result<(), ClientError> { - #[derive(Serialize)] - struct PingParams { - channel: &'static str, - } + async fn ping(&self) -> Result<(), ClientError> { self.request( "ping", PingParams { channel: ROOT_RESOURCE_URI, }, - weak, ) .await } @@ -340,32 +377,9 @@ impl Drop for PendingRequest<'_> { } } -/// Non-owning handle for client-correlated keepalive requests. -/// -/// Supplied to [`Transport::bind_client`] before transport I/O starts. -/// Cloning the handle or awaiting [`Self::ping`] never keeps the client's -/// background driver alive. Pings share the client's normal request IDs, -/// response correlation, and configured request timeout. -#[derive(Clone)] -pub struct WeakPingHandle { - shared: Weak, -} - -impl WeakPingHandle { - /// Send the same root-channel request as [`Client::ping`]. - /// - /// Returns [`ClientError::Shutdown`] after explicit shutdown, transport - /// closure, or dropping the last [`Client`], including for an unanswered - /// in-flight ping. Timeouts return [`ClientError::Cancelled`]; server errors - /// return [`ClientError::Rpc`]. Dropping this future removes its pending - /// response entry but does not retract a request already sent. - /// - /// The caller decides when keepalive is appropriate for the connection; - /// binding the handle does not start a ping or negotiate keepalive. - pub async fn ping(&self) -> Result<(), ClientError> { - let shared = self.shared.upgrade().ok_or(ClientError::Shutdown)?; - shared.ping(true).await - } +#[derive(Serialize)] +struct PingParams { + channel: &'static str, } // ─── Server-initiated request handling ─────────────────────────────────────── @@ -535,10 +549,22 @@ impl Client { } pub(crate) async fn connect_with_request_ids( - mut transport: T, + transport: T, config: ClientConfig, request_ids: Arc, ) -> Result { + if let Some(keepalive) = config.keepalive { + if keepalive.idle_interval.is_zero() + || keepalive.liveness_timeout <= keepalive.idle_interval + || Instant::now() + .checked_add(keepalive.liveness_timeout) + .is_none() + { + return Err(TransportError::Protocol( + "keepalive requires a nonzero idle interval and a representable liveness timeout greater than the idle interval".into() + ).into()); + } + } let (outbound_tx, outbound_rx) = mpsc::channel::(64); let (all_events_tx, _) = broadcast::channel::(config.subscription_buffer); let (closed, _) = watch::channel(false); @@ -554,9 +580,6 @@ impl Client { server_request_handler: std::sync::Mutex::new(None), }); - transport.bind_client(WeakPingHandle { - shared: Arc::downgrade(&shared), - }); let handle = tokio::spawn(drive_transport(transport, shared.clone(), outbound_rx)); let reader = Arc::new(DriveHandle { handle: Mutex::new(Some(handle)), @@ -570,11 +593,10 @@ impl Client { /// Gracefully shut down the client. /// - /// In-flight normal requests retain the `-32000` [`ClientError::Rpc`] - /// shutdown error; weak pings resolve with [`ClientError::Shutdown`]. + /// In-flight requests retain the `-32000` [`ClientError::Rpc`] + /// shutdown error. Automatic keepalive ends with the transport driver. pub async fn shutdown(&self) { self.shared.stop_requests(Some("client shut down")); - let _ = self.shared.outbound.send(Outbound::Shutdown).await; } /// Send a JSON-RPC request and await its result. @@ -583,7 +605,7 @@ impl Client { P: Serialize, R: DeserializeOwned, { - self.shared.request(method, params, false).await + self.shared.request(method, params).await } /// Send a JSON-RPC notification (fire-and-forget). @@ -660,7 +682,7 @@ impl Client { /// server responds regardless of whether `initialize` has completed or any /// subscriptions are held. pub async fn ping(&self) -> Result<(), ClientError> { - self.shared.ping(false).await + self.shared.ping().await } /// Subscribe to a URI and obtain a handle that streams @@ -954,27 +976,44 @@ async fn drive_transport( mut outbound: mpsc::Receiver, ) { let _requests = DriverRequests(shared.clone()); + let mut closed = shared.closed.subscribe(); + let mut last_received = Instant::now(); + let mut heartbeat = None; loop { - tokio::select! { - outbound_msg = outbound.recv() => { - match outbound_msg { - Some(Outbound::Message(msg)) => { - if let Ok(wire) = TransportMessage::encode(&msg) { - if let Err(err) = transport.send(wire).await { - tracing::warn!(?err, "transport send failed"); - break; - } - } - } - Some(Outbound::Shutdown) | None => { - let _ = transport.close().await; - break; - } + if *closed.borrow() { + break; + } + let deadline = shared.config.keepalive.map(|policy| { + last_received + + if heartbeat.is_some() { + policy.liveness_timeout + } else { + policy.idle_interval } + }); + let mut event = tokio::select! { + _ = async { let _ = closed.wait_for(|closed| *closed).await; } => break, + inbound = transport.recv() => DriverEvent::Inbound(inbound), + _ = keepalive_deadline(deadline) => DriverEvent::Idle, + outbound = outbound.recv() => DriverEvent::Outbound(outbound), + }; + // Preserve receive-first expiry semantics without starving outbound RPCs. + if matches!(event, DriverEvent::Idle) { + if let Some(inbound) = tokio::select! { + biased; + inbound = transport.recv() => Some(inbound), + _ = async {} => None, + } { + event = DriverEvent::Inbound(inbound); } - inbound = transport.recv() => { + } + match event { + DriverEvent::Inbound(inbound) => { match inbound { Ok(Some(wire)) => { + last_received = Instant::now(); + // Any inbound traffic proves liveness, not just a ping response. + heartbeat.take(); match wire.into_parsed() { Ok(msg) => dispatch_inbound(&shared, msg).await, Err(err) => tracing::warn!(?err, "malformed frame"), @@ -987,13 +1026,62 @@ async fn drive_transport( } } } + DriverEvent::Idle => { + if heartbeat.is_some() { + tracing::warn!("connection liveness timeout"); + break; + } + let (request, pending, response) = match shared.prepare_request( + "ping", + PingParams { + channel: ROOT_RESOURCE_URI, + }, + ) { + Ok(request) => request, + Err(err) => { + tracing::warn!(?err, "keepalive request failed"); + break; + } + }; + heartbeat = Some((pending, response)); + if let Err(err) = send_wire( + &mut transport, + request, + &mut closed, + shared.config.keepalive.map(|p| p.liveness_timeout), + ) + .await + { + tracing::warn!(?err, "keepalive transport send failed"); + break; + } + } + DriverEvent::Outbound(outbound_msg) => match outbound_msg { + Some(Outbound::Message(msg)) => { + if let Err(err) = send_wire( + &mut transport, + msg, + &mut closed, + shared.config.keepalive.map(|p| p.liveness_timeout), + ) + .await + { + tracing::warn!(?err, "transport send failed"); + break; + } + } + None => break, + }, } } + drop(heartbeat); + drop(outbound); // Teardown: close everything and fail outstanding requests. shared.stop_requests(Some("transport closed")); let mut subs = shared.subscriptions.lock().await; subs.clear(); + drop(subs); // Drop the top-level fan-out sender so any active // `ClientEventStream::recv()` resolves with `None` rather than // hanging forever (the `Sender` would otherwise stay alive inside @@ -1001,6 +1089,42 @@ async fn drive_transport( if let Ok(mut guard) = shared.all_events.lock() { guard.take(); } + match tokio::time::timeout(Duration::from_secs(5), transport.close()).await { + Ok(Ok(())) => {} + Ok(Err(err)) => tracing::warn!(?err, "transport close failed"), + Err(_) => tracing::warn!("transport close timed out"), + } +} + +enum DriverEvent { + Inbound(Result, TransportError>), + Outbound(Option), + Idle, +} + +async fn keepalive_deadline(deadline: Option) { + match deadline { + Some(deadline) => tokio::time::sleep_until(deadline).await, + None => std::future::pending().await, + } +} + +async fn send_wire( + transport: &mut T, + msg: JsonRpcMessage, + closed: &mut watch::Receiver, + write_timeout: Option, +) -> Result<(), TransportError> { + let wire = TransportMessage::encode(&msg)?; + // Transport's &mut methods cannot read during a send. Bound a stalled write + // separately, rather than misreporting queued inbound traffic as silence. + let deadline = write_timeout.map(|timeout| Instant::now() + timeout); + tokio::select! { + biased; + _ = closed.wait_for(|closed| *closed) => Err(TransportError::Closed), + _ = keepalive_deadline(deadline) => Err(TransportError::Protocol("transport write timed out".into())), + result = transport.send(wire) => result, + } } struct DriverRequests(Arc); diff --git a/clients/rust/crates/ahp/src/client_tests.rs b/clients/rust/crates/ahp/src/client_tests.rs index 32d73fa7a..c9956ffd7 100644 --- a/clients/rust/crates/ahp/src/client_tests.rs +++ b/clients/rust/crates/ahp/src/client_tests.rs @@ -6,14 +6,9 @@ use crate::BoxedTransport; struct TestTransport { sent: mpsc::Sender, received: mpsc::Receiver, - bound: Option, } impl Transport for TestTransport { - fn bind_client(&mut self, ping: WeakPingHandle) { - self.bound = Some(ping); - } - async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { self.sent .send(message) @@ -27,7 +22,7 @@ impl Transport for TestTransport { } async fn client( - config: ClientConfig, + timeout: Option, ) -> ( Client, mpsc::Receiver, @@ -35,139 +30,117 @@ async fn client( ) { let (sent, rx) = mpsc::channel(1); let (tx, received) = mpsc::channel(1); - let transport = TestTransport { - sent, - received, - bound: None, - }; - ( - Client::connect(BoxedTransport::new(transport), config) - .await - .unwrap(), - rx, - tx, + let client = Client::connect( + BoxedTransport::new(TestTransport { sent, received }), + ClientConfig { + default_request_timeout: timeout, + keepalive: None, + ..ClientConfig::default() + }, ) + .await + .unwrap(); + (client, rx, tx) } #[tokio::test] async fn cancelled_and_timed_out_requests_remove_pending_entries() { - let (client, mut sent, _received) = client(ClientConfig { - default_request_timeout: None, - ..ClientConfig::default() - }) - .await; - let weak = WeakPingHandle { - shared: Arc::downgrade(&client.shared), - }; - for is_weak in [false, true] { - let mut request = Box::pin(async { - if is_weak { - weak.ping().await - } else { - client.ping().await - } - }); - tokio::select! { - result = &mut request => panic!("premature result: {result:?}"), - _ = sent.recv() => {}, - } - assert_eq!(client.shared.pending.lock().unwrap().len(), 1); - drop(request); - assert!(client.shared.pending.lock().unwrap().is_empty()); + let (client, mut sent, _received) = client(None).await; + let mut request = Box::pin(client.ping()); + tokio::select! { + result = &mut request => panic!("premature result: {result:?}"), + _ = sent.recv() => {}, } + assert_eq!(client.shared.pending.lock().unwrap().len(), 1); + drop(request); + assert!(client.shared.pending.lock().unwrap().is_empty()); drop(client); - assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) - .await - .unwrap() - .is_none()); - assert!( - weak.shared.upgrade().is_none(), - "Shared leaked after driver drop" - ); + assert!(sent.recv().await.is_none()); - let (client, mut sent, _received) = self::client(ClientConfig { - default_request_timeout: Some(Duration::ZERO), - ..ClientConfig::default() - }) - .await; - let weak = WeakPingHandle { - shared: Arc::downgrade(&client.shared), - }; - assert!(matches!(weak.ping().await, Err(ClientError::Cancelled))); - assert!(client.shared.pending.lock().unwrap().is_empty()); - assert!(sent.recv().await.is_some()); + let (client, mut sent, _received) = self::client(Some(Duration::ZERO)).await; assert!(matches!(client.ping().await, Err(ClientError::Cancelled))); assert!(client.shared.pending.lock().unwrap().is_empty()); assert!(sent.recv().await.is_some()); drop(client); - assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) - .await - .unwrap() - .is_none()); -} - -#[tokio::test] -async fn owner_drop_clears_unanswered_weak_request_and_releases_shared_allocator() { - let (client, mut sent, _received) = client(ClientConfig { - default_request_timeout: None, - ..ClientConfig::default() - }) - .await; - let weak = WeakPingHandle { - shared: Arc::downgrade(&client.shared), - }; - let ids = Arc::downgrade(&client.shared.request_ids); - let mut request = Box::pin(weak.ping()); - tokio::select! { - result = &mut request => panic!("premature result: {result:?}"), - _ = sent.recv() => {}, - } - drop(client); - assert!(matches!(request.await, Err(ClientError::Shutdown))); - assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) - .await - .unwrap() - .is_none()); - assert!(weak.shared.upgrade().is_none()); - assert!(ids.upgrade().is_none()); + assert!(sent.recv().await.is_none()); } #[tokio::test] async fn request_id_exhaustion_never_wraps_or_enqueues_another_request() { - let (client, mut sent, _received) = client(ClientConfig { - default_request_timeout: Some(Duration::ZERO), - ..ClientConfig::default() - }) - .await; + let (client, mut sent, _received) = client(Some(Duration::ZERO)).await; *client.shared.request_ids.next.lock().unwrap() = Some(u64::MAX); - let weak = WeakPingHandle { - shared: Arc::downgrade(&client.shared), - }; - assert!(matches!(weak.ping().await, Err(ClientError::Cancelled))); + assert!(matches!(client.ping().await, Err(ClientError::Cancelled))); let JsonRpcMessage::Request(request) = sent.recv().await.unwrap().into_parsed().unwrap() else { panic!("expected request"); }; assert_eq!(request.id, u64::MAX); - for is_weak in [false, true] { - let result = if is_weak { - weak.ping().await - } else { - client.ping().await - }; + for _ in 0..2 { assert!( - matches!(result, Err(ClientError::Transport(TransportError::Protocol(message))) if message == "request ID space exhausted") + matches!(client.ping().await, Err(ClientError::Transport(TransportError::Protocol(message))) if message == "request ID space exhausted") ); assert!(client.shared.pending.lock().unwrap().is_empty()); assert!(sent.try_recv().is_err()); } - assert!(client.shared.request_ids.next.lock().unwrap().is_none()); client.shutdown().await; - assert!(matches!(weak.ping().await, Err(ClientError::Shutdown))); + assert!(matches!(client.ping().await, Err(ClientError::Shutdown))); drop(client); - assert!(tokio::time::timeout(Duration::from_secs(2), sent.recv()) + assert!(sent.recv().await.is_none()); +} + +#[tokio::test(start_paused = true)] +async fn automatic_ping_and_driver_do_not_outlive_client_or_retain_pending_entries() { + let (sent, mut rx) = mpsc::channel(1); + let (_tx, received) = mpsc::channel(1); + let client = Client::connect(TestTransport { sent, received }, ClientConfig::default()) .await - .unwrap() - .is_none()); + .unwrap(); + let shared = Arc::downgrade(&client.shared); + let ids = Arc::downgrade(&client.shared.request_ids); + let request = rx.recv().await.unwrap().into_parsed().unwrap(); + assert!(matches!(request, JsonRpcMessage::Request(request) if request.method == "ping")); + assert_eq!(client.shared.pending.lock().unwrap().len(), 1); + let clone = client.clone(); + drop(client); + assert!(shared.upgrade().is_some()); + drop(clone); + assert!(rx.recv().await.is_none()); + assert!(shared.upgrade().is_none()); + assert!(ids.upgrade().is_none()); +} + +#[tokio::test(start_paused = true)] +async fn inbound_activity_retires_automatic_pending_ping() { + let (sent, mut rx) = mpsc::channel(1); + let (tx, received) = mpsc::channel(1); + let client = Client::connect(TestTransport { sent, received }, ClientConfig::default()) + .await + .unwrap(); + assert!(rx.recv().await.is_some()); + assert_eq!(client.shared.pending.lock().unwrap().len(), 1); + tx.send(TransportMessage::Text( + r#"{"jsonrpc":"2.0","method":"activity"}"#.into(), + )) + .await + .unwrap(); + for _ in 0..10 { + tokio::task::yield_now().await; + } + assert!(client.shared.pending.lock().unwrap().is_empty()); + client.shutdown().await; + assert!(rx.recv().await.is_none()); +} + +#[tokio::test(start_paused = true)] +async fn automatic_ping_exhaustion_fails_connection_without_id_reuse() { + let (sent, mut rx) = mpsc::channel(1); + let (_tx, received) = mpsc::channel(1); + let client = Client::connect(TestTransport { sent, received }, ClientConfig::default()) + .await + .unwrap(); + *client.shared.request_ids.next.lock().unwrap() = None; + assert!(rx.recv().await.is_none()); + assert!(client.shared.pending.lock().unwrap().is_empty()); + assert!(matches!(client.ping().await, Err(ClientError::Shutdown))); } struct BlockedTransport { @@ -192,54 +165,107 @@ impl Drop for BlockedTransport { } } -#[tokio::test] -async fn owner_drop_cancels_weak_ping_blocked_on_full_outbound_queue() { +#[tokio::test(start_paused = true)] +async fn stalled_send_is_interrupted_by_shutdown_or_liveness_timeout() { + for shutdown in [false, true] { + let (entered, send_started) = oneshot::channel(); + let (dropped, transport_dropped) = oneshot::channel(); + let client = Client::connect( + BlockedTransport { + entered: Some(entered), + dropped: Some(dropped), + }, + ClientConfig::default(), + ) + .await + .unwrap(); + let mut events = client.events(); + client.notify("block", ()).await.unwrap(); + send_started.await.unwrap(); + if shutdown { + // Full outbound queue must not trap graceful shutdown either. + for _ in 0..64 { + client.notify("queued", ()).await.unwrap(); + } + + client.shutdown().await; + } else { + tokio::time::advance(Duration::from_secs(90)).await; + } + assert!(events.recv().await.is_none()); + transport_dropped.await.unwrap(); + assert!(client.shared.pending.lock().unwrap().is_empty()); + } +} + +struct StalledCloseTransport { + dropped: Option>, +} + +impl Transport for StalledCloseTransport { + async fn send(&mut self, _: TransportMessage) -> Result<(), TransportError> { + Ok(()) + } + + async fn recv(&mut self) -> Result, TransportError> { + std::future::pending().await + } + + async fn close(&mut self) -> Result<(), TransportError> { + std::future::pending().await + } +} + +impl Drop for StalledCloseTransport { + fn drop(&mut self) { + self.dropped.take().unwrap().send(()).unwrap(); + } +} + +#[tokio::test(start_paused = true)] +async fn stalled_close_is_bounded_without_trapping_shutdown_or_receivers() { + let (dropped, mut transport_dropped) = oneshot::channel(); + let client = Client::connect( + StalledCloseTransport { + dropped: Some(dropped), + }, + ClientConfig::default(), + ) + .await + .unwrap(); + let mut events = client.events(); + client.shutdown().await; + assert!(events.recv().await.is_none()); + assert!(transport_dropped.try_recv().is_err()); + tokio::time::advance(Duration::from_secs(5)).await; + transport_dropped.await.unwrap(); + assert!(client.shared.pending.lock().unwrap().is_empty()); +} + +#[tokio::test(start_paused = true)] +async fn blocked_write_watchdog_starts_with_send_not_the_last_inbound_deadline() { let (entered, send_started) = oneshot::channel(); - let (dropped, transport_dropped) = oneshot::channel(); + let (dropped, mut transport_dropped) = oneshot::channel(); let client = Client::connect( - BoxedTransport::new(BlockedTransport { + BlockedTransport { entered: Some(entered), dropped: Some(dropped), - }), - ClientConfig { - default_request_timeout: None, - ..ClientConfig::default() }, + ClientConfig::default(), ) .await .unwrap(); - let weak = WeakPingHandle { - shared: Arc::downgrade(&client.shared), - }; + tokio::time::advance(Duration::from_secs(25)).await; client.notify("block", ()).await.unwrap(); - tokio::time::timeout(Duration::from_secs(2), send_started) - .await - .unwrap() - .unwrap(); - for _ in 0..64 { - client.notify("queued", ()).await.unwrap(); + send_started.await.unwrap(); + tokio::time::advance(Duration::from_secs(65)).await; + for _ in 0..10 { + tokio::task::yield_now().await; } - assert_eq!(client.shared.outbound.capacity(), 0); - let mut request = Box::pin(weak.ping()); - std::future::poll_fn(|cx| { - assert!(request.as_mut().poll(cx).is_pending()); - std::task::Poll::Ready(()) - }) - .await; - assert_eq!(client.shared.pending.lock().unwrap().len(), 1); - drop(client); - assert!(weak - .shared - .upgrade() - .unwrap() - .pending - .lock() - .unwrap() - .is_empty()); - assert!(matches!(request.await, Err(ClientError::Shutdown))); - tokio::time::timeout(Duration::from_secs(2), transport_dropped) - .await - .unwrap() - .unwrap(); - assert!(weak.shared.upgrade().is_none()); + assert!( + transport_dropped.try_recv().is_err(), + "write must not reuse old inbound-silence deadline" + ); + tokio::time::advance(Duration::from_secs(25)).await; + transport_dropped.await.unwrap(); } diff --git a/clients/rust/crates/ahp/src/hosts/runtime.rs b/clients/rust/crates/ahp/src/hosts/runtime.rs index 9d8e48af2..147e4e0a9 100644 --- a/clients/rust/crates/ahp/src/hosts/runtime.rs +++ b/clients/rust/crates/ahp/src/hosts/runtime.rs @@ -318,21 +318,6 @@ impl HostRuntime { } }; - // Refresh session summaries from `listSessions` — cheap on first - // connect, kept in sync by notifications afterward. Failures are - // non-fatal: we just leave the cache as-is and log. - let summaries: Result = client - .request( - "listSessions", - ListSessionsParams { - channel: ROOT_RESOURCE_URI.to_string(), - meta: None, - limit: None, - cursor: None, - }, - ) - .await; - // Bump generation and install the new client. let new_generation = { let mut state = self.shared.lock().await; @@ -361,14 +346,6 @@ impl HostRuntime { .clone() .unwrap_or_default(); } - if let Ok(list) = summaries { - state.session_summaries.clear(); - for summary in list.items { - state - .session_summaries - .insert(summary.resource.clone(), summary); - } - } state.generation }; @@ -456,11 +433,33 @@ impl HostRuntime { } async fn run_connection(&mut self, mut events: crate::ClientEventStream) -> InnerOutcome { + let (client, generation) = { + let state = self.shared.lock().await; + ( + state.current_client.clone().expect("connected client"), + state.generation, + ) + }; + let refresh = client.request::<_, ListSessionsResult>( + "listSessions", + ListSessionsParams { + channel: ROOT_RESOURCE_URI.to_string(), + meta: None, + limit: None, + cursor: None, + }, + ); + tokio::pin!(refresh); + let mut refreshing = true; + let mut updates = BTreeMap::new(); loop { tokio::select! { _ = self.shutdown_signal.notified() => return InnerOutcome::Shutdown, ev = events.recv() => match ev { Some(event) => { + if refreshing { + record_session_update(&mut updates, &event.event); + } self.handle_event(event).await; } None => return InnerOutcome::Disconnected, @@ -484,6 +483,40 @@ impl HostRuntime { let _ = reply.send(result); } }, + result = &mut refresh, if refreshing => { + // Reconcile the finite queued prefix before replacing the cache. + // Later events are applied normally after this refresh commits. + for event in events.drain_buffered() { + record_session_update(&mut updates, &event.event); + self.handle_event(event).await; + } + refreshing = false; + match result { + Ok(list) => { + let mut summaries: BTreeMap<_, _> = list.items.into_iter() + .map(|summary| (summary.resource.clone(), summary)).collect(); + for (uri, update) in std::mem::take(&mut updates) { + match update { + SessionUpdate::Added(summary) => { summaries.insert(uri, summary); } + SessionUpdate::Removed => { summaries.remove(&uri); } + SessionUpdate::Changed(changes) => { + if let Some(summary) = summaries.get_mut(&uri) { + apply_summary_changes(summary, &changes); + } + } + } + } + let mut state = self.shared.lock().await; + if state.generation == generation && state.current_client.is_some() { + state.session_summaries = summaries; + } + } + Err(err) => tracing::warn!( + host_id = %self.config.id, ?err, "session refresh failed" + ), + } + updates.clear(); + }, } } } @@ -748,6 +781,73 @@ fn apply_summary_changes( if let Some(v) = &changes.changes { existing.changes = Some(v.clone()); } + if let Some(v) = &changes.meta { + existing.meta = Some(v.clone()); + } + if let Some(v) = &changes.chats { + existing.chats = Some(v.clone()); + } + if let Some(v) = &changes.default_chat { + existing.default_chat = Some(v.clone()); + } +} + +enum SessionUpdate { + Added(SessionSummary), + Removed, + Changed(ahp_types::notifications::PartialSessionSummary), +} + +fn record_session_update(updates: &mut BTreeMap, event: &SubscriptionEvent) { + match event { + SubscriptionEvent::SessionAdded(n) => { + updates.insert( + n.summary.resource.clone(), + SessionUpdate::Added(n.summary.clone()), + ); + } + SubscriptionEvent::SessionRemoved(n) => { + updates.insert(n.session.clone(), SessionUpdate::Removed); + } + SubscriptionEvent::SessionSummaryChanged(n) => { + match updates + .entry(n.session.clone()) + .or_insert_with(|| SessionUpdate::Changed(n.changes.clone())) + { + SessionUpdate::Added(summary) => apply_summary_changes(summary, &n.changes), + SessionUpdate::Removed => {} + SessionUpdate::Changed(changes) => merge_summary_changes(changes, &n.changes), + } + } + _ => {} + } +} + +fn merge_summary_changes( + existing: &mut ahp_types::notifications::PartialSessionSummary, + changes: &ahp_types::notifications::PartialSessionSummary, +) { + macro_rules! merge_fields { + ($($field:ident),+ $(,)?) => { + $(if changes.$field.is_some() { existing.$field = changes.$field.clone(); })+ + }; + } + merge_fields!( + provider, + title, + status, + activity, + origin, + modified_at, + created_at, + project, + working_directories, + annotations, + changes, + meta, + chats, + default_chat + ); } // ─── Random helpers (no external dep on `rand`) ───────────────────────────── diff --git a/clients/rust/crates/ahp/src/hosts/types.rs b/clients/rust/crates/ahp/src/hosts/types.rs index ec47aaa12..fde405c28 100644 --- a/clients/rust/crates/ahp/src/hosts/types.rs +++ b/clients/rust/crates/ahp/src/hosts/types.rs @@ -249,9 +249,10 @@ pub struct HostHandle { pub subscriptions: Vec, /// Trigger characters from `InitializeResult.completionTriggerCharacters`. pub completion_trigger_characters: Vec, - /// Cached session summaries keyed by URI. Seeded by `listSessions` - /// after each connect and kept fresh by + /// Cached session summaries keyed by URI. Refreshed concurrently by + /// `listSessions` after handshake readiness and kept fresh by /// `root/sessionAdded`/`Removed`/`SummaryChanged` notifications. + /// May be empty or retain the prior cache while discovery is pending. pub session_summaries: Vec, /// Generation counter — bumped on every `connect` or `reconnect`. /// [`HostClientHandle`]s carry the generation they were issued at, diff --git a/clients/rust/crates/ahp/src/lib.rs b/clients/rust/crates/ahp/src/lib.rs index 9ee30d67e..04e8194f9 100644 --- a/clients/rust/crates/ahp/src/lib.rs +++ b/clients/rust/crates/ahp/src/lib.rs @@ -141,9 +141,8 @@ //! dropped; [`Client::shutdown`] requests graceful transport closure. //! Dropping the last client resolves in-flight requests with //! [`ClientError::Shutdown`]. Explicit shutdown and transport closure retain -//! the normal `-32000` [`ClientError::Rpc`] teardown error. Non-owning -//! [`WeakPingHandle`] requests resolve with [`ClientError::Shutdown`] for -//! each of these lifecycle endings. +//! the normal `-32000` [`ClientError::Rpc`] teardown error. Automatic idle +//! keepalive is owned by the same driver and ends with it. #![forbid(unsafe_code)] #![warn(missing_docs)] @@ -159,9 +158,9 @@ pub mod transport; pub use ahp_types; pub use client::{ - Client, ClientConfig, ClientEvent, ClientEventStream, DispatchHandle, ResourceRequestHandlers, - ServerRequestFuture, ServerRequestHandler, SessionSubscription, SubscriptionEvent, - WeakPingHandle, + Client, ClientConfig, ClientEvent, ClientEventStream, DispatchHandle, KeepaliveConfig, + ResourceRequestHandlers, ServerRequestFuture, ServerRequestHandler, SessionSubscription, + SubscriptionEvent, }; pub use error::{ClientError, TransportError}; pub use multi_host_state_mirror::{HostedResourceKey, MultiHostStateMirror}; diff --git a/clients/rust/crates/ahp/src/transport.rs b/clients/rust/crates/ahp/src/transport.rs index 3a1e4d790..3bf3a552e 100644 --- a/clients/rust/crates/ahp/src/transport.rs +++ b/clients/rust/crates/ahp/src/transport.rs @@ -64,7 +64,6 @@ use std::future::Future; use std::pin::Pin; -use crate::client::WeakPingHandle; use crate::error::TransportError; use ahp_types::messages::JsonRpcMessage; @@ -113,15 +112,6 @@ impl TransportMessage { /// client sends indefinitely until the underlying connection closes, /// and `recv` signals closure by returning `None`. pub trait Transport: Send + 'static { - /// Bind a non-owning ping handle before the client starts transport I/O. - /// - /// Called once by [`crate::Client::connect`], not on subsequent - /// initialize/reconnect requests on the same client. A transport may pass - /// the handle to its keepalive task without keeping the client alive. - /// The transport remains responsible for keepalive timing and negotiation. - /// The default implementation is a no-op. - fn bind_client(&mut self, _ping: WeakPingHandle) {} - /// Send a single message. /// /// Errors returned here are typically fatal for the transport @@ -164,9 +154,6 @@ pub trait Transport: Send + 'static { /// allocation cost (typically: registries that hold one transport per /// host). pub trait DynTransport: Send + 'static { - /// Object-safe analogue of [`Transport::bind_client`]. - fn bind_client(&mut self, _ping: WeakPingHandle) {} - /// Object-safe analogue of [`Transport::send`]. fn send<'a>( &'a mut self, @@ -185,10 +172,6 @@ pub trait DynTransport: Send + 'static { } impl DynTransport for T { - fn bind_client(&mut self, ping: WeakPingHandle) { - ::bind_client(self, ping); - } - fn send<'a>( &'a mut self, msg: TransportMessage, @@ -264,10 +247,6 @@ impl std::fmt::Debug for BoxedTransport { } impl Transport for BoxedTransport { - fn bind_client(&mut self, ping: WeakPingHandle) { - self.inner.bind_client(ping); - } - fn send( &mut self, msg: TransportMessage, diff --git a/clients/rust/crates/ahp/tests/connection_lifecycle.rs b/clients/rust/crates/ahp/tests/connection_lifecycle.rs new file mode 100644 index 000000000..5aa3f947a --- /dev/null +++ b/clients/rust/crates/ahp/tests/connection_lifecycle.rs @@ -0,0 +1,689 @@ +#![allow(clippy::panic, clippy::unwrap_used)] + +use std::collections::VecDeque; +use std::future::Future; +use std::sync::Arc; +use std::time::Duration; + +use ahp::hosts::{HostConfig, HostEvent, HostId, MultiHostClient, ReconnectPolicy}; +use ahp::{ + BoxedTransport, Client, ClientConfig, ClientError, KeepaliveConfig, Transport, TransportError, + TransportMessage, +}; +use ahp_types::messages::{ + JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcSuccessResponse, JsonRpcVersion, +}; +use serde_json::{json, Value}; +use tokio::sync::{mpsc, Mutex}; + +struct MemTransport { + tx: mpsc::Sender, + rx: mpsc::Receiver, +} + +struct Peer { + tx: mpsc::Sender, + rx: mpsc::Receiver, +} + +fn pair() -> (MemTransport, Peer) { + let (to_peer, rx) = mpsc::channel(16); + let (tx, from_peer) = mpsc::channel(16); + ( + MemTransport { + tx: to_peer, + rx: from_peer, + }, + Peer { tx, rx }, + ) +} + +impl Transport for MemTransport { + async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { + self.tx + .send(message) + .await + .map_err(|_| TransportError::Closed) + } + + async fn recv(&mut self) -> Result, TransportError> { + Ok(self.rx.recv().await) + } +} + +async fn bounded(future: impl Future) -> T { + tokio::time::timeout(Duration::from_secs(2), future) + .await + .expect("fixture timed out") +} + +async fn settle() { + for _ in 0..20 { + tokio::task::yield_now().await; + } +} + +impl Peer { + async fn request(&mut self) -> JsonRpcRequest { + let JsonRpcMessage::Request(request) = bounded(self.rx.recv()) + .await + .expect("driver closed") + .into_parsed() + .unwrap() + else { + panic!("expected request") + }; + request + } + + async fn reply(&self, id: u64, result: Value) { + self.tx + .send(TransportMessage::Parsed(JsonRpcMessage::SuccessResponse( + JsonRpcSuccessResponse { + jsonrpc: JsonRpcVersion::V2, + id, + result, + }, + ))) + .await + .unwrap(); + } + + async fn notify(&self, method: &str, params: Value) { + self.tx + .send(TransportMessage::Parsed(JsonRpcMessage::Notification( + JsonRpcNotification { + jsonrpc: JsonRpcVersion::V2, + method: method.into(), + params: Some(params), + }, + ))) + .await + .unwrap(); + } +} + +fn config(keepalive: bool) -> ClientConfig { + ClientConfig { + default_request_timeout: None, + keepalive: keepalive.then_some(KeepaliveConfig { + idle_interval: Duration::from_secs(10), + liveness_timeout: Duration::from_secs(30), + }), + ..ClientConfig::default() + } +} + +#[tokio::test(start_paused = true)] +async fn idle_ping_uses_normal_ids_and_continued_silence_closes_connection() { + let (transport, mut peer) = pair(); + let client = Client::connect(BoxedTransport::new(transport), config(true)) + .await + .unwrap(); + let mut events = client.events(); + let request = client.request::<_, Value>("ordinary", ()); + tokio::pin!(request); + let ordinary = tokio::select! { + result = &mut request => panic!("premature result: {result:?}"), + request = peer.request() => request, + }; + tokio::time::advance(Duration::from_secs(9)).await; + settle().await; + assert!(peer.rx.try_recv().is_err()); + tokio::time::advance(Duration::from_secs(1)).await; + let ping = peer.request().await; + assert_eq!(ping.method, "ping"); + assert_eq!(ping.id, ordinary.id + 1); + assert_eq!(ping.params.unwrap()["channel"], "ahp-root://"); + tokio::time::advance(Duration::from_secs(19)).await; + settle().await; + assert!(peer.rx.try_recv().is_err(), "at most one unanswered ping"); + tokio::time::advance(Duration::from_secs(1)).await; + assert!(bounded(peer.rx.recv()).await.is_none()); + assert!(bounded(events.recv()).await.is_none()); + assert!(matches!(bounded(request).await, Err(ClientError::Rpc(error)) if error.code == -32000)); +} + +#[tokio::test(start_paused = true)] +async fn inbound_wire_traffic_suppresses_pings_and_retires_unanswered_ping() { + let (transport, mut peer) = pair(); + let client = Client::connect(transport, config(true)).await.unwrap(); + for _ in 0..5 { + tokio::time::advance(Duration::from_secs(9)).await; + peer.notify("activity", json!({})).await; + settle().await; + assert!(peer.rx.try_recv().is_err()); + } + + tokio::time::advance(Duration::from_secs(10)).await; + let first = peer.request().await; + assert_eq!(first.method, "ping"); + peer.notify("activity", json!({})).await; + settle().await; + tokio::time::advance(Duration::from_secs(10)).await; + let second = peer.request().await; + assert_eq!(second.id, first.id + 1); + peer.reply(first.id, Value::Null).await; + peer.reply(second.id, Value::Null).await; + settle().await; + tokio::time::advance(Duration::from_secs(9)).await; + settle().await; + assert!(peer.rx.try_recv().is_err()); + client.shutdown().await; + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn continuously_ready_inbound_traffic_does_not_starve_outbound_requests() { + struct BusyTransport { + sent: mpsc::UnboundedSender, + } + impl Transport for BusyTransport { + async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { + self.sent.send(message).map_err(|_| TransportError::Closed) + } + async fn recv(&mut self) -> Result, TransportError> { + tokio::task::consume_budget().await; + Ok(Some(TransportMessage::Text( + r#"{"jsonrpc":"2.0","method":"activity"}"#.into(), + ))) + } + } + let (sent, mut rx) = mpsc::unbounded_channel(); + let client = Client::connect(BusyTransport { sent }, config(true)) + .await + .unwrap(); + client.notify("outbound", ()).await.unwrap(); + let message = bounded(rx.recv()).await.unwrap().into_parsed().unwrap(); + assert!( + matches!(message, JsonRpcMessage::Notification(message) if message.method == "outbound") + ); + client.shutdown().await; + assert!(bounded(rx.recv()).await.is_none()); +} + +#[tokio::test(start_paused = true)] +async fn outbound_traffic_does_not_prove_inbound_liveness_and_keepalive_can_be_disabled() { + for enabled in [false, true] { + let (transport, mut peer) = pair(); + let client = Client::connect(transport, config(enabled)).await.unwrap(); + for _ in 0..3 { + client.notify("outbound", ()).await.unwrap(); + assert!(bounded(peer.rx.recv()).await.is_some()); + tokio::time::advance(Duration::from_secs(9)).await; + settle().await; + if enabled && !peer.rx.is_empty() { + assert_eq!(peer.request().await.method, "ping"); + } + } + tokio::time::advance(Duration::from_secs(3)).await; + settle().await; + if enabled { + assert!(bounded(peer.rx.recv()).await.is_none()); + } else { + assert!(peer.rx.try_recv().is_err()); + client.shutdown().await; + assert!(bounded(peer.rx.recv()).await.is_none()); + } + } +} + +#[tokio::test] +async fn invalid_keepalive_policy_is_rejected_before_transport_io() { + for policy in [ + KeepaliveConfig { + idle_interval: Duration::ZERO, + liveness_timeout: Duration::from_secs(1), + }, + KeepaliveConfig { + idle_interval: Duration::from_secs(1), + liveness_timeout: Duration::from_secs(1), + }, + KeepaliveConfig { + idle_interval: Duration::from_secs(2), + liveness_timeout: Duration::from_secs(1), + }, + KeepaliveConfig { + idle_interval: Duration::from_secs(1), + liveness_timeout: Duration::MAX, + }, + ] { + let (transport, mut peer) = pair(); + let result = Client::connect( + transport, + ClientConfig { + keepalive: Some(policy), + ..config(false) + }, + ) + .await; + assert!(matches!( + result, + Err(ClientError::Transport(TransportError::Protocol(_))) + )); + assert!(peer.rx.recv().await.is_none()); + } +} + +async fn add_host( + multi: &MultiHostClient, + transports: Vec>, + keepalive: bool, +) { + let transports = Arc::new(Mutex::new(VecDeque::from(transports))); + let config = HostConfig::new("host", "Host", move |_| { + let transports = transports.clone(); + async move { + transports + .lock() + .await + .pop_front() + .unwrap() + .map(BoxedTransport::new) + } + }) + .with_client_config(config(keepalive)) + .with_reconnect_policy(ReconnectPolicy::immediate_forever()); + multi.add_host(config).await.unwrap(); +} + +async fn handshake(peer: &mut Peer, reconnect: bool) -> JsonRpcRequest { + let request = peer.request().await; + assert_eq!( + request.method, + if reconnect { "reconnect" } else { "initialize" } + ); + peer.reply(request.id, if reconnect { + json!({"type": "replay", "actions": [], "missing": []}) + } else { + json!({"protocolVersion": ahp_types::PROTOCOL_VERSION, "serverSeq": 10, "snapshots": []}) + }).await; + request +} + +async fn connected(events: &mut ahp::hosts::HostEventStream) { + loop { + if matches!( + bounded(events.recv()).await, + Some(HostEvent::Connected { .. }) + ) { + break; + } + } +} + +fn summary(uri: &str, title: &str) -> Value { + json!({"resource": uri, "provider": "test", "title": title, "status": 0, + "createdAt": "1970-01-01T00:00:00Z", "modifiedAt": "1970-01-01T00:00:00Z"}) +} + +#[tokio::test(start_paused = true)] +async fn handshake_publishes_client_and_keepalive_works_while_discovery_is_pending() { + let (transport, mut peer) = pair(); + let multi = MultiHostClient::new(); + let mut events = multi.host_events(); + add_host(&multi, vec![Ok(transport)], true).await; + handshake(&mut peer, false).await; + connected(&mut events).await; + let discovery = peer.request().await; + assert_eq!(discovery.method, "listSessions"); + let handle = multi + .client(&HostId::new("host")) + .await + .expect("ready before discovery"); + assert!(multi + .host(&HostId::new("host")) + .await + .unwrap() + .state + .is_connected()); + let ordinary = handle.request::<_, Value>("ordinary", ()); + tokio::pin!(ordinary); + let request = tokio::select! { + result = &mut ordinary => panic!("premature result: {result:?}"), + request = peer.request() => request, + }; + peer.reply(request.id, json!("usable")).await; + assert_eq!(bounded(ordinary).await.unwrap(), json!("usable")); + tokio::time::advance(Duration::from_secs(10)).await; + let ping = peer.request().await; + assert_eq!(ping.method, "ping"); + assert_ne!(ping.id, discovery.id); + peer.reply(ping.id, Value::Null).await; + settle().await; + bounded(multi.remove_host(&HostId::new("host"))) + .await + .unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn delayed_discovery_preserves_queued_additions_changes_and_removals() { + let (transport, mut peer) = pair(); + let multi = MultiHostClient::new(); + let mut events = multi.events(); + let mut hosts = multi.host_events(); + add_host(&multi, vec![Ok(transport)], false).await; + let init = peer.request().await; + // Notifications before the handshake completes must already be captured. + peer.notify( + "root/sessionAdded", + json!({"channel":"ahp-root://", "summary":summary("added","first")}), + ) + .await; + peer.reply( + init.id, + json!({"protocolVersion": ahp_types::PROTOCOL_VERSION, "serverSeq":10, "snapshots":[]}), + ) + .await; + connected(&mut hosts).await; + let discovery = peer.request().await; + bounded(events.recv()).await.unwrap(); + for (method, params) in [ + ( + "root/sessionSummaryChanged", + json!({"session":"added", "changes":{"title":"new"}, "channel":"ahp-root://"}), + ), + ( + "root/sessionSummaryChanged", + json!({"session":"changed", "changes":{"title":"fresh", "_meta":{"test":true}}, "channel":"ahp-root://"}), + ), + ( + "root/sessionSummaryChanged", + json!({"session":"changed", "changes":{"activity":"busy"}, "channel":"ahp-root://"}), + ), + ( + "root/sessionRemoved", + json!({"session":"removed", "channel":"ahp-root://"}), + ), + ( + "root/sessionAdded", + json!({"summary":summary("readded","new"), "channel":"ahp-root://"}), + ), + ( + "root/sessionRemoved", + json!({"session":"readded", "channel":"ahp-root://"}), + ), + ( + "root/sessionAdded", + json!({"summary":summary("readded","latest"), "channel":"ahp-root://"}), + ), + ] { + peer.notify(method, params).await; + bounded(events.recv()).await.unwrap(); + } + peer.reply( + discovery.id, + json!({"items":[summary("changed","old"),summary("removed","old"),summary("added","old")]}), + ) + .await; + bounded(async { + loop { + let sessions = multi + .host(&HostId::new("host")) + .await + .unwrap() + .session_summaries; + if sessions.len() == 3 && sessions.iter().any(|s| s.resource == "changed") { + assert!(!sessions.iter().any(|s| s.resource == "removed")); + assert_eq!( + sessions + .iter() + .find(|s| s.resource == "added") + .unwrap() + .title, + "new" + ); + assert_eq!( + sessions + .iter() + .find(|s| s.resource == "readded") + .unwrap() + .title, + "latest" + ); + let changed = sessions.iter().find(|s| s.resource == "changed").unwrap(); + assert_eq!(changed.title, "fresh"); + assert_eq!(changed.activity.as_deref(), Some("busy")); + assert_eq!(changed.meta.as_ref().unwrap()["test"], json!(true)); + break; + } + tokio::task::yield_now().await; + } + }) + .await; + bounded(multi.remove_host(&HostId::new("host"))) + .await + .unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn failed_refresh_keeps_host_ready_and_live_cache() { + let (transport, mut peer) = pair(); + let multi = MultiHostClient::new(); + let mut events = multi.events(); + add_host(&multi, vec![Ok(transport)], false).await; + handshake(&mut peer, false).await; + let discovery = peer.request().await; + peer.notify( + "root/sessionAdded", + json!({"channel":"ahp-root://", "summary":summary("live","Live")}), + ) + .await; + bounded(events.recv()).await.unwrap(); + // A bad result is a nonfatal refresh error, not a readiness failure. + peer.reply(discovery.id, json!({"bad":"result"})).await; + settle().await; + assert!(multi.client(&HostId::new("host")).await.is_some()); + assert_eq!( + multi + .host(&HostId::new("host")) + .await + .unwrap() + .session_summaries[0] + .title, + "Live" + ); + bounded(multi.remove_host(&HostId::new("host"))) + .await + .unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn continuous_session_updates_allow_discovery_and_supervisor_commands_to_progress() { + struct BusyHost { + replies: VecDeque, + } + impl Transport for BusyHost { + async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { + let JsonRpcMessage::Request(request) = message.into_parsed()? else { + panic!("expected request"); + }; + let result = match request.method.as_str() { + "initialize" => json!({ + "protocolVersion": ahp_types::PROTOCOL_VERSION, + "serverSeq": 10, + "snapshots": [] + }), + "listSessions" => json!({"items":[summary("listed","Listed")]}), + method => panic!("unexpected request: {method}"), + }; + self.replies + .push_back(TransportMessage::Parsed(JsonRpcMessage::SuccessResponse( + JsonRpcSuccessResponse { + jsonrpc: JsonRpcVersion::V2, + id: request.id, + result, + }, + ))); + Ok(()) + } + + async fn recv(&mut self) -> Result, TransportError> { + Ok(Some(self.replies.pop_front().unwrap_or_else(|| { + TransportMessage::Parsed(JsonRpcMessage::Notification(JsonRpcNotification { + jsonrpc: JsonRpcVersion::V2, + method: "root/sessionSummaryChanged".into(), + params: Some(json!({ + "channel": "ahp-root://", + "session": "busy", + "changes": {"title": "Live"} + })), + })) + }))) + } + } + let (replacement, mut peer) = pair(); + let transports = Arc::new(Mutex::new(VecDeque::from([ + BoxedTransport::new(BusyHost { + replies: VecDeque::new(), + }), + BoxedTransport::new(replacement), + ]))); + let multi = MultiHostClient::new(); + let host = HostId::new("host"); + let config = HostConfig::new("host", "Host", move |_| { + let transports = transports.clone(); + async move { Ok(transports.lock().await.pop_front().unwrap()) } + }) + .with_client_config(config(false)) + .with_reconnect_policy(ReconnectPolicy::immediate_forever()); + multi.add_host(config).await.unwrap(); + bounded(async { + loop { + if multi + .host(&host) + .await + .unwrap() + .session_summaries + .iter() + .any(|summary| summary.resource == "listed") + { + break; + } + tokio::task::yield_now().await; + } + }) + .await; + bounded(multi.reconnect_host(&host)).await.unwrap(); + handshake(&mut peer, true).await; + assert_eq!(peer.request().await.method, "listSessions"); + bounded(multi.remove_host(&host)).await.unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); +} + +#[tokio::test] +async fn superseded_refresh_and_old_handshake_reply_cannot_cross_connection_epoch() { + let (first, mut old) = pair(); + let (second, mut peer) = pair(); + let (third, mut retry) = pair(); + let multi = MultiHostClient::new(); + let mut events = multi.host_events(); + add_host( + &multi, + vec![ + Err(TransportError::Closed), + Ok(first), + Ok(second), + Ok(third), + ], + false, + ) + .await; + let abandoned = old.request().await; + assert_eq!(abandoned.id, 1); + drop(old.tx); + assert!(bounded(old.rx.recv()).await.is_none()); + let init = peer.request().await; + peer.reply( + abandoned.id, + json!({"protocolVersion":ahp_types::PROTOCOL_VERSION,"serverSeq":99,"snapshots":[]}), + ) + .await; + settle().await; + assert!( + peer.rx.try_recv().is_err(), + "stale response must not finish handshake" + ); + assert_eq!(init.id, 2); + peer.reply( + init.id, + json!({"protocolVersion":ahp_types::PROTOCOL_VERSION,"serverSeq":10,"snapshots":[]}), + ) + .await; + connected(&mut events).await; + let old_refresh = peer.request().await; + assert_eq!(old_refresh.method, "listSessions"); + multi.reconnect_host(&HostId::new("host")).await.unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); + let reconnect = retry.request().await; + assert_eq!(reconnect.method, "reconnect"); + assert!(reconnect.id > old_refresh.id); + retry + .reply(old_refresh.id, json!({"items":[summary("old","Stale")]})) + .await; + settle().await; + assert!(retry.rx.try_recv().is_err()); + retry + .reply( + reconnect.id, + json!({"type":"replay","actions":[],"missing":[]}), + ) + .await; + connected(&mut events).await; + let refresh = retry.request().await; + retry + .reply(refresh.id, json!({"items":[summary("new","Current")]})) + .await; + bounded(async { + loop { + let host = multi.host(&HostId::new("host")).await.unwrap(); + if !host.session_summaries.is_empty() { + assert_eq!(host.session_summaries.len(), 1); + assert_eq!(host.session_summaries[0].title, "Current"); + break; + } + tokio::task::yield_now().await; + } + }) + .await; + bounded(multi.remove_host(&HostId::new("host"))) + .await + .unwrap(); + assert!(bounded(retry.rx.recv()).await.is_none()); +} + +#[tokio::test(start_paused = true)] +async fn liveness_failure_reconnects_and_restarts_keepalive_before_discovery_completes() { + let (first, mut old) = pair(); + let (second, mut peer) = pair(); + let multi = MultiHostClient::new(); + let mut events = multi.host_events(); + add_host(&multi, vec![Ok(first), Ok(second)], true).await; + handshake(&mut old, false).await; + connected(&mut events).await; + let discovery = old.request().await; + assert_eq!(discovery.method, "listSessions"); + tokio::time::advance(Duration::from_secs(10)).await; + let ping = old.request().await; + assert_eq!(ping.method, "ping"); + tokio::time::advance(Duration::from_secs(20)).await; + assert!(bounded(old.rx.recv()).await.is_none()); + let reconnect = handshake(&mut peer, true).await; + assert!(reconnect.id > ping.id); + connected(&mut events).await; + let discovery = peer.request().await; + assert_eq!(discovery.method, "listSessions"); + assert!(multi.client(&HostId::new("host")).await.is_some()); + tokio::time::advance(Duration::from_secs(10)).await; + let ping = peer.request().await; + assert_eq!(ping.method, "ping"); + peer.reply(ping.id, Value::Null).await; + settle().await; + bounded(multi.remove_host(&HostId::new("host"))) + .await + .unwrap(); + assert!(bounded(peer.rx.recv()).await.is_none()); +} diff --git a/clients/rust/crates/ahp/tests/weak_ping.rs b/clients/rust/crates/ahp/tests/weak_ping.rs deleted file mode 100644 index 7c2131120..000000000 --- a/clients/rust/crates/ahp/tests/weak_ping.rs +++ /dev/null @@ -1,589 +0,0 @@ -#![allow(clippy::panic, clippy::unwrap_used)] - -use std::collections::VecDeque; -use std::future::Future; -use std::pin::Pin; -use std::sync::atomic::{AtomicUsize, Ordering}; -use std::sync::Arc; -use std::time::Duration; - -use ahp::hosts::{HostConfig, HostEvent, HostId, MultiHostClient, ReconnectPolicy}; -use ahp::{ - BoxedTransport, Client, ClientConfig, ClientError, DynTransport, Transport, TransportError, - TransportMessage, WeakPingHandle, -}; -use ahp_types::messages::{ - JsonRpcError, JsonRpcErrorResponse, JsonRpcMessage, JsonRpcRequest, JsonRpcSuccessResponse, - JsonRpcVersion, -}; -use serde_json::{json, Value}; -use tokio::sync::{mpsc, Mutex}; - -struct BoundTransport { - tx: mpsc::Sender, - rx: mpsc::Receiver, - bindings: mpsc::UnboundedSender, - bind_count: Arc, -} - -struct Peer { - tx: mpsc::Sender, - rx: mpsc::Receiver, - bindings: mpsc::UnboundedReceiver, - bind_count: Arc, -} - -fn pair() -> (BoundTransport, Peer) { - let (to_peer, rx) = mpsc::channel(16); - let (tx, from_peer) = mpsc::channel(16); - let (bindings, bound) = mpsc::unbounded_channel(); - let bind_count = Arc::new(AtomicUsize::new(0)); - ( - BoundTransport { - tx: to_peer, - rx: from_peer, - bindings, - bind_count: bind_count.clone(), - }, - Peer { - tx, - rx, - bindings: bound, - bind_count, - }, - ) -} - -impl Transport for BoundTransport { - fn bind_client(&mut self, ping: WeakPingHandle) { - self.bind_count.fetch_add(1, Ordering::SeqCst); - self.bindings.send(ping).unwrap(); - } - - async fn send(&mut self, message: TransportMessage) -> Result<(), TransportError> { - assert_eq!(self.bind_count.load(Ordering::SeqCst), 1); - self.tx - .send(message) - .await - .map_err(|_| TransportError::Closed) - } - - async fn recv(&mut self) -> Result, TransportError> { - Ok(self.rx.recv().await) - } -} - -struct LegacyDynTransport { - tx: mpsc::Sender, - rx: mpsc::Receiver, -} - -impl DynTransport for LegacyDynTransport { - fn send<'a>( - &'a mut self, - message: TransportMessage, - ) -> Pin> + Send + 'a>> { - Box::pin(async move { - self.tx - .send(message) - .await - .map_err(|_| TransportError::Closed) - }) - } - - fn recv<'a>( - &'a mut self, - ) -> Pin, TransportError>> + Send + 'a>> - { - Box::pin(async move { Ok(self.rx.recv().await) }) - } - - fn close<'a>( - &'a mut self, - ) -> Pin> + Send + 'a>> { - Box::pin(async { Ok(()) }) - } -} - -#[tokio::test] -async fn legacy_object_safe_transport_needs_no_binding_implementation() { - let (transport, mut peer) = pair(); - let legacy = LegacyDynTransport { - tx: transport.tx, - rx: transport.rx, - }; - let client = Client::connect(BoxedTransport::from_dyn(Box::new(legacy)), no_timeout()) - .await - .unwrap(); - let mut request = Box::pin(client.ping()); - let sent = tokio::select! { - result = &mut request => panic!("premature result: {result:?}"), - sent = peer.request() => sent, - }; - peer.reply(sent.id, Value::Null).await; - bounded(request).await.unwrap(); - drop(client); - assert!(bounded(peer.rx.recv()).await.is_none()); -} - -async fn bounded(future: impl Future) -> T { - tokio::time::timeout(Duration::from_secs(2), future) - .await - .expect("fixture timed out") -} - -impl Peer { - async fn request(&mut self) -> JsonRpcRequest { - let wire = bounded(self.rx.recv()).await.expect("driver closed"); - let JsonRpcMessage::Request(request) = wire.into_parsed().unwrap() else { - panic!("expected request") - }; - request - } - - async fn reply(&self, id: u64, result: Value) { - bounded( - self.tx - .send(TransportMessage::Parsed(JsonRpcMessage::SuccessResponse( - JsonRpcSuccessResponse { - jsonrpc: JsonRpcVersion::V2, - id, - result, - }, - ))), - ) - .await - .unwrap(); - } - - async fn bound(&mut self) -> WeakPingHandle { - bounded(self.bindings.recv()).await.unwrap() - } -} - -fn no_timeout() -> ClientConfig { - ClientConfig { - default_request_timeout: None, - ..ClientConfig::default() - } -} - -#[tokio::test] -async fn boxed_binding_precedes_handshakes_and_is_not_repeated() { - for from_dyn in [false, true] { - let (transport, mut peer) = pair(); - let transport = if from_dyn { - let dynamic: Box = Box::new(transport); - BoxedTransport::from_dyn(dynamic) - } else { - BoxedTransport::new(transport) - }; - let client = Client::connect(transport, no_timeout()).await.unwrap(); - let ping = peer - .bindings - .try_recv() - .expect("bind before connect returns"); - assert!(peer.rx.try_recv().is_err(), "binding must not send a ping"); - - for method in ["initialize", "reconnect", "initialize"] { - let request = client.request::<_, Value>(method, json!({})); - tokio::pin!(request); - let sent = tokio::select! { - request = &mut request => panic!("premature result: {request:?}"), - sent = peer.request() => sent, - }; - assert_eq!(sent.method, method); - peer.reply(sent.id, json!({"handshake": method})).await; - assert_eq!(bounded(request).await.unwrap()["handshake"], method); - } - assert_eq!(peer.bind_count.load(Ordering::SeqCst), 1); - assert!(peer.bindings.try_recv().is_err()); - drop(client); - assert!(bounded(peer.rx.recv()).await.is_none()); - assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); - } -} - -#[tokio::test] -async fn weak_and_normal_pings_share_ids_and_out_of_order_correlation() { - let (transport, mut peer) = pair(); - let client = Client::connect(BoxedTransport::new(transport), no_timeout()) - .await - .unwrap(); - let ping = peer.bound().await; - let ordinary = client.request::<_, Value>("listSessions", json!({"channel": "ahp-root://"})); - let heartbeat = ping.ping(); - let normal_ping = client.ping(); - tokio::pin!(ordinary, heartbeat, normal_ping); - - let requests = async { - let mut requests = Vec::new(); - for _ in 0..3 { - requests.push(peer.request().await); - } - requests - }; - let requests = tokio::select! { - result = &mut ordinary => panic!("premature result: {result:?}"), - result = &mut heartbeat => panic!("premature result: {result:?}"), - result = &mut normal_ping => panic!("premature result: {result:?}"), - requests = requests => requests, - }; - let mut ids: Vec<_> = requests.iter().map(|request| request.id).collect(); - ids.sort_unstable(); - assert_eq!(ids, [1, 2, 3]); - let discovery = requests - .iter() - .find(|r| r.method == "listSessions") - .unwrap(); - for request in requests.iter().rev().filter(|r| r.method == "ping") { - assert_eq!(request.params.as_ref().unwrap()["channel"], "ahp-root://"); - peer.reply(request.id, Value::Null).await; - } - bounded(heartbeat).await.unwrap(); - bounded(normal_ping).await.unwrap(); - // Discovery is still unanswered while both heartbeat responses resolve. - assert!(tokio::time::timeout(Duration::ZERO, &mut ordinary) - .await - .is_err()); - peer.reply(discovery.id, json!({"items": []})).await; - assert_eq!(bounded(ordinary).await.unwrap(), json!({"items": []})); - client.shutdown().await; - assert!(bounded(peer.rx.recv()).await.is_none()); -} - -#[tokio::test] -async fn weak_ping_works_while_managed_host_discovery_blocks_client_access() { - let (transport, mut peer) = pair(); - let transport = Arc::new(Mutex::new(Some(transport))); - let config = HostConfig::new("discovering", "Discovering host", move |_| { - let transport = transport.clone(); - async move { Ok(BoxedTransport::new(transport.lock().await.take().unwrap())) } - }) - .with_client_config(no_timeout()) - .with_reconnect_policy(ReconnectPolicy::disabled()); - let multi = MultiHostClient::new(); - multi.add_host(config).await.unwrap(); - let ping = peer.bound().await; - let initialize = peer.request().await; - assert_eq!(initialize.method, "initialize"); - peer.reply( - initialize.id, - json!({"protocolVersion": ahp_types::PROTOCOL_VERSION, "serverSeq": 0, "snapshots": []}), - ) - .await; - let discovery = peer.request().await; - assert_eq!(discovery.method, "listSessions"); - assert!(multi.client(&HostId::from("discovering")).await.is_none()); - - let heartbeat = ping.ping(); - tokio::pin!(heartbeat); - let sent = tokio::select! { - result = &mut heartbeat => panic!("premature ping: {result:?}"), - sent = peer.request() => sent, - }; - assert_eq!(sent.method, "ping"); - assert_ne!(sent.id, initialize.id); - assert_ne!(sent.id, discovery.id); - peer.reply(sent.id, Value::Null).await; - bounded(heartbeat).await.unwrap(); - assert!(multi.client(&HostId::from("discovering")).await.is_none()); - bounded(multi.remove_host(&HostId::from("discovering"))) - .await - .unwrap(); - assert!(bounded(peer.rx.recv()).await.is_none()); - assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); -} - -#[tokio::test] -async fn unanswered_weak_ping_does_not_retain_last_client_or_driver() { - let (transport, mut peer) = pair(); - let client = Client::connect(BoxedTransport::new(transport), no_timeout()) - .await - .unwrap(); - let clone = client.clone(); - let ping = peer.bound().await; - let heartbeat = ping.ping(); - tokio::pin!(heartbeat); - let sent = tokio::select! { - result = &mut heartbeat => panic!("premature ping: {result:?}"), - sent = peer.request() => sent, - }; - assert_eq!(sent.method, "ping"); - drop(client); - assert!(peer.rx.try_recv().is_err()); - drop(clone); - assert!(matches!( - bounded(heartbeat).await, - Err(ClientError::Shutdown) - )); - assert!(bounded(peer.rx.recv()).await.is_none(), "driver leaked"); - assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); -} - -#[tokio::test] -async fn shutdown_and_transport_close_preserve_normal_errors_but_cancel_weak_ping() { - for explicit in [false, true] { - let (transport, mut peer) = pair(); - let client = Client::connect(transport, no_timeout()).await.unwrap(); - let ping = peer.bound().await; - let normal = client.ping(); - let weak = ping.ping(); - tokio::pin!(normal, weak); - let requests = async { - peer.request().await; - peer.request().await; - }; - tokio::select! { - result = &mut normal => panic!("premature ping: {result:?}"), - result = &mut weak => panic!("premature ping: {result:?}"), - _ = requests => {}, - } - if explicit { - client.shutdown().await; - } else { - drop(peer.tx); - } - let Err(ClientError::Rpc(error)) = bounded(normal).await else { - panic!("normal ping must retain its existing RPC teardown error"); - }; - assert_eq!(error.code, -32000); - assert_eq!( - error.message, - if explicit { - "client shut down" - } else { - "transport closed" - } - ); - assert!(matches!(bounded(weak).await, Err(ClientError::Shutdown))); - assert!(matches!(ping.ping().await, Err(ClientError::Shutdown))); - assert!(bounded(peer.rx.recv()).await.is_none()); - } -} - -#[tokio::test] -async fn weak_ping_preserves_server_error_and_configured_timeout() { - let (transport, mut peer) = pair(); - let client = Client::connect(transport, no_timeout()).await.unwrap(); - let ping = peer.bound().await; - let heartbeat = ping.ping(); - tokio::pin!(heartbeat); - let request = tokio::select! { - result = &mut heartbeat => panic!("premature ping: {result:?}"), - request = peer.request() => request, - }; - peer.tx - .send(TransportMessage::Parsed(JsonRpcMessage::ErrorResponse( - JsonRpcErrorResponse { - jsonrpc: JsonRpcVersion::V2, - id: request.id, - error: JsonRpcError { - code: -32601, - message: "unsupported ping".into(), - data: None, - }, - }, - ))) - .await - .unwrap(); - let Err(ClientError::Rpc(error)) = bounded(heartbeat).await else { - panic!("expected server RPC error"); - }; - assert_eq!(error.code, -32601); - drop(client); - assert!(bounded(peer.rx.recv()).await.is_none()); - - let (transport, mut peer) = pair(); - let client = Client::connect( - transport, - ClientConfig { - default_request_timeout: Some(Duration::ZERO), - ..ClientConfig::default() - }, - ) - .await - .unwrap(); - let ping = peer.bound().await; - assert!(matches!( - bounded(ping.ping()).await, - Err(ClientError::Cancelled) - )); - assert_eq!(peer.request().await.method, "ping"); - drop(client); - assert!(bounded(peer.rx.recv()).await.is_none()); -} - -#[tokio::test] -async fn managed_host_ids_survive_factory_failure_handshake_failure_and_reconnect() { - let (first, mut first_peer) = pair(); - let (second, mut second_peer) = pair(); - let (third, mut third_peer) = pair(); - let transports = Arc::new(Mutex::new(VecDeque::from([ - Err(TransportError::Closed), - Ok(first), - Ok(second), - Ok(third), - ]))); - let attempts = Arc::new(AtomicUsize::new(0)); - let factory_attempts = attempts.clone(); - let config = HostConfig::new("retained", "Retained delivery host", move |_| { - let transports = transports.clone(); - factory_attempts.fetch_add(1, Ordering::SeqCst); - async move { - transports - .lock() - .await - .pop_front() - .unwrap() - .map(BoxedTransport::new) - } - }) - .with_client_config(no_timeout()) - .with_reconnect_policy(ReconnectPolicy::immediate_forever()); - let multi = MultiHostClient::new(); - let mut events = multi.host_events(); - multi.add_host(config).await.unwrap(); - - let first_ping = first_peer.bound().await; - let abandoned = first_peer.request().await; - assert_eq!(abandoned.method, "initialize"); - assert_eq!(abandoned.id, 1); - drop(first_peer.tx); - assert!(bounded(first_peer.rx.recv()).await.is_none()); - assert!(matches!( - first_ping.ping().await, - Err(ClientError::Shutdown) - )); - - let second_ping = second_peer.bound().await; - let retry = second_peer.request().await; - assert_eq!(retry.method, "initialize"); - assert_eq!(retry.id, 2); - assert_eq!( - retry.params.as_ref().unwrap()["clientId"], - abandoned.params.as_ref().unwrap()["clientId"] - ); - let init = - json!({"protocolVersion": ahp_types::PROTOCOL_VERSION, "serverSeq": 10, "snapshots": []}); - second_peer.reply(abandoned.id, init.clone()).await; - let heartbeat = second_ping.ping(); - tokio::pin!(heartbeat); - let barrier = tokio::select! { - result = &mut heartbeat => panic!("premature result: {result:?}"), - sent = second_peer.request() => sent, - }; - assert_eq!( - barrier.method, "ping", - "stale handshake must not start discovery" - ); - assert_eq!(barrier.id, 3); - second_peer.reply(barrier.id, Value::Null).await; - bounded(heartbeat).await.unwrap(); - assert!(second_peer.rx.try_recv().is_err()); - assert!(multi.client(&HostId::from("retained")).await.is_none()); - - second_peer.reply(retry.id, init).await; - let discovery = second_peer.request().await; - assert_eq!(discovery.method, "listSessions"); - assert_eq!(discovery.id, 4); - second_peer.reply(discovery.id, json!({"items": []})).await; - loop { - if matches!( - bounded(events.recv()).await, - Some(HostEvent::Connected { .. }) - ) { - break; - } - } - let unanswered = second_ping.ping(); - tokio::pin!(unanswered); - let old_ping = tokio::select! { - result = &mut unanswered => panic!("premature result: {result:?}"), - sent = second_peer.request() => sent, - }; - assert_eq!(old_ping.id, 5); - drop(second_peer.tx); - assert!(matches!( - bounded(unanswered).await, - Err(ClientError::Shutdown) - )); - assert!(bounded(second_peer.rx.recv()).await.is_none()); - - let third_ping = third_peer.bound().await; - let reconnect = third_peer.request().await; - assert_eq!(reconnect.method, "reconnect"); - assert_eq!(reconnect.id, 6); - assert_eq!( - reconnect.params.as_ref().unwrap()["clientId"], - retry.params.as_ref().unwrap()["clientId"] - ); - third_peer.reply(old_ping.id, Value::Null).await; - let heartbeat = third_ping.ping(); - tokio::pin!(heartbeat); - let barrier = tokio::select! { - result = &mut heartbeat => panic!("premature result: {result:?}"), - sent = third_peer.request() => sent, - }; - assert_eq!( - barrier.method, "ping", - "stale ping must not satisfy reconnect" - ); - assert_eq!(barrier.id, 7); - third_peer.reply(barrier.id, Value::Null).await; - bounded(heartbeat).await.unwrap(); - assert!(third_peer.rx.try_recv().is_err()); - third_peer - .reply( - reconnect.id, - json!({"type": "replay", "actions": [], "missing": []}), - ) - .await; - let discovery = third_peer.request().await; - assert_eq!(discovery.method, "listSessions"); - assert_eq!(discovery.id, 8); - third_peer.reply(discovery.id, json!({"items": []})).await; - loop { - if matches!( - bounded(events.recv()).await, - Some(HostEvent::Connected { .. }) - ) { - break; - } - } - assert_eq!(attempts.load(Ordering::SeqCst), 4); - for count in [ - &first_peer.bind_count, - &second_peer.bind_count, - &third_peer.bind_count, - ] { - assert_eq!(count.load(Ordering::SeqCst), 1); - } - bounded(multi.remove_host(&HostId::from("retained"))) - .await - .unwrap(); - assert!(bounded(third_peer.rx.recv()).await.is_none()); - assert!(matches!( - third_ping.ping().await, - Err(ClientError::Shutdown) - )); -} - -#[tokio::test] -async fn independent_managed_hosts_start_independent_id_sequences() { - let multi = MultiHostClient::new(); - for id in ["first", "second"] { - let (transport, mut peer) = pair(); - let transport = Arc::new(Mutex::new(Some(transport))); - let config = HostConfig::new(id, id, move |_| { - let transport = transport.clone(); - async move { Ok(BoxedTransport::new(transport.lock().await.take().unwrap())) } - }) - .with_client_config(no_timeout()) - .with_reconnect_policy(ReconnectPolicy::disabled()); - multi.add_host(config).await.unwrap(); - assert_eq!(peer.request().await.id, 1); - bounded(multi.remove_host(&HostId::from(id))).await.unwrap(); - assert!(bounded(peer.rx.recv()).await.is_none()); - } -} diff --git a/docs/.changes/20261001-rust-connection-lifecycle.json b/docs/.changes/20261001-rust-connection-lifecycle.json new file mode 100644 index 000000000..6b8a3ee5c --- /dev/null +++ b/docs/.changes/20261001-rust-connection-lifecycle.json @@ -0,0 +1,5 @@ +{ + "type": "added", + "message": "Rust clients provide configurable automatic idle keepalive; managed hosts become ready after the handshake and refresh sessions concurrently without losing live updates, while preserving request IDs across reconnect attempts.", + "targets": ["rust"] +} diff --git a/docs/.changes/20261001-rust-weak-ping-binding.json b/docs/.changes/20261001-rust-weak-ping-binding.json deleted file mode 100644 index 156373912..000000000 --- a/docs/.changes/20261001-rust-weak-ping-binding.json +++ /dev/null @@ -1,5 +0,0 @@ -{ - "type": "added", - "message": "Rust transport client binding supplies a non-owning `WeakPingHandle` for client-correlated keepalive during discovery; managed hosts retain request IDs across reconnect attempts, with deterministic weak-ping shutdown and cancelled-request cleanup.", - "targets": ["rust"] -}