diff --git a/clients/rust/MULTI_HOST.md b/clients/rust/MULTI_HOST.md index bbe13617..b4e9a83a 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 c4f31084..5f87d70e 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 34302790..41c54c84 100644 --- a/clients/rust/crates/ahp/README.md +++ b/clients/rust/crates/ahp/README.md @@ -65,6 +65,66 @@ impl Transport for MyTransport { See `tests/client_roundtrip.rs` for a complete in-memory example. +### 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 +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 4b9e5f05..ace1eaf1 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,11 @@ 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 tokio::time::Instant; -use crate::error::ClientError; +use crate::error::{ClientError, TransportError}; use crate::transport::{Transport, TransportMessage}; /// Default size of a per-subscription broadcast channel. Consumers that @@ -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,9 +229,36 @@ pub struct DispatchHandle { // ─── Internal plumbing ─────────────────────────────────────────────────────── type PendingMap = HashMap>>; +type PreparedRequest<'a> = ( + JsonRpcMessage, + PendingRequest<'a>, + oneshot::Receiver>, +); + +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 +269,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 @@ -211,7 +279,107 @@ struct Shared { enum Outbound { Message(JsonRpcMessage), - 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, + })); + } + } + } + + 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 req = JsonRpcMessage::Request(JsonRpcRequest { + jsonrpc: JsonRpcVersion::V2, + id, + method: method.into(), + 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 *self.closed.borrow() { + return Err(ClientError::Shutdown); + } + pending.insert(id, tx); + } + Ok((req, PendingRequest { shared: self, id }, rx)) + } + + 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), + } + } + + async fn ping(&self) -> Result<(), ClientError> { + self.request( + "ping", + PingParams { + channel: ROOT_RESOURCE_URI, + }, + ) + .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); + } +} + +#[derive(Serialize)] +struct PingParams { + channel: &'static str, } // ─── Server-initiated request handling ─────────────────────────────────────── @@ -354,10 +522,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(); @@ -373,41 +545,58 @@ impl Client { 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( + 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); 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), }); 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 requests retain the `-32000` [`ClientError::Rpc`] + /// shutdown error. Automatic keepalive ends with the transport driver. pub async fn shutdown(&self) { - 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, - })); - } + self.shared.stop_requests(Some("client shut down")); } /// Send a JSON-RPC request and await its result. @@ -416,53 +605,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).await } /// Send a JSON-RPC notification (fire-and-forget). @@ -539,17 +682,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().await } /// Subscribe to a URI and obtain a handle that streams @@ -842,27 +975,45 @@ async fn drive_transport( shared: Arc, 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"), @@ -875,20 +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, + }, } } - // 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, - })); - } + 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 @@ -896,17 +1089,71 @@ 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); + +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 +1267,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 00000000..c9956ffd --- /dev/null +++ b/clients/rust/crates/ahp/src/client_tests.rs @@ -0,0 +1,271 @@ +#![allow(clippy::panic, clippy::unwrap_used)] + +use super::*; +use crate::BoxedTransport; + +struct TestTransport { + sent: mpsc::Sender, + received: mpsc::Receiver, +} + +impl Transport for TestTransport { + 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( + timeout: Option, +) -> ( + Client, + mpsc::Receiver, + mpsc::Sender, +) { + let (sent, rx) = mpsc::channel(1); + let (tx, received) = mpsc::channel(1); + 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(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!(sent.recv().await.is_none()); + + 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!(sent.recv().await.is_none()); +} + +#[tokio::test] +async fn request_id_exhaustion_never_wraps_or_enqueues_another_request() { + let (client, mut sent, _received) = client(Some(Duration::ZERO)).await; + *client.shared.request_ids.next.lock().unwrap() = Some(u64::MAX); + 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 _ in 0..2 { + assert!( + 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()); + } + client.shutdown().await; + assert!(matches!(client.ping().await, Err(ClientError::Shutdown))); + drop(client); + 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(); + 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 { + 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(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, mut transport_dropped) = oneshot::channel(); + let client = Client::connect( + BlockedTransport { + entered: Some(entered), + dropped: Some(dropped), + }, + ClientConfig::default(), + ) + .await + .unwrap(); + tokio::time::advance(Duration::from_secs(25)).await; + client.notify("block", ()).await.unwrap(); + send_started.await.unwrap(); + tokio::time::advance(Duration::from_secs(65)).await; + for _ in 0..10 { + tokio::task::yield_now().await; + } + 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 976d0399..d665309e 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 @@ -310,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; @@ -353,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 }; @@ -448,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, @@ -476,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(); + }, } } } @@ -740,9 +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 ec47aaa1..fde405c2 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 01322940..04e8194f 100644 --- a/clients/rust/crates/ahp/src/lib.rs +++ b/clients/rust/crates/ahp/src/lib.rs @@ -138,8 +138,11 @@ //! //! 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. Automatic idle +//! keepalive is owned by the same driver and ends with it. #![forbid(unsafe_code)] #![warn(missing_docs)] @@ -155,8 +158,9 @@ pub mod transport; pub use ahp_types; pub use client::{ - Client, ClientConfig, ClientEvent, ClientEventStream, DispatchHandle, ResourceRequestHandlers, - ServerRequestFuture, ServerRequestHandler, SessionSubscription, SubscriptionEvent, + 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/tests/connection_lifecycle.rs b/clients/rust/crates/ahp/tests/connection_lifecycle.rs new file mode 100644 index 00000000..c939737f --- /dev/null +++ b/clients/rust/crates/ahp/tests/connection_lifecycle.rs @@ -0,0 +1,694 @@ +#![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}, + "chats":[{"resource":"ahp-chat:/latest", "title":"Latest", "status":97}]}, "channel":"ahp-root://"}), + ), + ( + "root/sessionSummaryChanged", + json!({"session":"changed", "changes":{"activity":"busy", "defaultChat":"ahp-chat:/latest"}, "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)); + let chat = &changed.chats.as_ref().unwrap()[0]; + assert_eq!(chat.resource, "ahp-chat:/latest"); + assert_eq!(chat.status, Some(97)); + assert_eq!(changed.default_chat.as_deref(), Some("ahp-chat:/latest")); + 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/docs/.changes/20261001-rust-connection-lifecycle.json b/docs/.changes/20261001-rust-connection-lifecycle.json new file mode 100644 index 00000000..6b8a3ee5 --- /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"] +}