From 3e132c33d1e32d34279168ad23145433e7025b7e Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 25 Aug 2026 14:28:37 +0100 Subject: [PATCH 1/9] feat(acp): add unstable v2 session injection --- Cargo.lock | 3 +- Cargo.toml | 2 +- src/agent-client-protocol/CHANGELOG.md | 3 + src/agent-client-protocol/Cargo.toml | 4 + .../src/schema/v2_impls.rs | 30 +++++ src/agent-client-protocol/src/session/v2.rs | 38 ++++++ .../tests/schema_session_inject.rs | 84 +++++++++++++ src/agent-client-protocol/tests/session_v2.rs | 112 ++++++++++++++++++ 8 files changed, 273 insertions(+), 3 deletions(-) create mode 100644 src/agent-client-protocol/tests/schema_session_inject.rs diff --git a/Cargo.lock b/Cargo.lock index 74c2bcbb..897ca11f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -140,8 +140,7 @@ dependencies = [ [[package]] name = "agent-client-protocol-schema" version = "1.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ca98360c7bb8cc97d7acd49e2a8a851c3f7bee6b2f0535036d8ab86b5fcd223d" +source = "git+https://github.com/danielkov/agent-client-protocol?rev=af121986c3a7e6a1fd5176485d7c809b5654c088#af121986c3a7e6a1fd5176485d7c809b5654c088" dependencies = [ "anyhow", "derive_more", diff --git a/Cargo.toml b/Cargo.toml index c5b8a21e..24f04bfa 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -38,7 +38,7 @@ agent-client-protocol-trace-viewer = { path = "src/agent-client-protocol-trace-v yopo = { package = "agent-client-protocol-yopo", path = "src/yopo" } # Protocol -agent-client-protocol-schema = { version = "=1.7.0", features = ["tracing"] } +agent-client-protocol-schema = { git = "https://github.com/danielkov/agent-client-protocol", rev = "af121986c3a7e6a1fd5176485d7c809b5654c088", features = ["tracing"] } # Core async runtime tokio = { version = "1.52", default-features = false } diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index 4003afe4..a49de91a 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -4,6 +4,9 @@ ### Added +- *(unstable-v2)* Expose pending session injection through the + `unstable_session_inject` feature, including typed JSON-RPC dispatch and + `V2Session` helpers to inject, replace, and revoke typed content. - *(unstable-v2)* Add runnable draft-v2 agent and one-shot client examples. The agent implements the complete baseline session lifecycle; the client handles permissions, projects chunk and snapshot updates by message ID, and waits for diff --git a/src/agent-client-protocol/Cargo.toml b/src/agent-client-protocol/Cargo.toml index 1eb857aa..ffbf4903 100644 --- a/src/agent-client-protocol/Cargo.toml +++ b/src/agent-client-protocol/Cargo.toml @@ -43,6 +43,10 @@ unstable_mcp_over_acp = ["agent-client-protocol-schema/unstable_mcp_over_acp"] unstable_plan_operations = ["agent-client-protocol-schema/unstable_plan_operations"] unstable_session_compaction = ["agent-client-protocol-schema/unstable_session_compaction"] unstable_session_fork = ["agent-client-protocol-schema/unstable_session_fork"] +unstable_session_inject = [ + "unstable_protocol_v2", + "agent-client-protocol-schema/unstable_session_inject", +] unstable_tool_call_name = ["agent-client-protocol-schema/unstable_tool_call_name"] unstable_protocol_v2 = ["agent-client-protocol-schema/unstable_protocol_v2"] diff --git a/src/agent-client-protocol/src/schema/v2_impls.rs b/src/agent-client-protocol/src/schema/v2_impls.rs index 8cae486b..bc6fb38e 100644 --- a/src/agent-client-protocol/src/schema/v2_impls.rs +++ b/src/agent-client-protocol/src/schema/v2_impls.rs @@ -263,6 +263,24 @@ impl_v2_jsonrpc_request!( "session/set_config_option" ); impl_v2_jsonrpc_request!(v2::PromptRequest, v2::PromptResponse, "session/prompt"); +#[cfg(feature = "unstable_session_inject")] +impl_v2_jsonrpc_request!( + v2::InjectSessionRequest, + v2::InjectSessionResponse, + "session/inject" +); +#[cfg(feature = "unstable_session_inject")] +impl_v2_jsonrpc_request!( + v2::RevokeInjectSessionRequest, + v2::RevokeInjectSessionResponse, + "session/revoke_inject" +); +#[cfg(feature = "unstable_session_inject")] +impl_v2_jsonrpc_request!( + v2::ReplaceInjectSessionRequest, + v2::ReplaceInjectSessionResponse, + "session/replace_inject" +); #[cfg(feature = "unstable_mcp_over_acp")] impl_v2_jsonrpc_request!(v2::MessageMcpRequest, v2::MessageMcpResponse, "mcp/message"); @@ -316,6 +334,12 @@ impl_v2_jsonrpc_request_enum!(v2::ClientRequest { CloseSessionRequest => "session/close", SetSessionConfigOptionRequest => "session/set_config_option", PromptRequest => "session/prompt", + #[cfg(feature = "unstable_session_inject")] + InjectSessionRequest => "session/inject", + #[cfg(feature = "unstable_session_inject")] + RevokeInjectSessionRequest => "session/revoke_inject", + #[cfg(feature = "unstable_session_inject")] + ReplaceInjectSessionRequest => "session/replace_inject", #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpRequest => "mcp/message", [ext] ExtMethodRequest, @@ -340,6 +364,12 @@ impl_v2_jsonrpc_response_enum!(v2::AgentResponse { CloseSessionResponse => "session/close", SetSessionConfigOptionResponse => "session/set_config_option", PromptResponse => "session/prompt", + #[cfg(feature = "unstable_session_inject")] + InjectSessionResponse => "session/inject", + #[cfg(feature = "unstable_session_inject")] + RevokeInjectSessionResponse => "session/revoke_inject", + #[cfg(feature = "unstable_session_inject")] + ReplaceInjectSessionResponse => "session/replace_inject", #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpResponse => "mcp/message", [ext] ExtMethodResponse, diff --git a/src/agent-client-protocol/src/session/v2.rs b/src/agent-client-protocol/src/session/v2.rs index 4d78c2fc..eba50c22 100644 --- a/src/agent-client-protocol/src/session/v2.rs +++ b/src/agent-client-protocol/src/session/v2.rs @@ -809,6 +809,44 @@ where ) } + /// Inject content for pending delivery to this session. + #[cfg(feature = "unstable_session_inject")] + pub fn inject( + &self, + mode: v2::SessionInjectMode, + content: Vec, + ) -> SentRequest { + self.connection.send_request_to( + Agent, + v2::InjectSessionRequest::new(self.session_id.clone(), mode, content), + ) + } + + /// Revoke a pending injected message. + #[cfg(feature = "unstable_session_inject")] + pub fn revoke_inject( + &self, + message_id: impl Into, + ) -> SentRequest { + self.connection.send_request_to( + Agent, + v2::RevokeInjectSessionRequest::new(self.session_id.clone(), message_id), + ) + } + + /// Replace the content of a pending injected message. + #[cfg(feature = "unstable_session_inject")] + pub fn replace_inject( + &self, + message_id: impl Into, + content: Vec, + ) -> SentRequest { + self.connection.send_request_to( + Agent, + v2::ReplaceInjectSessionRequest::new(self.session_id.clone(), message_id, content), + ) + } + /// Ask the agent to cancel the session's current foreground work. /// /// This is independent from cancelling a prompt's [`SentRequest`]. diff --git a/src/agent-client-protocol/tests/schema_session_inject.rs b/src/agent-client-protocol/tests/schema_session_inject.rs new file mode 100644 index 00000000..2f2cac41 --- /dev/null +++ b/src/agent-client-protocol/tests/schema_session_inject.rs @@ -0,0 +1,84 @@ +#![cfg(feature = "unstable_session_inject")] + +use agent_client_protocol::{JsonRpcMessage, JsonRpcResponse, schema::v2}; +use serde_json::json; + +#[test] +fn v2_session_inject_requests_serialize_and_dispatch() { + let content = vec![v2::ContentBlock::Text(v2::TextContent::new("steer now"))]; + let inject = + v2::InjectSessionRequest::new("session-1", v2::SessionInjectMode::Steer, content.clone()); + let untyped = inject.to_untyped_message().unwrap(); + assert_eq!(untyped.method, "session/inject"); + assert_eq!( + untyped.params, + json!({ + "sessionId": "session-1", + "mode": "steer", + "content": [{ "type": "text", "text": "steer now" }] + }) + ); + assert!(matches!( + v2::ClientRequest::parse_message(untyped.method(), untyped.params()).unwrap(), + v2::ClientRequest::InjectSessionRequest(_) + )); + assert!(matches!( + v2::AgentResponse::from_value( + "session/inject", + json!({ "messageId": "message-1" }), + ) + .unwrap(), + v2::AgentResponse::InjectSessionResponse(response) + if response.message_id == v2::MessageId::new("message-1") + )); + + let revoke = v2::RevokeInjectSessionRequest::new("session-1", v2::MessageId::new("message-1")); + assert_eq!( + revoke.to_untyped_message().unwrap().params, + json!({ "sessionId": "session-1", "messageId": "message-1" }) + ); + assert!(matches!( + v2::ClientRequest::parse_message( + "session/revoke_inject", + &json!({ "sessionId": "session-1", "messageId": "message-1" }), + ) + .unwrap(), + v2::ClientRequest::RevokeInjectSessionRequest(_) + )); + assert!(matches!( + v2::AgentResponse::from_value("session/revoke_inject", json!({})).unwrap(), + v2::AgentResponse::RevokeInjectSessionResponse(_) + )); + + let replace = + v2::ReplaceInjectSessionRequest::new("session-1", v2::MessageId::new("message-1"), content); + assert_eq!( + replace.to_untyped_message().unwrap().params, + json!({ + "sessionId": "session-1", + "messageId": "message-1", + "content": [{ "type": "text", "text": "steer now" }] + }) + ); + assert!(matches!( + v2::ClientRequest::parse_message( + "session/replace_inject", + &json!({ + "sessionId": "session-1", + "messageId": "message-1", + "content": [{ "type": "text", "text": "steer now" }] + }), + ) + .unwrap(), + v2::ClientRequest::ReplaceInjectSessionRequest(_) + )); + assert!(matches!( + v2::AgentResponse::from_value( + "session/replace_inject", + json!({ "messageId": "message-1" }), + ) + .unwrap(), + v2::AgentResponse::ReplaceInjectSessionResponse(response) + if response.message_id == v2::MessageId::new("message-1") + )); +} diff --git a/src/agent-client-protocol/tests/session_v2.rs b/src/agent-client-protocol/tests/session_v2.rs index 87943dad..e828f12a 100644 --- a/src/agent-client-protocol/tests/session_v2.rs +++ b/src/agent-client-protocol/tests/session_v2.rs @@ -1004,6 +1004,118 @@ async fn v2_session_commands_cover_configuration_and_close() { .expect("v2 session command test failed"); } +#[cfg(feature = "unstable_session_inject")] +#[tokio::test(flavor = "current_thread")] +async fn v2_session_inject_helpers_preserve_typed_content_and_message_ids() { + let session_id = v2::SessionId::new("inject-session"); + let agent_session_id = session_id.clone(); + + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _connection: V2ConnectionTo| { + responder.respond(initialize_response(request.protocol_version)) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |_request: v2::NewSessionRequest, + responder: Responder, + _connection: V2ConnectionTo| { + responder.respond(v2::NewSessionResponse::new(agent_session_id.clone())) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::InjectSessionRequest, + responder: Responder, + _connection: V2ConnectionTo| { + assert_eq!(request.session_id, v2::SessionId::new("inject-session")); + assert_eq!(request.mode, v2::SessionInjectMode::Steer); + assert_eq!( + request.content, + vec![v2::ContentBlock::Text(v2::TextContent::new("steer"))] + ); + responder.respond(v2::InjectSessionResponse::new(v2::MessageId::new( + "message-1", + ))) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::RevokeInjectSessionRequest, + responder: Responder, + _connection: V2ConnectionTo| { + assert_eq!(request.session_id, v2::SessionId::new("inject-session")); + assert_eq!(request.message_id, v2::MessageId::new("message-1")); + responder.respond(v2::RevokeInjectSessionResponse::new()) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::ReplaceInjectSessionRequest, + responder: Responder, + _connection: V2ConnectionTo| { + assert_eq!(request.session_id, v2::SessionId::new("inject-session")); + assert_eq!(request.message_id, v2::MessageId::new("message-1")); + assert_eq!( + request.content, + vec![v2::ContentBlock::Text(v2::TextContent::new("replacement"))] + ); + responder.respond(v2::ReplaceInjectSessionResponse::new(request.message_id)) + }, + agent_client_protocol::on_receive_request!(), + ); + + let client = Client.v2().connect_with(agent, async move |connection| { + connection + .send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + implementation(), + )) + .block_task() + .await?; + + let opened = connection + .build_session(cwd()?) + .start_session() + .block_task() + .await?; + let (session, _) = opened.into_parts(); + + let injected = session + .inject( + v2::SessionInjectMode::Steer, + vec![v2::ContentBlock::Text(v2::TextContent::new("steer"))], + ) + .block_task() + .await?; + assert_eq!(injected.message_id, v2::MessageId::new("message-1")); + + let replaced = session + .replace_inject( + injected.message_id, + vec![v2::ContentBlock::Text(v2::TextContent::new("replacement"))], + ) + .block_task() + .await?; + assert_eq!(replaced.message_id, v2::MessageId::new("message-1")); + + session + .revoke_inject(replaced.message_id) + .block_task() + .await?; + Ok(()) + }); + + tokio::time::timeout(TIMEOUT, client) + .await + .expect("v2 session inject helper test timed out") + .expect("v2 session inject helper test failed"); +} + #[tokio::test(flavor = "current_thread")] async fn dropping_v2_session_does_not_unregister_update_handling() { let session_id = v2::SessionId::new("dropped-handle"); From c1ba83b2bffd47a6ce49aa87ea0dca7c3f83ec73 Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 25 Aug 2026 14:30:13 +0100 Subject: [PATCH 2/9] build(acp): pin reviewed session injection schema --- Cargo.lock | 2 +- Cargo.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 897ca11f..1df89af3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -140,7 +140,7 @@ dependencies = [ [[package]] name = "agent-client-protocol-schema" version = "1.7.0" -source = "git+https://github.com/danielkov/agent-client-protocol?rev=af121986c3a7e6a1fd5176485d7c809b5654c088#af121986c3a7e6a1fd5176485d7c809b5654c088" +source = "git+https://github.com/danielkov/agent-client-protocol?rev=6e7e044f9464c4fd652d90699a09e9edc8b3bbad#6e7e044f9464c4fd652d90699a09e9edc8b3bbad" dependencies = [ "anyhow", "derive_more", diff --git a/Cargo.toml b/Cargo.toml index 24f04bfa..21417a8c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -38,7 +38,7 @@ agent-client-protocol-trace-viewer = { path = "src/agent-client-protocol-trace-v yopo = { package = "agent-client-protocol-yopo", path = "src/yopo" } # Protocol -agent-client-protocol-schema = { git = "https://github.com/danielkov/agent-client-protocol", rev = "af121986c3a7e6a1fd5176485d7c809b5654c088", features = ["tracing"] } +agent-client-protocol-schema = { git = "https://github.com/danielkov/agent-client-protocol", rev = "6e7e044f9464c4fd652d90699a09e9edc8b3bbad", features = ["tracing"] } # Core async runtime tokio = { version = "1.52", default-features = false } From 2f039993d1d6ed8da35b38c31f54a7cbb7338c70 Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 25 Aug 2026 18:35:01 +0100 Subject: [PATCH 3/9] feat(acp): add tracked response receipts --- src/agent-client-protocol/src/jsonrpc.rs | 281 +++++++++++++++--- .../src/jsonrpc/incoming_actor.rs | 1 + .../src/jsonrpc/outgoing_actor.rs | 130 ++++++-- src/agent-client-protocol/src/lib.rs | 4 +- .../tests/jsonrpc_batch.rs | 269 ++++++++++++++++- 5 files changed, 622 insertions(+), 63 deletions(-) diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 9cff5a6d..5068cea4 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -2692,6 +2692,120 @@ pub fn is_cancel_request_notification(notification: &N) } } +/// Resolves when a tracked JSON-RPC response frame is accepted by the outgoing +/// transport-frame queue. +/// +/// This receipt does not indicate that any bytes were written to the transport. +/// For a response in a batch, it resolves only after every response in the batch +/// is ready and the aggregate batch frame is accepted by the queue. It resolves +/// with an error if enqueueing fails or the outgoing actor tears down first. +/// +/// # Batch handler deadlocks +/// +/// A batch cannot be enqueued until all of its handlers return. Do not await a +/// receipt from inside a batch handler. Return from the handler first, or spawn +/// receipt-dependent side effects so the handler can return immediately. +#[must_use = "a response receipt must be awaited to observe enqueue completion"] +pub struct ResponseReceipt { + receiver: oneshot::Receiver>, +} + +impl std::fmt::Debug for ResponseReceipt { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ResponseReceipt") + .finish_non_exhaustive() + } +} + +impl std::future::Future for ResponseReceipt { + type Output = Result<(), crate::Error>; + + fn poll( + mut self: std::pin::Pin<&mut Self>, + context: &mut std::task::Context<'_>, + ) -> std::task::Poll { + match std::pin::Pin::new(&mut self.receiver).poll(context) { + std::task::Poll::Ready(Ok(result)) => std::task::Poll::Ready(result), + std::task::Poll::Ready(Err(_)) => { + std::task::Poll::Ready(Err(response_receipt_teardown_error())) + } + std::task::Poll::Pending => std::task::Poll::Pending, + } + } +} + +struct ResponseReceiptSender { + state: Arc, +} + +impl std::fmt::Debug for ResponseReceiptSender { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ResponseReceiptSender") + .finish_non_exhaustive() + } +} + +struct ResponseReceiptState { + sender: Mutex>>>, +} + +impl ResponseReceiptSender { + fn channel() -> (Self, ResponseReceipt) { + let (sender, receiver) = oneshot::channel(); + ( + Self { + state: Arc::new(ResponseReceiptState { + sender: Mutex::new(Some(sender)), + }), + }, + ResponseReceipt { receiver }, + ) + } + + fn resolve(self, result: Result<(), crate::Error>) { + self.state.resolve(result); + } +} + +impl ResponseReceiptState { + fn resolve(&self, result: Result<(), crate::Error>) { + if let Some(sender) = self + .sender + .lock() + .expect("response receipt mutex poisoned") + .take() + { + drop(sender.send(result)); + } + } +} + +impl Drop for ResponseReceiptState { + fn drop(&mut self) { + if let Some(sender) = self + .sender + .get_mut() + .expect("response receipt mutex poisoned") + .take() + { + drop(sender.send(Err(response_receipt_teardown_error()))); + } + } +} + +fn response_receipt_teardown_error() -> crate::Error { + crate::util::internal_error( + "outgoing JSON-RPC actor stopped before the response frame was enqueued", + ) +} + +struct CompletedResponseFrame { + frame: TransportFrame, + receipts: Vec, +} + /// Messages send to be serialized over the transport. #[derive(Clone)] enum ResponseDestination { @@ -2719,6 +2833,7 @@ impl ResponseDestination { responses: (0..slot_count).map(|_| None).collect(), abandoned: (0..slot_count).map(|_| None).collect(), active_handler_attempts: (0..slot_count).map(|_| 0).collect(), + receipts: (0..slot_count).map(|_| None).collect(), dispatch_complete: false, emitted: false, })); @@ -2737,14 +2852,18 @@ impl ResponseDestination { ) } - fn complete(self, response: RawJsonRpcMessage) -> Option { + fn complete( + self, + response: RawJsonRpcMessage, + receipt: Option, + ) -> Option { match self { - Self::Individual(slot) => slot.complete(response), - Self::Batch(slot) => slot.complete(response).map(batch_response_frame), + Self::Individual(slot) => slot.complete(response, receipt), + Self::Batch(slot) => slot.complete(response, receipt).map(batch_response_frame), } } - fn abandon(self, fallback: RawJsonRpcMessage) -> Option { + fn abandon(self, fallback: RawJsonRpcMessage) -> Option { match self { Self::Individual(_) => None, Self::Batch(slot) => slot.abandon(fallback).map(batch_response_frame), @@ -2769,7 +2888,7 @@ impl ResponseDestination { }) } - fn finish_handler_attempt(self) -> Option { + fn finish_handler_attempt(self) -> Option { match self { Self::Individual(_) => None, Self::Batch(slot) => slot.finish_handler_attempt().map(batch_response_frame), @@ -2783,21 +2902,33 @@ struct IndividualResponseSlot { } impl IndividualResponseSlot { - fn complete(self, response: RawJsonRpcMessage) -> Option { + fn complete( + self, + response: RawJsonRpcMessage, + receipt: Option, + ) -> Option { if self.completed.swap(true, Ordering::AcqRel) { tracing::warn!("Ignoring duplicate completion of JSON-RPC request"); return None; } - Some(TransportFrame::Single(response)) + Some(CompletedResponseFrame { + frame: TransportFrame::Single(response), + receipts: receipt.into_iter().collect(), + }) } } -fn batch_response_frame(responses: Vec) -> TransportFrame { - TransportFrame::Batch( - TransportBatch::from_messages(responses) - .expect("a completed JSON-RPC response batch is non-empty"), - ) +fn batch_response_frame( + (responses, receipts): (Vec, Vec), +) -> CompletedResponseFrame { + CompletedResponseFrame { + frame: TransportFrame::Batch( + TransportBatch::from_messages(responses) + .expect("a completed JSON-RPC response batch is non-empty"), + ), + receipts, + } } #[derive(Clone)] @@ -2814,7 +2945,7 @@ impl std::fmt::Debug for BatchDispatchCompletion { } impl BatchDispatchCompletion { - fn complete(self) -> Option { + fn complete(self) -> Option { let mut state = self .state .lock() @@ -2841,23 +2972,25 @@ fn promote_abandoned_response(state: &mut BatchResponseState, index: usize) { } } -fn take_completed_batch(state: &mut BatchResponseState) -> Option> { +fn take_completed_batch( + state: &mut BatchResponseState, +) -> Option<(Vec, Vec)> { if !state.dispatch_complete || state.remaining != 0 || state.emitted { return None; } state.emitted = true; - Some( - state - .responses - .iter_mut() - .map(|response| { - response - .take() - .expect("completed JSON-RPC batch has every response slot") - }) - .collect(), - ) + let responses = state + .responses + .iter_mut() + .map(|response| { + response + .take() + .expect("completed JSON-RPC batch has every response slot") + }) + .collect(); + let receipts = state.receipts.iter_mut().filter_map(Option::take).collect(); + Some((responses, receipts)) } #[derive(Clone)] @@ -2884,7 +3017,9 @@ impl BatchResponseSlot { state.active_handler_attempts[self.index] += 1; } - fn finish_handler_attempt(self) -> Option> { + fn finish_handler_attempt( + self, + ) -> Option<(Vec, Vec)> { let mut state = self .state .lock() @@ -2898,7 +3033,11 @@ impl BatchResponseSlot { take_completed_batch(&mut state) } - fn complete(self, response: RawJsonRpcMessage) -> Option> { + fn complete( + self, + response: RawJsonRpcMessage, + receipt: Option, + ) -> Option<(Vec, Vec)> { let mut state = self .state .lock() @@ -2924,11 +3063,15 @@ impl BatchResponseSlot { state.abandoned[self.index] = None; state.responses[self.index] = Some(response); + state.receipts[self.index] = receipt; state.remaining -= 1; take_completed_batch(&mut state) } - fn abandon(self, fallback: RawJsonRpcMessage) -> Option> { + fn abandon( + self, + fallback: RawJsonRpcMessage, + ) -> Option<(Vec, Vec)> { let mut state = self .state .lock() @@ -2959,6 +3102,7 @@ struct BatchResponseState { responses: Vec>, abandoned: Vec>, active_handler_attempts: Vec, + receipts: Vec>, dispatch_complete: bool, emitted: bool, } @@ -3138,6 +3282,8 @@ enum OutgoingMessage { response: Result, destination: ResponseDestination, + + receipt: Option, }, /// Send an Error Response that cannot be correlated to a request ID. @@ -4452,7 +4598,13 @@ pub struct Responder { /// /// For incoming requests: serializes to JSON and sends over the wire. /// For incoming responses: sends to the waiting oneshot channel. - send_fn: Box) -> Result<(), crate::Error> + Send>, + send_fn: Box< + dyn FnOnce( + Result, + Option, + ) -> Result<(), crate::Error> + + Send, + >, /// Completes an abandoned batch slot unless an explicit response disarms it. drop_guard: ResponderDropGuard, @@ -4533,17 +4685,20 @@ impl Responder { id, cancellation, destination, - send_fn: Box::new(move |response: Result| { - send_raw_message( - &message_tx, - OutgoingMessage::Response { - id: id_clone, - method: method_clone, - response, - destination: send_destination, - }, - ) - }), + send_fn: Box::new( + move |response: Result, receipt| { + send_raw_message( + &message_tx, + OutgoingMessage::Response { + id: id_clone, + method: method_clone, + response, + destination: send_destination, + receipt, + }, + ) + }, + ), drop_guard, } } @@ -4619,9 +4774,9 @@ impl Responder { id: self.id, cancellation: self.cancellation, destination: self.destination, - send_fn: Box::new(move |input: Result| { + send_fn: Box::new(move |input: Result, receipt| { let t_value = wrap_fn(&method, input); - (self.send_fn)(t_value) + (self.send_fn)(t_value, receipt) }), drop_guard: self.drop_guard, } @@ -4634,7 +4789,47 @@ impl Responder { ) -> Result<(), crate::Error> { tracing::debug!(id = ?self.id, "respond called"); self.drop_guard.disarm(); - (self.send_fn)(response) + (self.send_fn)(response, None) + } + + /// Respond to the JSON-RPC request with either a value (`Ok`) or an error (`Err`) + /// and return a receipt for enqueue completion. + /// + /// The receipt resolves successfully when the response's single + /// [`TransportFrame`], or its aggregate batch frame, is accepted by the + /// outgoing transport-frame queue. It does not wait for bytes to be written. + /// It resolves with an error if enqueueing fails or the outgoing actor tears + /// down first. + /// + /// # Errors + /// + /// Returns an error immediately if the response cannot enter the outgoing + /// protocol queue. After this method returns a receipt, enqueue or actor + /// teardown failures are reported by awaiting that receipt. + /// + /// # Batch handler deadlocks + /// + /// A batch frame cannot be enqueued until all batch handlers return. Do not + /// await the receipt inside a batch handler. Return first, or spawn any side + /// effect that awaits the receipt so the handler can return immediately. + pub fn respond_with_result_tracked( + mut self, + response: Result, + ) -> Result { + tracing::debug!(id = ?self.id, "tracked respond called"); + let (sender, receipt) = ResponseReceiptSender::channel(); + self.drop_guard.disarm(); + (self.send_fn)(response, Some(sender))?; + Ok(receipt) + } + + /// Respond to the JSON-RPC request with a value and return a receipt for + /// enqueue completion. + /// + /// See [`respond_with_result_tracked`](Self::respond_with_result_tracked) for + /// receipt semantics and the batch-handler deadlock warning. + pub fn respond_tracked(self, response: T) -> Result { + self.respond_with_result_tracked(Ok(response)) } /// Respond to the JSON-RPC request with a value. diff --git a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs index b65fccab..a4ab7a3d 100644 --- a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs @@ -647,6 +647,7 @@ fn handle_handler_error( method: reply_target.method, response: Err(error), destination: reply_target.destination, + receipt: None, }, ), Some(HandlerErrorTarget::Response(reply_target)) => { diff --git a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs index effe164c..0d115c36 100644 --- a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs @@ -1,9 +1,14 @@ // Types re-exported from crate root +use std::sync::{Arc, Weak}; + use futures::StreamExt as _; use futures::channel::mpsc; use crate::jsonrpc::protocol_compat::ProtocolCompat; -use crate::jsonrpc::{OutgoingMessage, PendingReplies, RawJsonRpcMessage, TransportFrame}; +use crate::jsonrpc::{ + CompletedResponseFrame, OutgoingMessage, PendingReplies, RawJsonRpcMessage, + ResponseReceiptSender, ResponseReceiptState, TransportFrame, response_receipt_teardown_error, +}; use crate::schema::v1::RequestId; pub type OutgoingMessageTx = mpsc::UnboundedSender; @@ -17,6 +22,47 @@ pub(crate) fn send_raw_message( .map_err(crate::util::internal_error) } +#[derive(Default)] +struct ResponseReceiptRegistry { + pending: Vec>, +} + +impl ResponseReceiptRegistry { + fn register(&mut self, sender: &ResponseReceiptSender) { + self.pending.retain(|state| state.strong_count() != 0); + self.pending.push(Arc::downgrade(&sender.state)); + } +} + +impl Drop for ResponseReceiptRegistry { + fn drop(&mut self) { + for state in self.pending.drain(..).filter_map(|state| state.upgrade()) { + state.resolve(Err(response_receipt_teardown_error())); + } + } +} + +fn enqueue_completed_response( + transport_tx: &mpsc::UnboundedSender, + completed: CompletedResponseFrame, +) -> Result<(), crate::Error> { + match transport_tx.unbounded_send(completed.frame) { + Ok(()) => { + for receipt in completed.receipts { + receipt.resolve(Ok(())); + } + Ok(()) + } + Err(error) => { + let error = crate::Error::into_internal_error(error); + for receipt in completed.receipts { + receipt.resolve(Err(error.clone())); + } + Err(error) + } + } +} + /// Outgoing protocol actor: Converts application-level OutgoingMessage to protocol-level RawJsonRpcMessage. /// /// This actor handles JSON-RPC protocol semantics: @@ -31,12 +77,13 @@ pub(super) async fn outgoing_protocol_actor( protocol_compat: ProtocolCompat, ) -> Result<(), crate::Error> { let mut drain_waiters = Vec::new(); + let mut receipt_registry = ResponseReceiptRegistry::default(); while let Some(message) = outgoing_rx.next().await { tracing::debug!(?message, "outgoing_protocol_actor"); // Create the message to be sent over the transport - let (json_rpc_message, destination) = match message { + let (json_rpc_message, destination, receipt) = match message { OutgoingMessage::CloseAfterDraining { done } => { // Reject later sends while preserving every message that was // already accepted into this receiver's buffer. @@ -46,17 +93,13 @@ pub(super) async fn outgoing_protocol_actor( } OutgoingMessage::BatchDispatchComplete { completion } => { if let Some(frame) = completion.complete() { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + enqueue_completed_response(&transport_tx, frame)?; } continue; } OutgoingMessage::BatchHandlerAttemptComplete { destination } => { if let Some(frame) = destination.finish_handler_attempt() { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + enqueue_completed_response(&transport_tx, frame)?; } continue; } @@ -79,9 +122,7 @@ pub(super) async fn outgoing_protocol_actor( ); let fallback = RawJsonRpcMessage::response(id, fallback); if let Some(frame) = destination.abandon(fallback) { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + enqueue_completed_response(&transport_tx, frame)?; } continue; } @@ -180,14 +221,23 @@ pub(super) async fn outgoing_protocol_actor( method, response, destination, + receipt, } => match protocol_compat.outgoing_response_to(&id, &method, response) { Ok(value) => { tracing::debug!(?id, "Sending success response"); - (RawJsonRpcMessage::response(id, Ok(value)), destination) + ( + RawJsonRpcMessage::response(id, Ok(value)), + destination, + receipt, + ) } Err(error) => { tracing::warn!(?id, %method, ?error, "Sending error response"); - (RawJsonRpcMessage::response(id, Err(error)), destination) + ( + RawJsonRpcMessage::response(id, Err(error)), + destination, + receipt, + ) } }, OutgoingMessage::UncorrelatedErrorResponse { error, destination } => { @@ -196,14 +246,16 @@ pub(super) async fn outgoing_protocol_actor( ( RawJsonRpcMessage::response(RequestId::Null, Err(error)), destination, + None, ) } }; - if let Some(frame) = destination.complete(json_rpc_message) { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some(receipt) = receipt.as_ref() { + receipt_registry.register(receipt); + } + if let Some(frame) = destination.complete(json_rpc_message, receipt) { + enqueue_completed_response(&transport_tx, frame)?; } } @@ -216,3 +268,47 @@ pub(super) async fn outgoing_protocol_actor( } Ok(()) } + +#[cfg(test)] +mod tests { + use futures::executor::block_on; + use futures::future::join; + + use super::*; + + #[test] + fn actor_teardown_fails_every_registered_response_receipt() { + let (first_sender, first_receipt) = ResponseReceiptSender::channel(); + let (second_sender, second_receipt) = ResponseReceiptSender::channel(); + let mut registry = ResponseReceiptRegistry::default(); + registry.register(&first_sender); + registry.register(&second_sender); + + drop(registry); + + let (first_result, second_result) = block_on(join(first_receipt, second_receipt)); + assert!(first_result.is_err()); + assert!(second_result.is_err()); + drop((first_sender, second_sender)); + } + + #[test] + fn response_transport_queue_failure_fails_every_receipt() { + let (transport_tx, transport_rx) = mpsc::unbounded(); + drop(transport_rx); + let (first_sender, first_receipt) = ResponseReceiptSender::channel(); + let (second_sender, second_receipt) = ResponseReceiptSender::channel(); + let completed = CompletedResponseFrame { + frame: TransportFrame::Single(RawJsonRpcMessage::response( + RequestId::Null, + Ok(serde_json::Value::Null), + )), + receipts: vec![first_sender, second_sender], + }; + + assert!(enqueue_completed_response(&transport_tx, completed).is_err()); + let (first_result, second_result) = block_on(join(first_receipt, second_receipt)); + assert!(first_result.is_err()); + assert!(second_result.is_err()); + } +} diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index 4f71dbde..b94cddf8 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -125,8 +125,8 @@ pub use jsonrpc::{ HandleConnectionClose, HandleDispatchFrom, Handled, INCOMING_TRANSPORT_CLOSED_REASON, IntoHandled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, Lines, NullClose, NullHandler, RawConnectionContext, RawJsonRpcMessage, RawJsonRpcParams, Responder, - ResponseRouter, SentRequest, TransportBatch, TransportBatchEntry, TransportFrame, - UntypedMessage, is_incoming_transport_closed, + ResponseReceipt, ResponseRouter, SentRequest, TransportBatch, TransportBatchEntry, + TransportFrame, UntypedMessage, is_incoming_transport_closed, run::{ChainRun, NullRun, RunWithConnectionTo}, }; pub use jsonrpc::{RequestCancellation, is_cancel_request_notification}; diff --git a/src/agent-client-protocol/tests/jsonrpc_batch.rs b/src/agent-client-protocol/tests/jsonrpc_batch.rs index 97df192e..ecb812b9 100644 --- a/src/agent-client-protocol/tests/jsonrpc_batch.rs +++ b/src/agent-client-protocol/tests/jsonrpc_batch.rs @@ -17,7 +17,8 @@ use std::{ use agent_client_protocol::{ Agent, ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, HandleDispatchFrom, Handled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, - RawJsonRpcMessage, Responder, TransportBatch, TransportBatchEntry, TransportFrame, + RawJsonRpcMessage, Responder, ResponseReceipt, TransportBatch, TransportBatchEntry, + TransportFrame, role::{Role, UntypedRole}, schema::ProtocolVersion, schema::v1, @@ -237,6 +238,79 @@ async fn next_deferred_response( .expect("deferred responder channel closed unexpectedly") } +async fn next_response_receipt( + rx: &mut mpsc::UnboundedReceiver<(String, ResponseReceipt)>, +) -> (String, ResponseReceipt) { + tokio::time::timeout(TIMEOUT, rx.next()) + .await + .expect("timed out waiting for response receipt") + .expect("response receipt channel closed unexpectedly") +} + +fn start_tracked_server( + responder_tx: mpsc::UnboundedSender<(String, Responder)>, + receipt_tx: mpsc::UnboundedSender<(String, ResponseReceipt)>, +) -> ( + DuplexStream, + BufReader, + JoinHandle>, +) { + let (peer_writer, sdk_reader) = tokio::io::duplex(8192); + let (sdk_writer, peer_reader) = tokio::io::duplex(8192); + let transport = ByteStreams::new(sdk_writer.compat_write(), sdk_reader.compat()); + + let server = UntypedRole.builder().on_receive_request( + async move |request: TestRequest, + responder: Responder, + connection: ConnectionTo| { + let message = request.message; + if message == "drop responder" { + drop(responder); + return Ok(()); + } + if message.starts_with("deferred") { + return responder_tx + .unbounded_send((message, responder)) + .map_err(agent_client_protocol::Error::into_internal_error); + } + if message == "tracked then notification" { + let receipt = responder.respond_tracked(TestResponse { + result: "tracked response".into(), + })?; + let notification_connection = connection.clone(); + connection.spawn(async move { + receipt.await?; + notification_connection.send_notification(TestNotification { + message: "after tracked response".into(), + }) + })?; + return Ok(()); + } + if message == "untracked" { + return responder.respond(TestResponse { + result: "untracked response".into(), + }); + } + + let response = if message == "tracked error" { + Err(agent_client_protocol::Error::internal_error().data("tracked error")) + } else { + Ok(TestResponse { + result: format!("echo: {message}"), + }) + }; + let receipt = responder.respond_with_result_tracked(response)?; + receipt_tx + .unbounded_send((message, receipt)) + .map_err(agent_client_protocol::Error::into_internal_error) + }, + agent_client_protocol::on_receive_request!(), + ); + + let server_task = tokio::task::spawn_local(server.connect_to(transport)); + (peer_writer, BufReader::new(peer_reader), server_task) +} + fn start_deferred_server( responder_tx: mpsc::UnboundedSender<(String, Responder)>, ) -> ( @@ -281,6 +355,199 @@ async fn finish_server( .expect("server connection failed"); } +#[tokio::test(flavor = "current_thread")] +async fn tracked_individual_response_orders_notification_and_leaves_untracked_unchanged() { + tokio::task::LocalSet::new() + .run_until(async { + let (responder_tx, _responder_rx) = mpsc::unbounded(); + let (receipt_tx, _receipt_rx) = mpsc::unbounded(); + let (mut peer_writer, mut peer_reader, server_task) = + start_tracked_server(responder_tx, receipt_tx); + + write_json_line( + &mut peer_writer, + &json!({ + "jsonrpc": "2.0", + "id": 1, + "method": "test/echo", + "params": { "message": "tracked then notification" } + }), + ) + .await; + + let response = read_json_line(&mut peer_reader).await; + assert_eq!(response["id"], json!(1)); + assert_eq!(response["result"], json!({ "result": "tracked response" })); + let notification = read_json_line(&mut peer_reader).await; + assert_eq!(notification["method"], json!("test/notify")); + assert_eq!( + notification["params"], + json!({ "message": "after tracked response" }) + ); + + write_json_line( + &mut peer_writer, + &json!({ + "jsonrpc": "2.0", + "id": 2, + "method": "test/echo", + "params": { "message": "untracked" } + }), + ) + .await; + let untracked = read_json_line(&mut peer_reader).await; + assert_eq!(untracked["id"], json!(2)); + assert_eq!( + untracked["result"], + json!({ "result": "untracked response" }) + ); + + finish_server(peer_writer, server_task).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn tracked_error_response_receipt_resolves_when_frame_is_enqueued() { + tokio::task::LocalSet::new() + .run_until(async { + let (responder_tx, _responder_rx) = mpsc::unbounded(); + let (receipt_tx, mut receipt_rx) = mpsc::unbounded(); + let (mut peer_writer, mut peer_reader, server_task) = + start_tracked_server(responder_tx, receipt_tx); + + write_json_line( + &mut peer_writer, + &json!({ + "jsonrpc": "2.0", + "id": 3, + "method": "test/echo", + "params": { "message": "tracked error" } + }), + ) + .await; + + let (message, receipt) = next_response_receipt(&mut receipt_rx).await; + assert_eq!(message, "tracked error"); + receipt + .await + .expect("error response frame should be enqueued"); + let response = read_json_line(&mut peer_reader).await; + assert_eq!(response["id"], json!(3)); + assert_eq!(response["error"]["code"], json!(-32603)); + assert_eq!(response["error"]["data"], json!("tracked error")); + + finish_server(peer_writer, server_task).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn tracked_batch_receipts_resolve_together_after_aggregate_is_ready() { + tokio::task::LocalSet::new() + .run_until(async { + let (responder_tx, mut responder_rx) = mpsc::unbounded(); + let (receipt_tx, mut receipt_rx) = mpsc::unbounded(); + let (mut peer_writer, mut peer_reader, server_task) = + start_tracked_server(responder_tx, receipt_tx); + + write_json_line( + &mut peer_writer, + &json!([ + { + "jsonrpc": "2.0", + "id": 4, + "method": "test/echo", + "params": { "message": "immediate tracked" } + }, + { + "jsonrpc": "2.0", + "id": 5, + "method": "test/echo", + "params": { "message": "deferred tracked" } + } + ]), + ) + .await; + + let (message, mut first_receipt) = next_response_receipt(&mut receipt_rx).await; + assert_eq!(message, "immediate tracked"); + let (message, responder) = next_deferred_response(&mut responder_rx).await; + assert_eq!(message, "deferred tracked"); + assert!( + tokio::time::timeout(Duration::from_millis(25), &mut first_receipt) + .await + .is_err(), + "a batch receipt resolved before its sibling response was ready" + ); + + let second_receipt = responder + .respond_tracked(TestResponse { + result: "echo: deferred tracked".into(), + }) + .expect("deferred tracked response should be accepted"); + let (first_result, second_result) = join(first_receipt, second_receipt).await; + first_result.expect("first batch receipt should resolve"); + second_result.expect("second batch receipt should resolve"); + + let response = read_json_line(&mut peer_reader).await; + let responses = response + .as_array() + .expect("tracked batch should emit one aggregate array"); + assert_eq!(responses.len(), 2); + + finish_server(peer_writer, server_task).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn tracked_batch_receipt_resolves_when_sibling_responder_is_abandoned() { + tokio::task::LocalSet::new() + .run_until(async { + let (responder_tx, _responder_rx) = mpsc::unbounded(); + let (receipt_tx, mut receipt_rx) = mpsc::unbounded(); + let (mut peer_writer, mut peer_reader, server_task) = + start_tracked_server(responder_tx, receipt_tx); + + write_json_line( + &mut peer_writer, + &json!([ + { + "jsonrpc": "2.0", + "id": 6, + "method": "test/echo", + "params": { "message": "tracked sibling" } + }, + { + "jsonrpc": "2.0", + "id": 7, + "method": "test/echo", + "params": { "message": "drop responder" } + } + ]), + ) + .await; + + let (message, receipt) = next_response_receipt(&mut receipt_rx).await; + assert_eq!(message, "tracked sibling"); + receipt + .await + .expect("tracked sibling should resolve with abandoned fallback batch"); + let response = read_json_line(&mut peer_reader).await; + let responses = response + .as_array() + .expect("abandoned sibling should still emit one aggregate array"); + assert_eq!(responses.len(), 2); + assert!(responses.iter().any(|response| { + response["id"] == json!(7) && response["error"]["code"] == json!(-32603) + })); + + finish_server(peer_writer, server_task).await; + }) + .await; +} + #[tokio::test(flavor = "current_thread")] async fn mixed_batch_returns_one_array_with_each_response_bearing_entry() { tokio::task::LocalSet::new() From 5b54f8067029095e3461f88c14fd7262676f1c9d Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 29 Sep 2026 01:14:07 +0100 Subject: [PATCH 4/9] feat(http): add opt-in bounded ACP transport admission --- .github/workflows/ci.yml | 2 +- md/http-transport.md | 96 ++ md/transport-architecture.md | 90 ++ src/agent-client-protocol-http/CHANGELOG.md | 8 + src/agent-client-protocol-http/README.md | 21 + .../src/bounded_server.rs | 1314 +++++++++++++++ src/agent-client-protocol-http/src/client.rs | 12 + .../src/client_limits.rs | 1408 +++++++++++++++++ .../src/http_server.rs | 13 +- src/agent-client-protocol-http/src/lib.rs | 7 +- src/agent-client-protocol-http/src/server.rs | 17 + .../tests/bounded_interop.rs | 329 ++++ src/agent-client-protocol/CHANGELOG.md | 4 + src/agent-client-protocol/README.md | 16 + src/agent-client-protocol/src/bounded.rs | 721 +++++++++ src/agent-client-protocol/src/component.rs | 55 + src/agent-client-protocol/src/jsonrpc.rs | 534 ++++++- .../src/jsonrpc/dynamic_handler.rs | 38 + .../src/jsonrpc/incoming_actor.rs | 43 +- .../src/jsonrpc/outgoing_actor.rs | 111 +- .../src/jsonrpc/task_actor.rs | 22 +- src/agent-client-protocol/src/lib.rs | 5 + 22 files changed, 4789 insertions(+), 77 deletions(-) create mode 100644 src/agent-client-protocol-http/src/bounded_server.rs create mode 100644 src/agent-client-protocol-http/src/client_limits.rs create mode 100644 src/agent-client-protocol-http/tests/bounded_interop.rs create mode 100644 src/agent-client-protocol/src/bounded.rs diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index ba581494..49761bc9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,7 +4,7 @@ on: push: branches: [main] pull_request: - branches: [main] + branches: [main, feat/unstable-v2-session-inject] workflow_dispatch: permissions: diff --git a/md/http-transport.md b/md/http-transport.md index c8584e76..d1bd3ad7 100644 --- a/md/http-transport.md +++ b/md/http-transport.md @@ -135,3 +135,99 @@ and it will open a single bidirectional connection instead of using POST + SSE: let transport = HttpClient::new("ws://127.0.0.1:8080")?; my_client().connect_to(transport).await?; ``` + +## Opt-in Bounded HTTP/SSE + +Existing constructors preserve the compatibility-only unbounded transport. +Select finite admission explicitly on each endpoint: + +```rust +use agent_client_protocol_http::{AcpHttpServer, HttpClient, HttpClientLimits, ServerLimits}; + +let app = AcpHttpServer::new_bounded(|| my_agent(), ServerLimits::default())? + .into_router(); +let transport = HttpClient::new("http://127.0.0.1:8080/acp")? + .with_limits(HttpClientLimits::default())?; +my_client().connect_to(transport).await?; +``` + +These return `BoundedAcpHttpServer` and `BoundedHttpClient`. The server retains +`with_options(ServerOptions)` and `into_router()`. No fields were added to +`ServerOptions`. Both endpoints use the core bounded connection interface +directly, without forwarding through the legacy unbounded `Channel`. +Custom components must implement bounded extraction; unsupported legacy-only +components fail closed instead of silently losing the admission guarantee. +Raw adapters can use `BoundedHttpClient::into_bounded_channel_and_future()`; +they must poll the returned driver and retain each `ChargedFrame` until their +own consumption boundary. See [bounded core transports](./transport-architecture.md) +for producer admission and charged frame ownership. + +### Finite Defaults + +Zero limits are invalid. The limits structures are public and can be configured +before construction. Budgets are independent and conservative: fitting one +limit does not guarantee admission through every other limit. + +| Budget | Server default | Client default | +| --- | --- | --- | +| Core maximum serialized frame | 256 KiB | 1 MiB | +| Core reserved bytes / frames, each direction | 16 MiB / 256 | 16 MiB / 256 | +| Core pending requests / tasks | 256 / 256 | 256 / 256 | +| Active logical connections | 64 | One per transport | +| Concurrent POSTs / aggregate request-body reservation | 32 / 8 MiB | Separate request and response lanes | +| HTTP maximum request frame / batch entries | 256 KiB / 128 | Core frame limit | +| HTTP egress bytes / frames | 4 MiB / 64 per connection | 64 MiB / 128 HTTP reservations | +| Pending routed RPC entries | 256 per connection | 128 | +| Registered sessions / active SSE streams | 64 / 65 per connection | 8 streams, including connection stream | +| Queued plus active POSTs | Global concurrent POST admission | 32 request and 32 response-only | +| Response body / SSE line / event / chunk | Outbound frame and egress limits | 1 MiB each | + +Client reservations cover POST bodies, pending metadata, bounded response +workspaces, and SSE parser workspaces; each reservation also consumes a frame +slot. Server POST bodies reserve the maximum frame size before polling the body. +Server connection admission precedes factory invocation. Duplicate request IDs +consume separate pending entries; a successful POST does not imply RPC completion. +Response-only callback POSTs have an independent bounded client lane. Mixed +batches remain in request order. Registered server session mailboxes last until +connection termination and count against the configured session limit. + +### Ownership and Exhaustion + +Core producer admission happens before queueing. Charges survive intermediate +dequeue, routing, and body construction; the server releases yielded body +charges on the subsequent poll or body drop. Client bodies retain charges +through HTTP handoff. This bounds SDK-owned queued data and work, not an +application's cumulative output. A healthy consumer can process more than any +single budget over the connection's lifetime. + +Exhaustion is explicit and fail-fast, not an unbounded queue of waiting sends. +The server rejects pre-acceptance overload with an HTTP error; terminal errors +revoke producers and release pending routes, mailboxes, and local tasks. +POST/body admission exhaustion for an addressed connection terminates that +connection so callbacks cannot remain indefinitely blocked behind saturated +request bodies. Cancellation/drop releases local reservations; client teardown +does not guarantee a remote DELETE completed. + +Encoded-byte budgets are **not hard peak-heap limits**. Parsed JSON and bounded +serialization scratch add overhead; application conversion hooks can allocate +intermediate values before capped normalization. Allocator capacity, arbitrary +application task captures, and HTTP/TLS/socket buffers are outside encoded-byte +accounting. Core reservations and HTTP reservations are additional budgets, not +one shared process-memory counter. + +### Delivery and Recovery Boundaries + +A body poll transfers bytes to the HTTP stack. It is **not** proof of socket +flush, peer parsing, or ACP application consumption. An event already yielded +when SSE disconnects can be lost. No cursor/replay or exactly-once delivery is +provided, and no `Last-Event-ID` recovery is performed. + +The bounded client does not automatically retry accepted or uncertain POSTs; +streaming request bodies are non-replayable, including with custom reqwest +redirect/retry policies. A lost response can leave acceptance unknown. Existing +JSON-RPC IDs remain correlation IDs, not idempotency keys. + +The bounded path supports HTTP/SSE only. The bounded client rejects `ws`/`wss` +URLs, and bounded server WebSocket upgrades return HTTP 501. Legacy WebSocket +support is unchanged. ACP JSON-RPC messages and HTTP connection/session headers +are unchanged; bounded endpoints do not require a private protocol extension. diff --git a/md/transport-architecture.md b/md/transport-architecture.md index a5f7dce3..382699f6 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -67,6 +67,96 @@ boundary, although serializing a relayed batch may normalize whitespace. The protocol actor, not the parser, decides whether malformed input requires a response. +## Opt-in bounded core transport + +`Channel` and its public unbounded sender/receiver fields remain unchanged for +compatibility. They do **not** provide admission or memory bounds. New transports +can instead use `BoundedChannel`, selected through +`ConnectTo::into_transport_and_future` and `TransportChannel::Bounded`. +`DynConnectTo` preserves this selection. A transport that accepts a component +factory can call `into_bounded_channel_and_future(limits)` to give its `Builder` +a bounded endpoint directly, without an unbounded adapter pump. Attempting to +extract a bounded endpoint through the legacy `into_channel_and_future` method +returns a disconnected channel and a failing driver future; it never inserts a +legacy bridge. + +```rust +use agent_client_protocol::{BoundedChannel, ChannelLimits}; + +let limits = ChannelLimits::default(); +let (protocol_endpoint, transport_endpoint) = BoundedChannel::duplex(limits)?; +// Connect the Builder to protocol_endpoint. The transport uses +// transport_endpoint.tx.try_send_serialized(...) and receives ChargedFrame +// values from transport_endpoint.rx. +# Ok::<(), agent_client_protocol::Error>(()) +``` + +The default limits per direction are: + +| Limit | Default | +| --- | ---: | +| Maximum encoded frame | 1 MiB | +| Reserved encoded bytes | 16 MiB | +| Preparing, queued, and handed-off frames/control items | 256 | +| Pending outgoing requests | 256 | +| Queued/running tasks and registered dynamic handlers | 256 | + +Zero limits are invalid; the byte budget must fit one maximum-sized frame. +Admission reserves the **maximum frame size**, not the eventual encoded length, +so the default byte limit permits at most 16 simultaneous frame reservations. +This conservative reservation allows synchronous producer admission before a +message's size is known. The transport driver itself occupies one task slot. + +Bounded protocol connections use fail-fast admission before outgoing message +conversion/enqueue and before task/dynamic-handler enqueue. Pending requests +are registered under a capped registry lock. Overload is terminal: it wakes the +driver, rejects escaped producer handles, and drops queued/running work. There +is no queue of producer futures awaiting a permit. An accepted bounded +notification that later fails protocol preparation or raw-message conversion +fails the connection; it is not silently discarded. The legacy unbounded path +retains its prior log-and-continue behavior. Internal control messages +also require admission; overload during destructor-originated batch completion +terminates the connection rather than silently stranding the batch. + +`BoundedSender::close_channel()` gracefully closes one direction for every +sender clone. It rejects later sends without signaling terminal failure, keeps +already-enqueued transport frames available for drain, and delivers EOF after +the queue drains. The opposite direction remains open. Successful protocol +driver completion likewise preserves admitted transport frames; cancellation, +overload, and driver errors still make the connection terminal. + +A `ChargedFrame` owns serialized JSON and its RAII reservation. Dequeueing does +not release that reservation. `try_forward(frame)` transfers ownership; a +cross-budget handoff reserves destination capacity before releasing source +capacity. The protocol actor retains producer charges through conversion and +batch accumulation. Incoming dispatch, deferred messages, and response slots +retain their incoming charges while the core owns their data. A batch is still +one wire frame and must fit the configured maximum including array punctuation. +HTTP adapters must keep the charged frame alive through their declared body +handoff/drop point. **Body consumption is not a peer acknowledgment**: it does +not establish socket flush, peer parsing, ACP dispatch, replay, or exactly-once +delivery. + +These are encoded-data and work-count bounds, not a claim about exact process +heap usage. The bounded frame queue stores compact serialized bytes. Logical +outgoing values are normalized with capped serialization and decoding before +queueing, so caller-provided spare `String`/`Vec` capacity is not retained. +Decoded JSON, batch entries, request metadata, and temporary serialization have +additional structural overhead proportional to admitted data; pending request +metadata also has its separately capped entry count. Application-owned values, +future captures, retained copies, and intermediate allocations in +`JsonRpcMessage`/`JsonRpcResponse` conversion are outside this byte budget. +This includes **standard SDK conversion implementations** that clone an untyped +value or construct a raw `serde_json::Value`, not only user-provided hooks. +A conversion holds admission before it runs, but these existing traits can +produce a raw value before the core checks its encoded size. Admission bounds +the number of concurrent conversions; it does not bound a conversion's peak +allocation. The later wire serializer and retained normalized queue values are +capped. Applications needing a hard peak/process-heap limit must also bound +their conversion inputs and implementations. +Transport-specific HTTP body, session, stream, and POST concurrency budgets +remain the responsibility of the HTTP transport, not these core limits. + ## Actor Architecture ### Protocol Actors diff --git a/src/agent-client-protocol-http/CHANGELOG.md b/src/agent-client-protocol-http/CHANGELOG.md index c58d6bdd..85b8ab2e 100644 --- a/src/agent-client-protocol-http/CHANGELOG.md +++ b/src/agent-client-protocol-http/CHANGELOG.md @@ -2,6 +2,14 @@ ## [Unreleased] +### Added + +- Opt-in bounded HTTP/SSE client and server integrated with core bounded + producer admission. Configurable byte/frame, pending-work, POST, connection, + session, and stream limits fail explicitly on exhaustion while preserving + existing constructors and ACP wire shapes. Bounded transports do not provide + SSE replay or automatic POST retries. + ## [2.0.0](https://github.com/agentclientprotocol/rust-sdk/compare/agent-client-protocol-http-v1.3.0...agent-client-protocol-http-v2.0.0) - 2026-07-23 ### Breaking changes diff --git a/src/agent-client-protocol-http/README.md b/src/agent-client-protocol-http/README.md index b659472e..b4934c1c 100644 --- a/src/agent-client-protocol-http/README.md +++ b/src/agent-client-protocol-http/README.md @@ -18,4 +18,25 @@ with `CorsOptions::allow_origins(...)` to allow specific browser origins. Core SDK request cancellation support is forwarded through this transport. +## Opt-in bounded HTTP/SSE + +Use `AcpHttpServer::new_bounded(factory, ServerLimits::default())?` and +`HttpClient::new(url)?.with_limits(HttpClientLimits::default())?` to select the +bounded transport path. Existing constructors and `ServerOptions` literals keep +their compatibility behavior; legacy transports remain unbounded. + +The bounded path integrates directly with the core `BoundedChannel` and +producer admission. Finite defaults constrain serialized bytes, frame counts, +pending work, POSTs, sessions, and streams. Exhaustion fails explicitly rather +than retaining arbitrarily many waiters. Reservations survive dequeue and +remain held until the documented HTTP handoff or cancellation/drop. + +These limits are **not** peer-delivery acknowledgments or hard process-heap +limits. Application conversion allocations, allocator overhead, and +reqwest/TLS/socket buffers are outside encoded-byte accounting. HTTP/SSE is +supported; bounded WebSocket endpoints are rejected. There is no SSE replay, +cursor recovery, or automatic retry of accepted/uncertain POSTs. See the book's +[HTTP transport chapter](https://agentclientprotocol.github.io/rust-sdk/http-transport.html) +for limits, ownership, and overload semantics. + See the [documentation](https://docs.rs/agent-client-protocol-http) for usage examples. diff --git a/src/agent-client-protocol-http/src/bounded_server.rs b/src/agent-client-protocol-http/src/bounded_server.rs new file mode 100644 index 00000000..8c006b85 --- /dev/null +++ b/src/agent-client-protocol-http/src/bounded_server.rs @@ -0,0 +1,1314 @@ +//! Opt-in HTTP-only server. This deliberately never uses the legacy connection pumps. +use std::{ + collections::{HashMap, VecDeque}, + convert::Infallible, + sync::{Arc, Mutex, Weak}, +}; + +use agent_client_protocol::{ + BoundedChannel, BoundedSender, ChannelLimits, ChargedFrame, Client, ConnectTo, + RawJsonRpcMessage, TransportBatchEntry, TransportFrame, schema::v1::RequestId, +}; +use axum::{ + Router, + body::Body, + extract::State, + http::{HeaderMap, HeaderValue, Request, StatusCode, header}, + response::{IntoResponse, Response}, + routing::{delete, get, post}, +}; +use futures::{StreamExt, future::BoxFuture}; +use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore, oneshot}; + +use super::ServerOptions; +use crate::{ + connection::ResponseRoute, + http_server::{ + collect_route, initial_initialize_request, initialize_response_failed, + prepare_message_route, + }, + protocol::{HEADER_CONNECTION_ID, HEADER_SESSION_ID, JSON_MIME_TYPE, session_id_from_message}, +}; + +/// Finite limits for the opt-in HTTP server, independent of the core channel limits. +/// +/// Byte limits count encoded JSON/SSE bytes, not allocator overhead. POST bodies reserve +/// `max_frame_bytes` before their first poll. Parsed JSON and serialization scratch are +/// additional bounded multiples of that size. Each connection has its own egress budget; +/// total server egress is therefore at most `max_connections * max_egress_bytes`. +#[derive(Clone, Debug)] +pub struct ServerLimits { + /// Core producer/transport admission limits supplied to the component factory. + pub channel_limits: ChannelLimits, + pub max_connections: usize, + pub max_in_flight_posts: usize, + pub max_body_bytes: usize, + pub max_frame_bytes: usize, + pub max_batch_entries: usize, + pub max_pending_routes: usize, + pub max_registered_sessions: usize, + pub max_active_streams: usize, + pub max_egress_frames: usize, + pub max_egress_bytes: usize, +} + +impl Default for ServerLimits { + fn default() -> Self { + Self { + channel_limits: ChannelLimits { + max_frame_bytes: 256 * 1024, + ..ChannelLimits::default() + }, + max_connections: 64, + max_in_flight_posts: 32, + max_body_bytes: 8 * 1024 * 1024, + max_frame_bytes: 256 * 1024, + max_batch_entries: 128, + max_pending_routes: 256, + max_registered_sessions: 64, + max_active_streams: 65, + max_egress_frames: 64, + max_egress_bytes: 4 * 1024 * 1024, + } + } +} + +/// An invalid bounded server configuration. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct ServerLimitsError(pub &'static str); + +impl std::fmt::Display for ServerLimitsError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(self.0) + } +} +impl std::error::Error for ServerLimitsError {} + +impl ServerLimits { + pub fn validate(&self) -> Result<(), ServerLimitsError> { + self.channel_limits + .validate() + .map_err(|_| ServerLimitsError("invalid core channel limits"))?; + for value in [ + self.max_connections, + self.max_in_flight_posts, + self.max_body_bytes, + self.max_frame_bytes, + self.max_batch_entries, + self.max_pending_routes, + self.max_registered_sessions, + self.max_active_streams, + self.max_egress_frames, + self.max_egress_bytes, + ] { + if value == 0 || value > u32::MAX as usize || value > Semaphore::MAX_PERMITS { + return Err(ServerLimitsError( + "limits must be nonzero and fit semaphore counters", + )); + } + } + if self.max_body_bytes < self.max_frame_bytes + || self.max_egress_bytes < self.max_frame_bytes + 8 + { + return Err(ServerLimitsError( + "body and egress budgets must fit one maximum frame", + )); + } + if self + .max_connections + .checked_mul(self.max_egress_bytes) + .is_none() + { + return Err(ServerLimitsError("aggregate egress budget overflows usize")); + } + Ok(()) + } +} + +type Factory = dyn Fn( + ChannelLimits, + ) -> agent_client_protocol::Result<( + BoundedChannel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + )> + Send + + Sync; + +/// HTTP/SSE-only bounded server. WebSocket upgrades are explicitly rejected. +/// +/// Components are connected via `ConnectTo::into_bounded_channel_and_future`. Components +/// that only extract legacy channels fail closed; no unbounded adapter is introduced. +/// Admission errors before core enqueue return 429/413. An admitted frame is never retried; +/// later overload terminates the connection. Body polling is not peer acknowledgment. +/// POST admission is fail-fast. Saturation of POST/body slots also terminates the +/// addressed connection, so callback responses cannot remain indefinitely blocked behind +/// slow request bodies. Callers must not automatically retry uncertain POST results. +/// Registered session mailboxes persist until connection termination; replay is not provided. +pub struct BoundedAcpHttpServer { + state: Arc, + options: ServerOptions, +} + +impl std::fmt::Debug for BoundedAcpHttpServer { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BoundedAcpHttpServer") + .field("limits", &self.state.limits) + .field("options", &self.options) + .finish_non_exhaustive() + } +} + +impl BoundedAcpHttpServer { + pub(super) fn new(factory: F, limits: ServerLimits) -> Result + where + F: Fn() -> C + Send + Sync + 'static, + C: ConnectTo, + { + limits.validate()?; + Ok(Self { + state: Arc::new(Registry { + factory: Arc::new(move |limits| factory().into_bounded_channel_and_future(limits)), + connections: Mutex::new(HashMap::new()), + connection_slots: Arc::new(Semaphore::new(limits.max_connections)), + posts: Arc::new(Semaphore::new(limits.max_in_flight_posts)), + bodies: Arc::new(Semaphore::new(limits.max_body_bytes)), + limits, + }), + options: ServerOptions::default(), + }) + } + + #[must_use] + pub fn with_options(mut self, options: ServerOptions) -> Self { + self.options = options; + self + } + + pub fn into_router(self) -> Router { + let mut router = Router::new() + .route(&self.options.path, post(handle_post)) + .route(&self.options.path, get(handle_get)) + .route(&self.options.path, delete(handle_delete)) + .with_state(self.state); + if self.options.health_endpoint { + router = router.route("/health", get(super::health)); + } + if let Some(origin) = self.options.cors.allow_origin_layer() { + router = router.layer(super::default_cors(origin)); + } + router + } +} + +struct Registry { + factory: Arc, + limits: ServerLimits, + connections: Mutex>>, + connection_slots: Arc, + posts: Arc, + bodies: Arc, +} + +impl Drop for Registry { + fn drop(&mut self) { + for connection in self.connections.get_mut().unwrap().values() { + connection.close(); + } + } +} + +struct Connection { + id: String, + registry: Weak, + limits: ServerLimits, + inner: Mutex, + wake: Notify, + frame_slots: Arc, + byte_slots: Arc, + stream_slots: Arc, + // Held until tasks, requests and response bodies have released this connection. + _slot: OwnedSemaphorePermit, +} + +struct ConnectionState { + tx: Option, + closed: bool, + task: Option, + pending: VecDeque<(RequestId, ResponseRoute)>, + streams: HashMap, Mailbox>, +} + +#[derive(Default)] +struct Mailbox { + queue: VecDeque, + subscribed: bool, +} + +struct Envelope { + // Keep the core byte/frame admission across decode, routing and HTTP body handoff. + frame: ChargedFrame, + _frame_slot: OwnedSemaphorePermit, + _byte_slot: OwnedSemaphorePermit, +} + +impl Connection { + fn close(&self) { + { + let mut state = self.inner.lock().unwrap(); + state.closed = true; + if let Some(tx) = state.tx.take() { + tx.fail("HTTP connection terminated"); + } + state.pending.clear(); + state.streams.clear(); + if let Some(task) = state.task.take() { + task.abort(); + } + } + self.wake.notify_waiters(); + } + + fn terminate(&self) { + self.close(); + if let Some(registry) = self.registry.upgrade() { + registry.connections.lock().unwrap().remove(&self.id); + } + } + + fn envelope(&self, frame: ChargedFrame) -> Result { + let len = frame.as_bytes().len(); + if len > self.limits.max_frame_bytes || frame.as_bytes().contains(&b'\r') { + return Err(()); + } + // Reserve actual SSE framing too, before its allocation. JSON is UTF-8; the + // serializer normally emits one line, but raw/malformed frames may contain LF. + let text = std::str::from_utf8(frame.as_bytes()).map_err(|_| ())?; + let wire_len = len + .checked_add(8) + .and_then(|n| n.checked_add(text.matches('\n').count() * 6)) + .ok_or(())?; + let frames = self + .frame_slots + .clone() + .try_acquire_owned() + .map_err(|_| ())?; + let bytes = self + .byte_slots + .clone() + .try_acquire_many_owned(u32::try_from(wire_len).map_err(|_| ())?) + .map_err(|_| ())?; + Ok(Envelope { + frame, + _frame_slot: frames, + _byte_slot: bytes, + }) + } + + fn complete_initial_routes(&self, frame: &TransportFrame) { + let mut state = self.inner.lock().unwrap(); + let mut complete = |message: &RawJsonRpcMessage| { + if message.response_id().is_some() { + outbound_route(&mut state.pending, message); + } + }; + match frame { + TransportFrame::Single(message) => complete(message), + TransportFrame::Batch(batch) => { + for entry in batch.entries() { + if let TransportBatchEntry::Message(message) = entry { + complete(message); + } + } + } + TransportFrame::Malformed { .. } => {} + } + } + + fn route(&self, envelope: Envelope, frame: &TransportFrame) -> Result<(), ()> { + let mut state = self.inner.lock().unwrap(); + if state.closed { + return Err(()); + } + let route = match frame { + TransportFrame::Single(message) => outbound_route(&mut state.pending, message), + TransportFrame::Malformed { .. } => ResponseRoute::Connection, + TransportFrame::Batch(batch) => { + let mut common = None; + let mut mixed = false; + for entry in batch.entries() { + let route = match entry { + TransportBatchEntry::Message(message) => { + outbound_route(&mut state.pending, message) + } + TransportBatchEntry::Malformed { .. } => ResponseRoute::Connection, + }; + match &common { + None => common = Some(route), + Some(previous) if previous == &route => {} + Some(_) => mixed = true, + } + } + if mixed { + ResponseRoute::Connection + } else { + common.unwrap_or(ResponseRoute::Connection) + } + } + }; + let key = route_key(route); + ensure_stream(&mut state, key, self.limits.max_registered_sessions)? + .queue + .push_back(envelope); + drop(state); + self.wake.notify_waiters(); + Ok(()) + } +} + +fn route_key(route: ResponseRoute) -> Option { + match route { + ResponseRoute::Connection => None, + ResponseRoute::Session(id) => Some(id), + } +} + +fn ensure_stream( + state: &mut ConnectionState, + key: Option, + max: usize, +) -> Result<&mut Mailbox, ()> { + if !state.streams.contains_key(&key) && key.is_some() && state.streams.len() > max { + return Err(()); + } + Ok(state.streams.entry(key).or_default()) +} + +fn outbound_route( + pending: &mut VecDeque<(RequestId, ResponseRoute)>, + message: &RawJsonRpcMessage, +) -> ResponseRoute { + if let Some(id) = message.response_id() { + if let Some(index) = pending.iter().position(|(pending_id, _)| pending_id == id) { + return pending.remove(index).unwrap().1; + } + return ResponseRoute::Connection; + } + session_id_from_message(message).map_or(ResponseRoute::Connection, ResponseRoute::Session) +} + +struct Cleanup { + connection: Arc, + armed: bool, +} +impl Drop for Cleanup { + fn drop(&mut self) { + if self.armed { + self.connection.terminate(); + } + } +} + +fn header_value(headers: &HeaderMap, name: &str) -> Result, StatusCode> { + headers + .get(name) + .map(|value| { + if value.len() > 1024 { + return Err(StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } + value + .to_str() + .map(str::to_owned) + .map_err(|_| StatusCode::BAD_REQUEST) + }) + .transpose() +} + +fn check_batch(frame: &TransportFrame, max: usize) -> bool { + match frame { + TransportFrame::Batch(batch) => batch.entries().take(max + 1).count() <= max, + _ => true, + } +} + +// Serialize into a capped writer before core enqueue, including session-header injection. +fn encode_frame(frame: &TransportFrame, max: usize) -> Result { + struct Capped { + bytes: Vec, + max: usize, + } + impl std::io::Write for Capped { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.max - self.bytes.len() { + return Err(std::io::Error::other("frame limit exceeded")); + } + self.bytes.extend_from_slice(bytes); + Ok(bytes.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + let mut output = Capped { + bytes: Vec::new(), + max, + }; + let result = match frame { + TransportFrame::Single(message) => serde_json::to_writer(&mut output, message), + TransportFrame::Batch(batch) => serde_json::to_writer(&mut output, batch), + TransportFrame::Malformed { raw, .. } => { + if raw.len() > max { + return Err(StatusCode::PAYLOAD_TOO_LARGE); + } + return Ok(raw.clone()); + } + }; + result.map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?; + Ok(String::from_utf8(output.bytes).expect("JSON is UTF-8")) +} + +async fn handle_post(State(registry): State>, request: Request) -> Response { + match post_inner(registry, request).await { + Ok(response) => response, + Err(status) => status.into_response(), + } +} + +async fn post_inner( + registry: Arc, + request: Request, +) -> Result { + if !request + .headers() + .get(header::CONTENT_TYPE) + .and_then(|v| v.to_str().ok()) + .is_some_and(|v| v.starts_with(JSON_MIME_TYPE)) + { + return Err(StatusCode::UNSUPPORTED_MEDIA_TYPE); + } + let connection_id = header_value(request.headers(), HEADER_CONNECTION_ID)?; + let session_id = header_value(request.headers(), HEADER_SESSION_ID)?; + // Body shape is unknown until read: rather than let slow request bodies starve + // callback responses indefinitely, saturation terminates the addressed connection. + // The rejected POST itself has not been accepted and is never resubmitted here. + let saturated = || { + if let Some(id) = &connection_id { + let connection = registry.connections.lock().unwrap().get(id).cloned(); + if let Some(connection) = connection { + connection.terminate(); + } + } + StatusCode::TOO_MANY_REQUESTS + }; + let post_slot = registry + .posts + .clone() + .try_acquire_owned() + .map_err(|_| saturated())?; + let body_slot = registry + .bodies + .clone() + .try_acquire_many_owned( + u32::try_from(registry.limits.max_frame_bytes).expect("validated frame limit"), + ) + .map_err(|_| saturated())?; + // Admission before the first body read and before factory invocation. + let creating = connection_id.is_none(); + let connection_slot = if creating { + Some( + registry + .connection_slots + .clone() + .try_acquire_owned() + .map_err(|_| StatusCode::TOO_MANY_REQUESTS)?, + ) + } else { + None + }; + if request + .headers() + .get(header::CONTENT_LENGTH) + .and_then(|v| v.to_str().ok()) + .and_then(|v| v.parse::().ok()) + .is_some_and(|n| n > registry.limits.max_frame_bytes) + { + return Err(StatusCode::PAYLOAD_TOO_LARGE); + } + let body = axum::body::to_bytes(request.into_body(), registry.limits.max_frame_bytes) + .await + .map_err(|_| StatusCode::PAYLOAD_TOO_LARGE)?; + let text = std::str::from_utf8(&body).map_err(|_| StatusCode::BAD_REQUEST)?; + let mut frame = TransportFrame::parse_json(text); + if !check_batch(&frame, registry.limits.max_batch_entries) { + return Err(StatusCode::PAYLOAD_TOO_LARGE); + } + let initialize = initial_initialize_request(&frame).map(|(id, _)| id.clone()); + if creating != initialize.is_some() { + return Err(StatusCode::BAD_REQUEST); + } + let mut sessions = Vec::new(); + let mut pending = Vec::new(); + let mut prepare = |message: &mut RawJsonRpcMessage| -> Result<(), StatusCode> { + let route = prepare_message_route(message, session_id.as_deref()) + .map_err(|_| StatusCode::BAD_REQUEST)?; + collect_route(message, route, &mut sessions, &mut pending); + Ok(()) + }; + match &mut frame { + TransportFrame::Single(message) => prepare(message)?, + TransportFrame::Batch(batch) => { + for entry in batch.entries_mut() { + if let TransportBatchEntry::Message(message) = entry { + prepare(message)?; + } + } + } + TransportFrame::Malformed { .. } => {} + } + // Reject impossible route reservations before any factory work. Existing connection + // totals are checked transactionally below, immediately before acceptance. + if pending.len() > registry.limits.max_pending_routes { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + sessions.sort(); + sessions.dedup(); + if sessions.len() > registry.limits.max_registered_sessions { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + let encoded = encode_frame(&frame, registry.limits.max_frame_bytes)?; + drop(frame); + let mut initialization = None; + let connection = if let Some(connection_id) = connection_id { + registry + .connections + .lock() + .unwrap() + .get(&connection_id) + .cloned() + .ok_or(StatusCode::NOT_FOUND)? + } else { + let (BoundedChannel { tx, mut rx }, agent) = + (registry.factory)(registry.limits.channel_limits) + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + let mut streams = HashMap::new(); + streams.insert(None, Mailbox::default()); + let connection = Arc::new(Connection { + id: uuid::Uuid::new_v4().to_string(), + registry: Arc::downgrade(®istry), + limits: registry.limits.clone(), + inner: Mutex::new(ConnectionState { + tx: Some(tx), + closed: false, + task: None, + pending: VecDeque::new(), + streams, + }), + wake: Notify::new(), + frame_slots: Arc::new(Semaphore::new(registry.limits.max_egress_frames)), + byte_slots: Arc::new(Semaphore::new(registry.limits.max_egress_bytes)), + stream_slots: Arc::new(Semaphore::new(registry.limits.max_active_streams)), + _slot: connection_slot.unwrap(), + }); + let (init_tx, init_rx) = oneshot::channel(); + initialization = Some(init_rx); + registry + .connections + .lock() + .unwrap() + .insert(connection.id.clone(), connection.clone()); + let cleanup = Cleanup { + connection: connection.clone(), + armed: true, + }; + let init_id = initialize.unwrap(); + let task_connection = connection.clone(); + // Capture cleanup before spawning: dropping even a never-polled task cleans up. + let task = tokio::spawn(async move { + let _cleanup = cleanup; + let router = async move { + let mut init_tx = Some(init_tx); + while let Some(charged) = rx.next().await { + if charged.as_bytes().len() > task_connection.limits.max_frame_bytes { + break; + } + let frame = charged.decode(); + if !check_batch(&frame, task_connection.limits.max_batch_entries) { + break; + } + let Ok(envelope) = task_connection.envelope(charged) else { + break; + }; + if let Some(failed) = initialize_response_failed(&frame, &init_id) + && let Some(sender) = init_tx.take() + { + task_connection.complete_initial_routes(&frame); + if sender.send((envelope, failed)).is_err() { + break; + } + continue; + } + if task_connection.route(envelope, &frame).is_err() { + break; + } + } + }; + // A bare channel's driver is immediately successful. Success must not + // discard frames still owned by the channel or an escaped producer. + let agent = async move { + if agent.await.is_ok() { + futures::future::pending::<()>().await; + } + }; + futures::pin_mut!(agent, router); + let _ = futures::future::select(router, agent).await; + }); + { + let mut state = connection.inner.lock().unwrap(); + if state.closed { + task.abort(); + } else { + state.task = Some(task.abort_handle()); + } + } + connection + }; + let mut cleanup = Cleanup { + connection: connection.clone(), + armed: initialization.is_some(), + }; + { + // One lock makes route reservation, duplicate-ID ordering and enqueue transactional. + // There is no await/cancellation point between bookkeeping and core acceptance. + let mut state = connection.inner.lock().unwrap(); + if state.closed { + return Err(StatusCode::GONE); + } + if state + .pending + .len() + .checked_add(pending.len()) + .is_none_or(|n| n > connection.limits.max_pending_routes) + { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + sessions.sort(); + sessions.dedup(); + let new_sessions = sessions + .iter() + .filter(|id| !state.streams.contains_key(&Some((*id).clone()))) + .count(); + if state.streams.len() - 1 + new_sessions > connection.limits.max_registered_sessions { + return Err(StatusCode::TOO_MANY_REQUESTS); + } + // try_send is the acceptance boundary. A core error can close its bounded channel; + // terminate instead of implying that resubmission of this POST is safe. + if state + .tx + .as_ref() + .unwrap() + .try_send_serialized(&encoded) + .is_err() + { + drop(state); + connection.terminate(); + return Err(StatusCode::INTERNAL_SERVER_ERROR); + } + for id in sessions { + state.streams.entry(Some(id)).or_default(); + } + state.pending.extend(pending); + } + drop(encoded); + drop(body); + drop(body_slot); + drop(post_slot); + let Some(initialization) = initialization else { + return Ok(StatusCode::ACCEPTED.into_response()); + }; + let (envelope, failed) = initialization + .await + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + if failed { + connection.terminate(); + } else if connection.inner.lock().unwrap().closed { + return Err(StatusCode::INTERNAL_SERVER_ERROR); + } + // An unpolled/dropped initialization body still owns cleanup. Only yielding its bytes + // publishes success; this is an HTTP-stack handoff, not a peer acknowledgment. + let id = connection.id.clone(); + let stream = async_stream::stream! { + let envelope = envelope; + let bytes = envelope.frame.as_bytes().to_vec(); + cleanup.armed = false; + yield Ok::<_, Infallible>(bytes); + drop(envelope); + drop(cleanup); + }; + let mut response = Body::from_stream(stream).into_response(); + response.headers_mut().insert( + header::CONTENT_TYPE, + HeaderValue::from_static(JSON_MIME_TYPE), + ); + if !failed { + response + .headers_mut() + .insert(HEADER_CONNECTION_ID, HeaderValue::from_str(&id).unwrap()); + } + Ok(response) +} + +struct Lease { + connection: Arc, + key: Option, + _slot: OwnedSemaphorePermit, +} +impl Drop for Lease { + fn drop(&mut self) { + if let Some(mailbox) = self + .connection + .inner + .lock() + .unwrap() + .streams + .get_mut(&self.key) + { + mailbox.subscribed = false; + } + } +} + +async fn handle_get(State(registry): State>, request: Request) -> Response { + match get_inner(registry, request) { + Ok(response) => response, + Err(status) => status.into_response(), + } +} + +fn get_inner(registry: Arc, request: Request) -> Result { + if request.headers().contains_key(header::UPGRADE) { + return Err(StatusCode::NOT_IMPLEMENTED); + } + let id = + header_value(request.headers(), HEADER_CONNECTION_ID)?.ok_or(StatusCode::BAD_REQUEST)?; + let key = header_value(request.headers(), HEADER_SESSION_ID)?; + let connection = registry + .connections + .lock() + .unwrap() + .get(&id) + .cloned() + .ok_or(StatusCode::NOT_FOUND)?; + let slot = connection + .stream_slots + .clone() + .try_acquire_owned() + .map_err(|_| StatusCode::TOO_MANY_REQUESTS)?; + { + let mut state = connection.inner.lock().unwrap(); + if state.closed { + return Err(StatusCode::GONE); + } + let mailbox = ensure_stream( + &mut state, + key.clone(), + connection.limits.max_registered_sessions, + ) + .map_err(|()| StatusCode::TOO_MANY_REQUESTS)?; + if mailbox.subscribed { + return Err(StatusCode::CONFLICT); + } + mailbox.subscribed = true; + } + let lease = Lease { + connection, + key: key.clone(), + _slot: slot, + }; + let stream = async_stream::stream! { + let lease = lease; + loop { + // Register before checking the queue to avoid a lost notify_waiters race. + let notified = lease.connection.wake.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + let (envelope, closed) = { + let mut state = lease.connection.inner.lock().unwrap(); + let envelope = state.streams.get_mut(&lease.key).and_then(|m| m.queue.pop_front()); + (envelope, state.closed) + }; + if let Some(envelope) = envelope { + let text = std::str::from_utf8(envelope.frame.as_bytes()).expect("serialized JSON is UTF-8"); + let mut event = String::with_capacity(text.len() + 8 + text.matches('\n').count() * 6); + for line in text.split('\n') { event.push_str("data: "); event.push_str(line); event.push('\n'); } + event.push('\n'); + yield Ok::<_, Infallible>(event.into_bytes()); + // Retained through yield; release on next body poll or cancellation/drop. + drop(envelope); + } else if closed { break; } else { notified.await; } + } + }; + let mut response = Body::from_stream(stream).into_response(); + response.headers_mut().insert( + header::CONTENT_TYPE, + HeaderValue::from_static("text/event-stream"), + ); + response + .headers_mut() + .insert(header::CACHE_CONTROL, HeaderValue::from_static("no-cache")); + response + .headers_mut() + .insert(HEADER_CONNECTION_ID, HeaderValue::from_str(&id).unwrap()); + if let Some(key) = key { + response + .headers_mut() + .insert(HEADER_SESSION_ID, HeaderValue::from_str(&key).unwrap()); + } + Ok(response) +} + +async fn handle_delete(State(registry): State>, request: Request) -> Response { + let Ok(Some(id)) = header_value(request.headers(), HEADER_CONNECTION_ID) else { + return StatusCode::BAD_REQUEST.into_response(); + }; + let connection = registry.connections.lock().unwrap().remove(&id); + let Some(connection) = connection else { + return StatusCode::NOT_FOUND.into_response(); + }; + connection.close(); + StatusCode::ACCEPTED.into_response() +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + use tokio::time::{Duration, timeout}; + + fn registry(limits: ServerLimits) -> (Arc, BoundedChannel) { + let (agent, transport) = BoundedChannel::duplex(limits.channel_limits).unwrap(); + let transport = Mutex::new(Some(transport)); + let server = + BoundedAcpHttpServer::new(move || transport.lock().unwrap().take().unwrap(), limits) + .unwrap(); + (server.state, agent) + } + + fn post_request(body: impl Into, id: Option<&str>) -> Request { + let mut request = Request::builder() + .method("POST") + .header(header::CONTENT_TYPE, JSON_MIME_TYPE); + if let Some(id) = id { + request = request.header(HEADER_CONNECTION_ID, id); + } + request.body(body.into()).unwrap() + } + + async fn initialize(registry: Arc, agent: &mut BoundedChannel) -> (String, Response) { + let task = tokio::spawn(handle_post( + State(registry), + post_request( + r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}"#, + None, + ), + )); + let received = timeout(Duration::from_secs(2), agent.rx.next()) + .await + .unwrap() + .unwrap(); + assert!(matches!(received.decode(), TransportFrame::Single(_))); + drop(received); + agent + .tx + .try_send_serialized(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#) + .unwrap(); + let response = timeout(Duration::from_secs(2), task) + .await + .unwrap() + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + let id = response.headers()[HEADER_CONNECTION_ID] + .to_str() + .unwrap() + .to_owned(); + (id, response) + } + + #[test] + fn validates_finite_limits_and_capped_encoding() { + let limits = ServerLimits { + max_connections: 0, + ..ServerLimits::default() + }; + assert!(limits.validate().is_err()); + let frame = TransportFrame::parse_json(r#"{"jsonrpc":"2.0","method":"ping"}"#); + let encoded = encode_frame(&frame, 1024).unwrap(); + assert!(encode_frame(&frame, encoded.len()).is_ok()); + assert_eq!( + encode_frame(&frame, encoded.len() - 1), + Err(StatusCode::PAYLOAD_TOO_LARGE) + ); + } + + #[tokio::test] + async fn post_and_connection_admission_precede_body_poll_and_factory() { + let calls = Arc::new(AtomicUsize::new(0)); + let factory_calls = calls.clone(); + let state = BoundedAcpHttpServer::new( + move || -> BoundedChannel { + factory_calls.fetch_add(1, Ordering::SeqCst); + panic!("factory must not run"); + }, + ServerLimits { + max_connections: 1, + max_in_flight_posts: 1, + ..ServerLimits::default() + }, + ) + .unwrap() + .state; + let polls = Arc::new(AtomicUsize::new(0)); + for hold_connection in [false, true] { + let held = if hold_connection { + state.connection_slots.clone() + } else { + state.posts.clone() + } + .try_acquire_owned() + .unwrap(); + let body_polls = polls.clone(); + let body = Body::from_stream(futures::stream::poll_fn(move |_| { + body_polls.fetch_add(1, Ordering::SeqCst); + std::task::Poll::Ready(Some(Ok::<_, Infallible>("{}"))) + })); + let response = handle_post(State(state.clone()), post_request(body, None)).await; + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!(polls.load(Ordering::SeqCst), 0); + assert_eq!(calls.load(Ordering::SeqCst), 0); + drop(held); + } + assert_eq!(state.posts.available_permits(), 1); + assert_eq!( + state.bodies.available_permits(), + state.limits.max_body_bytes + ); + } + + #[tokio::test] + async fn dropping_unpolled_initialize_body_terminates_connection() { + let (state, mut agent) = registry(ServerLimits::default()); + let (id, response) = initialize(state.clone(), &mut agent).await; + assert!(state.connections.lock().unwrap().contains_key(&id)); + drop(response); + assert!(!state.connections.lock().unwrap().contains_key(&id)); + tokio::task::yield_now().await; + assert_eq!( + state.connection_slots.available_permits(), + state.limits.max_connections + ); + } + + #[tokio::test] + async fn dropped_initialization_request_cleans_up_without_detached_cleanup() { + let (state, mut agent) = registry(ServerLimits::default()); + let task = tokio::spawn(handle_post( + State(state.clone()), + post_request( + r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}"#, + None, + ), + )); + let received = timeout(Duration::from_secs(2), agent.rx.next()) + .await + .unwrap() + .unwrap(); + drop(received); + task.abort(); + drop(task.await); + tokio::task::yield_now().await; + assert!(state.connections.lock().unwrap().is_empty()); + assert_eq!( + state.posts.available_permits(), + state.limits.max_in_flight_posts + ); + assert_eq!( + state.bodies.available_permits(), + state.limits.max_body_bytes + ); + } + + #[tokio::test] + async fn sse_lease_and_egress_charge_survive_body_handoff_until_drop() { + let (state, mut agent) = registry(ServerLimits::default()); + let (id, response) = initialize(state.clone(), &mut agent).await; + axum::body::to_bytes(response.into_body(), 1024) + .await + .unwrap(); + let connection = state.connections.lock().unwrap().get(&id).unwrap().clone(); + agent + .tx + .try_send_serialized(r#"{"jsonrpc":"2.0","method":"notice"}"#) + .unwrap(); + for _ in 0..10 { + tokio::task::yield_now().await; + } + assert_eq!( + connection.frame_slots.available_permits(), + state.limits.max_egress_frames - 1 + ); + let get = || { + Request::builder() + .header(HEADER_CONNECTION_ID, &id) + .body(Body::empty()) + .unwrap() + }; + let response = get_inner(state.clone(), get()).unwrap(); + assert_eq!( + get_inner(state.clone(), get()).unwrap_err(), + StatusCode::CONFLICT + ); + let mut body = response.into_body().into_data_stream(); + let bytes = timeout(Duration::from_secs(2), body.next()) + .await + .unwrap() + .unwrap() + .unwrap(); + assert!(std::str::from_utf8(&bytes).unwrap().starts_with("data: ")); + assert_eq!( + connection.frame_slots.available_permits(), + state.limits.max_egress_frames - 1 + ); + drop(body); + assert_eq!( + connection.frame_slots.available_permits(), + state.limits.max_egress_frames + ); + assert_eq!( + connection.byte_slots.available_permits(), + state.limits.max_egress_bytes + ); + assert_eq!( + connection.stream_slots.available_permits(), + state.limits.max_active_streams + ); + assert!(get_inner(state.clone(), get()).is_ok()); + connection.terminate(); + } + + #[tokio::test] + async fn batched_initialize_releases_all_response_routes() { + let (state, mut agent) = registry(ServerLimits::default()); + let task = tokio::spawn(handle_post( + State(state.clone()), + post_request( + r#"[{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}},{"jsonrpc":"2.0","id":2,"method":"ping"}]"#, + None, + ), + )); + let received = timeout(Duration::from_secs(2), agent.rx.next()) + .await + .unwrap() + .unwrap(); + drop(received); + agent + .tx + .try_send_serialized( + r#"[{"jsonrpc":"2.0","id":1,"result":{}},{"jsonrpc":"2.0","id":2,"result":{}}]"#, + ) + .unwrap(); + let response = timeout(Duration::from_secs(2), task) + .await + .unwrap() + .unwrap(); + let id = response.headers()[HEADER_CONNECTION_ID] + .to_str() + .unwrap() + .to_owned(); + let connection = state.connections.lock().unwrap().get(&id).unwrap().clone(); + assert!(connection.inner.lock().unwrap().pending.is_empty()); + drop(response); + } + + #[tokio::test] + async fn duplicate_route_cap_is_atomic_and_unknown_session_gets_are_bounded() { + let (state, mut agent) = registry(ServerLimits { + max_pending_routes: 2, + max_registered_sessions: 1, + ..ServerLimits::default() + }); + let (id, response) = initialize(state.clone(), &mut agent).await; + axum::body::to_bytes(response.into_body(), 1024) + .await + .unwrap(); + let request = r#"[{"jsonrpc":"2.0","id":2,"method":"ping"},{"jsonrpc":"2.0","id":2,"method":"ping"}]"#; + let response = handle_post(State(state.clone()), post_request(request, Some(&id))).await; + assert_eq!(response.status(), StatusCode::ACCEPTED); + let connection = state.connections.lock().unwrap().get(&id).unwrap().clone(); + assert_eq!(connection.inner.lock().unwrap().pending.len(), 2); + let response = handle_post(State(state.clone()), post_request(request, Some(&id))).await; + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert_eq!(connection.inner.lock().unwrap().pending.len(), 2); + let get = |session| { + Request::builder() + .header(HEADER_CONNECTION_ID, &id) + .header(HEADER_SESSION_ID, session) + .body(Body::empty()) + .unwrap() + }; + drop(get_inner(state.clone(), get("one")).unwrap()); + assert_eq!( + get_inner(state.clone(), get("two")).unwrap_err(), + StatusCode::TOO_MANY_REQUESTS + ); + connection.terminate(); + assert!(connection.inner.lock().unwrap().pending.is_empty()); + } + + #[tokio::test] + async fn oversized_unknown_length_body_and_batch_restore_admission() { + let calls = Arc::new(AtomicUsize::new(0)); + let factory_calls = calls.clone(); + let state = BoundedAcpHttpServer::new( + move || -> BoundedChannel { + factory_calls.fetch_add(1, Ordering::SeqCst); + panic!("oversized bodies must not reach the factory"); + }, + ServerLimits { + max_frame_bytes: 128, + max_batch_entries: 1, + ..ServerLimits::default() + }, + ) + .unwrap() + .state; + let body = Body::from_stream(futures::stream::iter([Ok::<_, Infallible>(vec![ + b' '; + 129 + ])])); + assert_eq!( + handle_post(State(state.clone()), post_request(body, None)) + .await + .status(), + StatusCode::PAYLOAD_TOO_LARGE + ); + let batch = r#"[{"jsonrpc":"2.0","id":1,"method":"initialize"},{"jsonrpc":"2.0","id":2,"method":"ping"}]"#; + assert_eq!( + handle_post(State(state.clone()), post_request(batch, None)) + .await + .status(), + StatusCode::PAYLOAD_TOO_LARGE + ); + assert_eq!(calls.load(Ordering::SeqCst), 0); + assert_eq!( + state.posts.available_permits(), + state.limits.max_in_flight_posts + ); + assert_eq!( + state.bodies.available_permits(), + state.limits.max_body_bytes + ); + assert_eq!( + state.connection_slots.available_permits(), + state.limits.max_connections + ); + } + + #[tokio::test] + async fn accepted_egress_overflow_terminates_and_revokes_escaped_sender() { + let (state, mut agent) = registry(ServerLimits { + max_egress_frames: 1, + ..ServerLimits::default() + }); + let (id, response) = initialize(state.clone(), &mut agent).await; + axum::body::to_bytes(response.into_body(), 1024) + .await + .unwrap(); + let connection = state.connections.lock().unwrap().get(&id).unwrap().clone(); + for _ in 0..2 { + agent + .tx + .try_send_serialized(r#"{"jsonrpc":"2.0","method":"notice"}"#) + .unwrap(); + } + for _ in 0..10 { + tokio::task::yield_now().await; + } + assert!(state.connections.lock().unwrap().is_empty()); + assert!( + agent + .tx + .try_send_serialized(r#"{"jsonrpc":"2.0","method":"notice"}"#) + .is_err() + ); + assert_eq!( + connection.frame_slots.available_permits(), + state.limits.max_egress_frames + ); + assert_eq!( + connection.byte_slots.available_permits(), + state.limits.max_egress_bytes + ); + assert!(connection.inner.lock().unwrap().streams.is_empty()); + } + + #[tokio::test] + async fn response_body_saturation_terminates_instead_of_deadlocking_callback() { + let (state, mut agent) = registry(ServerLimits { + max_in_flight_posts: 1, + ..ServerLimits::default() + }); + let (id, response) = initialize(state.clone(), &mut agent).await; + axum::body::to_bytes(response.into_body(), 1024) + .await + .unwrap(); + let _held = state.posts.clone().try_acquire_owned().unwrap(); + let response = handle_post( + State(state.clone()), + post_request(r#"{"jsonrpc":"2.0","id":2,"result":{}}"#, Some(&id)), + ) + .await; + assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS); + assert!(state.connections.lock().unwrap().is_empty()); + assert!( + agent + .tx + .try_send_serialized(r#"{"jsonrpc":"2.0","method":"notice"}"#) + .is_err() + ); + } + + #[tokio::test] + async fn legacy_component_fails_closed_without_a_bridge() { + let state = BoundedAcpHttpServer::new( + || { + let (channel, peer) = agent_client_protocol::Channel::duplex(); + drop(peer); + channel + }, + ServerLimits::default(), + ) + .unwrap() + .state; + let response = timeout( + Duration::from_secs(2), + handle_post( + State(state.clone()), + post_request(r#"{"jsonrpc":"2.0","id":1,"method":"initialize"}"#, None), + ), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR); + assert!(state.connections.lock().unwrap().is_empty()); + } + + #[tokio::test] + async fn bounded_router_rejects_websocket_upgrade() { + let (state, _agent) = registry(ServerLimits::default()); + let request = Request::builder() + .header(header::UPGRADE, "websocket") + .body(Body::empty()) + .unwrap(); + assert_eq!( + get_inner(state, request).unwrap_err(), + StatusCode::NOT_IMPLEMENTED + ); + } +} diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 1c899872..74457b71 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -32,6 +32,10 @@ pub enum HttpClientError { Reqwest(#[from] reqwest::Error), } +#[path = "client_limits.rs"] +mod limits; +pub use limits::{BoundedHttpClient, HttpClientLimits}; + pub struct HttpClient { endpoint: url::Url, http: reqwest::Client, @@ -95,6 +99,14 @@ impl HttpClient { Ok(Self { endpoint, http }) } + /// Opt into fail-fast, bounded HTTP transport and core channel admission. + /// + /// Unlike the legacy constructors, this rejects WebSocket endpoints. No POST + /// retries or SSE replay are performed. See [`HttpClientLimits`] for scope. + pub fn with_limits(self, limits: HttpClientLimits) -> Result { + BoundedHttpClient::new(self, limits) + } + fn is_websocket(&self) -> bool { matches!(self.endpoint.scheme(), "ws" | "wss") } diff --git a/src/agent-client-protocol-http/src/client_limits.rs b/src/agent-client-protocol-http/src/client_limits.rs new file mode 100644 index 00000000..2085472a --- /dev/null +++ b/src/agent-client-protocol-http/src/client_limits.rs @@ -0,0 +1,1408 @@ +//! Separate opt-in HTTP path: never converts charged frames into legacy channels. +use super::*; +use agent_client_protocol::{BoundedChannel, ChannelLimits, ChargedFrame, TransportChannel}; +use std::task::Poll; + +/// Finite per-client limits. Every value must be nonzero. +/// +/// Byte accounting measures serialized wire bytes, not allocator overhead or +/// reqwest/TLS/socket buffers. The HTTP budget covers retained POST bodies and +/// pending metadata, capped response workspaces, and fixed SSE parser workspaces. +/// Core channel budgets apply separately in each direction, including producer +/// admission. Exhaustion is terminal and never waits for capacity. A full POST +/// lane fails rather than parking on the shared outgoing FIFO. Response-only +/// callback frames have an independent lane; mixed batches retain request order. +#[derive(Clone, Debug)] +pub struct HttpClientLimits { + /// Independent bounded core limits in each channel direction. + pub channel: ChannelLimits, + /// HTTP wire-byte reservations, including conservative metadata allowances. + pub max_buffered_bytes: usize, + /// Simultaneous HTTP reservations. Bodies, parser workspaces and metadata + /// each consume a slot; this is conservative relative to actual wire frames. + pub max_buffered_frames: usize, + /// Individual RPC entries awaiting a reply, including duplicate/null IDs. + pub max_pending_requests: usize, + /// Queued plus active request/mixed/notification POSTs (one active at a time). + pub max_request_posts: usize, + /// Independent queued plus active response-only POSTs (one active at a time). + pub max_response_posts: usize, + /// Includes the connection stream and establishing session streams. + pub max_sse_streams: usize, + /// Collected response bytes, for initialization and success/error POSTs. + pub max_response_bytes: usize, + /// Raw bytes in an incomplete SSE line, including comment/ignored lines. + pub max_sse_line_bytes: usize, + /// Includes comments, ignored fields, and delimiters, not only data fields. + /// CRLF is normalized to one delimiter byte for this limit. + pub max_sse_event_bytes: usize, + /// Largest accepted reqwest SSE chunk before SDK copying or parsing. + pub max_sse_chunk_bytes: usize, +} + +impl Default for HttpClientLimits { + fn default() -> Self { + Self { + channel: ChannelLimits::default(), + max_buffered_bytes: 64 * 1024 * 1024, + max_buffered_frames: 128, + max_pending_requests: 128, + max_request_posts: 32, + max_response_posts: 32, + max_sse_streams: 8, + max_response_bytes: 1024 * 1024, + max_sse_line_bytes: 1024 * 1024, + max_sse_event_bytes: 1024 * 1024, + max_sse_chunk_bytes: 1024 * 1024, + } + } +} + +fn failure(message: impl Into) -> AcpError { + AcpError::internal_error().data(format!("bounded HTTP: {}", message.into())) +} + +impl HttpClientLimits { + fn workspace(&self) -> Result { + self.max_sse_line_bytes + .checked_add(self.max_sse_event_bytes) + .and_then(|n| n.checked_add(self.max_sse_chunk_bytes)) + .ok_or_else(|| failure("SSE workspace arithmetic overflow")) + } + + fn validate(&self) -> Result<(), AcpError> { + if [ + self.max_buffered_bytes, + self.max_buffered_frames, + self.max_pending_requests, + self.max_request_posts, + self.max_response_posts, + self.max_sse_streams, + self.max_response_bytes, + self.max_sse_line_bytes, + self.max_sse_event_bytes, + self.max_sse_chunk_bytes, + ] + .contains(&0) + { + return Err(failure("limits must be nonzero")); + } + if self.workspace()? > self.max_buffered_bytes + || self.max_response_bytes > self.max_buffered_bytes + { + return Err(failure("workspace exceeds HTTP byte budget")); + } + Ok(()) + } +} + +/// An HTTP-only transport with bounded admission. Construct with +/// [`HttpClient::with_limits`]; legacy channel extraction fails closed. +/// +/// There are no background SSE tasks, observer mailboxes, or hidden unbounded +/// bridges. The transport future directly polls all streams and POSTs. Dropping +/// it releases local resources; it does not wait for an HTTP DELETE or promise +/// peer cleanup. No uncertain or accepted POST is ever retried. +pub struct BoundedHttpClient { + client: HttpClient, + limits: HttpClientLimits, + caller: BoundedChannel, + transport: BoundedChannel, +} + +impl std::fmt::Debug for BoundedHttpClient { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BoundedHttpClient") + .field("client", &self.client) + .field("limits", &self.limits) + .finish_non_exhaustive() + } +} + +impl BoundedHttpClient { + pub(super) fn new(client: HttpClient, limits: HttpClientLimits) -> Result { + limits.validate()?; + if !matches!(client.endpoint.scheme(), "http" | "https") { + return Err(failure("bounded client supports only http/https")); + } + let (caller, transport) = BoundedChannel::duplex(limits.channel)?; + Ok(Self { + client, + limits, + caller, + transport, + }) + } + + /// Extract the charged channel without a legacy adapter. Poll the returned + /// future to drive HTTP; dropping it cancels all local HTTP work. + #[must_use] + pub fn into_bounded_channel_and_future( + self, + ) -> (BoundedChannel, BoxFuture<'static, Result<(), AcpError>>) { + ( + self.caller, + run_bounded(self.client, self.limits, self.transport).boxed(), + ) + } +} + +impl ConnectTo for BoundedHttpClient { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), AcpError> { + let (channel, transport) = self.into_bounded_channel_and_future(); + let shutdown = channel.tx.clone(); + let application = client.connect_to(channel).boxed(); + match futures::future::select(application, transport).await { + futures::future::Either::Left((result, transport)) => { + result?; + shutdown.close_channel(); + transport.await + } + futures::future::Either::Right((result, application)) => { + result?; + // A final admitted initialization rejection must reach the + // application before its closed inbound queue is discarded. + application.await + } + } + } + + fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), AcpError>>) { + let (caller, other) = Channel::duplex(); + drop(other); + drop(self); + ( + caller, + async { + Err(failure( + "legacy channel extraction is disabled; use bounded transport extraction", + )) + } + .boxed(), + ) + } + + fn into_bounded_channel_and_future( + self, + limits: ChannelLimits, + ) -> Result<(BoundedChannel, BoxFuture<'static, Result<(), AcpError>>), AcpError> { + let limits = limits.validate()?; + if self.limits.channel != limits { + return Err(self + .caller + .tx + .fail("bounded HTTP client limits differ from requested limits")); + } + Ok(BoundedHttpClient::into_bounded_channel_and_future(self)) + } + + fn into_transport_and_future( + self, + ) -> (TransportChannel, BoxFuture<'static, Result<(), AcpError>>) { + let (channel, future) = self.into_bounded_channel_and_future(); + (TransportChannel::Bounded(channel), future) + } +} + +#[derive(Clone)] +struct Budget(Arc>); +struct Usage { + bytes: usize, + frames: usize, + max_bytes: usize, + max_frames: usize, +} +struct Lease { + budget: Budget, + bytes: usize, +} +impl Budget { + fn new(limits: &HttpClientLimits) -> Self { + Self(Arc::new(StdMutex::new(Usage { + bytes: 0, + frames: 0, + max_bytes: limits.max_buffered_bytes, + max_frames: limits.max_buffered_frames, + }))) + } + fn reserve(&self, bytes: usize) -> Result, AcpError> { + let mut usage = self.0.lock().expect("budget mutex poisoned"); + if bytes > usage.max_bytes.saturating_sub(usage.bytes) || usage.frames >= usage.max_frames { + return Err(failure("HTTP aggregate byte/frame limit exhausted")); + } + usage.bytes += bytes; + usage.frames += 1; + Ok(Arc::new(Lease { + budget: self.clone(), + bytes, + })) + } +} +impl Drop for Lease { + fn drop(&mut self) { + let mut usage = self.budget.0.lock().expect("budget mutex poisoned"); + usage.bytes -= self.bytes; + usage.frames -= 1; + } +} + +// Streaming bodies have no clone/replay implementation, preventing reqwest +// retry policies and 307/308 redirects from resending these POSTs. +fn non_replayable_body(bytes: Vec) -> reqwest::Body { + reqwest::Body::wrap_stream(futures::stream::once(async move { + Ok::<_, std::io::Error>(bytes) + })) +} + +async fn capped_body(mut response: reqwest::Response, cap: usize) -> Result, AcpError> { + if response.content_length().is_some_and(|n| n > cap as u64) { + return Err(failure("HTTP response body exceeds limit")); + } + let mut body = Vec::new(); + while let Some(chunk) = response.chunk().await.map_err(|e| failure(e.to_string()))? { + if chunk.len() > cap.saturating_sub(body.len()) { + return Err(failure("HTTP response body exceeds limit")); + } + body.extend_from_slice(&chunk); + } + Ok(body) +} + +/// Raw bytes are checked before UTF-8 conversion. CR, LF, CRLF, split UTF-8, +/// a leading BOM, ignored fields and comments all share bounded state. +struct Parser { + line: Vec, + data: Vec, + event_bytes: usize, + after_cr: bool, + first_line: bool, + has_data: bool, + max_line: usize, + max_event: usize, +} +impl Parser { + fn new(limits: &HttpClientLimits) -> Self { + Self { + line: Vec::new(), + data: Vec::new(), + event_bytes: 0, + after_cr: false, + first_line: true, + has_data: false, + max_line: limits.max_sse_line_bytes, + max_event: limits.max_sse_event_bytes, + } + } + fn byte(&mut self, byte: u8) -> Result>, AcpError> { + if self.after_cr { + self.after_cr = false; + if byte == b'\n' { + return Ok(None); + } + } + if self.event_bytes == self.max_event { + return Err(failure("SSE event limit exhausted")); + } + self.event_bytes += 1; + if byte != b'\r' && byte != b'\n' { + if self.line.len() == self.max_line { + return Err(failure("SSE line limit exhausted")); + } + self.line.push(byte); + return Ok(None); + } + self.after_cr = byte == b'\r'; + let line = std::str::from_utf8(&self.line).map_err(|_| failure("invalid SSE UTF-8"))?; + let line = if self.first_line { + line.strip_prefix('\u{feff}').unwrap_or(line) + } else { + line + }; + self.first_line = false; + if line.is_empty() { + self.line.clear(); + self.event_bytes = 0; + if self.has_data { + self.has_data = false; + self.data.pop(); // final data-field newline + return Ok(Some(std::mem::take(&mut self.data))); + } + return Ok(None); + } + let (field, value) = line.split_once(':').unwrap_or((line, "")); + if field == "data" { + let value = value.strip_prefix(' ').unwrap_or(value); + if value.len().saturating_add(1) > self.max_event.saturating_sub(self.data.len()) { + return Err(failure("SSE data limit exhausted")); + } + self.data.extend_from_slice(value.as_bytes()); + self.data.push(b'\n'); + self.has_data = true; + } + self.line.clear(); + Ok(None) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicUsize, Ordering}; + + fn parser_limits(line: usize, event: usize) -> HttpClientLimits { + HttpClientLimits { + max_sse_line_bytes: line, + max_sse_event_bytes: event, + ..Default::default() + } + } + + #[test] + fn parser_handles_split_utf8_bom_and_crlf_without_replay() { + let mut parser = Parser::new(&parser_limits(64, 128)); + let mut events = Vec::new(); + for byte in "\u{feff}data: hé\r\nid: ignored\r\n\r\ndata: next\n\n".bytes() { + if let Some(event) = parser.byte(byte).unwrap() { + events.push(event); + } + } + assert_eq!(events, vec!["hé".as_bytes(), b"next"]); + } + + #[test] + fn parser_caps_partial_lines_comments_events_and_utf8() { + let mut parser = Parser::new(&parser_limits(4, 32)); + for byte in b":abc" { + parser.byte(*byte).unwrap(); + } + assert!(parser.byte(b'd').is_err()); + let mut parser = Parser::new(&parser_limits(32, 8)); + for byte in b":a\n:b\n:c" { + parser.byte(*byte).unwrap(); + } + assert!(parser.byte(b'\n').is_err()); + let mut parser = Parser::new(&parser_limits(32, 32)); + parser.byte(0xc3).unwrap(); + assert!(parser.byte(b'\n').is_err()); + } + + #[test] + fn parser_exact_event_boundary_and_many_events() { + let mut parser = Parser::new(&parser_limits(7, 9)); + for _ in 0..1024 { + let mut event = None; + for byte in b"data: a\n\n" { + event = parser.byte(*byte).unwrap().or(event); + } + assert_eq!(event.unwrap(), b"a"); + } + let mut crlf = Parser::new(&parser_limits(7, 9)); + let mut event = None; + for byte in b"data: a\r\n\r\n" { + event = crlf.byte(*byte).unwrap().or(event); + } + assert_eq!(event.unwrap(), b"a"); + let mut parser = Parser::new(&parser_limits(7, 8)); + for byte in b"data: a\n" { + parser.byte(*byte).unwrap(); + } + assert!(parser.byte(b'\n').is_err()); + } + + #[test] + fn leases_survive_handoffs_and_release_exactly_once() { + let limits = HttpClientLimits { + max_buffered_bytes: 10, + max_buffered_frames: 1, + ..Default::default() + }; + let budget = Budget::new(&limits); + let lease = budget.reserve(10).unwrap(); + let pending_entry = lease.clone(); + drop(lease); + assert!(budget.reserve(1).is_err()); + drop(pending_entry); + let restored = budget.reserve(10).unwrap(); + assert!(budget.reserve(0).is_err()); + drop(restored); + assert_eq!(budget.0.lock().unwrap().bytes, 0); + assert_eq!(budget.0.lock().unwrap().frames, 0); + } + + #[tokio::test] + async fn success_and_error_bodies_share_the_same_cap() { + for status in [200, 500] { + let response = axum::http::Response::builder() + .status(status) + .body("1234") + .unwrap(); + assert_eq!(capped_body(response.into(), 4).await.unwrap(), b"1234"); + let response = axum::http::Response::builder() + .status(status) + .body("12345") + .unwrap(); + assert!(capped_body(response.into(), 4).await.is_err()); + } + } + + #[test] + fn bounded_extraction_validates_requested_limits() { + let configured = ChannelLimits::default(); + let client = || { + HttpClient::new("http://localhost") + .unwrap() + .with_limits(HttpClientLimits { + channel: configured, + ..Default::default() + }) + .unwrap() + }; + let (channel, future) = + ConnectTo::::into_bounded_channel_and_future(client(), configured).unwrap(); + assert_eq!(channel.limits(), configured); + drop((channel, future)); + let different = ChannelLimits { + max_pending_requests: configured.max_pending_requests + 1, + ..configured + }; + assert!(ConnectTo::::into_bounded_channel_and_future(client(), different).is_err()); + let zero = ChannelLimits { + max_buffered_frames: 0, + ..configured + }; + assert!(ConnectTo::::into_bounded_channel_and_future(client(), zero).is_err()); + } + + #[cfg(feature = "server")] + #[tokio::test] + async fn bounded_server_rejects_http_factory_with_different_core_limits() { + use tower::ServiceExt; + let calls = Arc::new(AtomicUsize::new(0)); + let factory_calls = calls.clone(); + let limits = crate::ServerLimits { + channel_limits: ChannelLimits { + max_pending_requests: 1, + ..Default::default() + }, + ..Default::default() + }; + let router = crate::AcpHttpServer::new_bounded( + move || { + factory_calls.fetch_add(1, Ordering::SeqCst); + HttpClient::new("http://localhost:9") + .unwrap() + .with_limits(HttpClientLimits::default()) + .unwrap() + }, + limits, + ) + .unwrap() + .into_router(); + let request = axum::http::Request::builder() + .method("POST") + .uri("/acp") + .header("Content-Type", "application/json") + .body(axum::body::Body::from(initialize().to_json().unwrap())) + .unwrap(); + let response = router.oneshot(request).await.unwrap(); + assert_eq!( + response.status(), + axum::http::StatusCode::INTERNAL_SERVER_ERROR + ); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + #[test] + fn invalid_limits_and_websockets_are_rejected() { + let limits = HttpClientLimits { + max_request_posts: 0, + ..Default::default() + }; + assert!( + HttpClient::new("http://localhost") + .unwrap() + .with_limits(limits) + .is_err() + ); + assert!( + HttpClient::new("ws://localhost") + .unwrap() + .with_limits(HttpClientLimits::default()) + .is_err() + ); + let limits = HttpClientLimits { + max_sse_line_bytes: usize::MAX, + ..Default::default() + }; + assert!( + HttpClient::new("http://localhost") + .unwrap() + .with_limits(limits) + .is_err() + ); + } + + async fn fixture() -> (String, Arc, tokio::task::JoinHandle<()>) { + use axum::{Router, body::Body, response::Response, routing::post}; + let count = Arc::new(AtomicUsize::new(0)); + let counter = count.clone(); + let app = Router::new().route( + "/acp", + post(move || { + let counter = counter.clone(); + async move { + if counter.fetch_add(1, Ordering::SeqCst) == 0 { + Response::builder() + .header(HEADER_CONNECTION_ID, "bounded-test") + .body(Body::from(r#"{"jsonrpc":"2.0","id":0,"result":{}}"#)) + .unwrap() + } else { + Response::builder().status(202).body(Body::empty()).unwrap() + } + } + }) + .get(|| async { + Response::builder() + .header("Content-Type", "text/event-stream") + .body(Body::from_stream(futures::stream::pending::< + Result, + >())) + .unwrap() + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + (format!("http://{address}/acp"), count, server) + } + + #[tokio::test] + async fn non_replayable_posts_ignore_redirect_and_retry_policies() { + use axum::{Router, body::Body, response::Response, routing::post}; + for status in [307, 308, 503] { + let count = Arc::new(AtomicUsize::new(0)); + let counter = count.clone(); + let app = Router::new().route( + "/", + post(move || { + let counter = counter.clone(); + async move { + counter.fetch_add(1, Ordering::SeqCst); + Response::builder() + .status(status) + .header("Location", "/") + .body(Body::empty()) + .unwrap() + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}/", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let http = reqwest::Client::builder() + .retry( + reqwest::retry::for_host("127.0.0.1") + .no_budget() + .max_retries_per_request(2) + .classify_fn(|response| { + if response.status().is_some_and(|status| { + status.is_server_error() || status.is_redirection() + }) { + response.retryable() + } else { + response.success() + } + }), + ) + .build() + .unwrap(); + let result = http + .post(url) + .body(non_replayable_body(b"request".to_vec())) + .send() + .await; + assert_eq!(result.unwrap().status().as_u16(), status); + assert_eq!(count.load(Ordering::SeqCst), 1); + server.abort(); + } + } + + #[tokio::test] + async fn producer_exhaustion_cancels_stalled_initialize() { + use axum::{Router, response::Response, routing::post}; + let started = Arc::new(tokio::sync::Notify::new()); + let notify = started.clone(); + let app = Router::new().route( + "/acp", + post(move || { + let notify = notify.clone(); + async move { + notify.notify_one(); + futures::future::pending::().await + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}/acp", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let limits = HttpClientLimits { + channel: ChannelLimits { + max_buffered_frames: 1, + ..Default::default() + }, + ..Default::default() + }; + let (channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(limits) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + tokio::time::timeout(std::time::Duration::from_secs(3), started.notified()) + .await + .unwrap(); + assert!(channel.tx.try_send(initialize()).is_err()); + assert!( + tokio::time::timeout(std::time::Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .is_err() + ); + server.abort(); + } + + #[tokio::test] + async fn response_post_bypasses_stalled_request_lane() { + use axum::{ + Router, + body::{Body, Bytes}, + response::Response, + routing::post, + }; + let start_callback = Arc::new(tokio::sync::Notify::new()); + let response_received = Arc::new(tokio::sync::Notify::new()); + let request_finished = Arc::new(tokio::sync::Notify::new()); + let start = start_callback.clone(); + let received = response_received.clone(); + let finished = request_finished.clone(); + let app = Router::new().route( + "/acp", + post(move |body: Bytes| { + let (start, received, finished) = + (start.clone(), received.clone(), finished.clone()); + async move { + let value: serde_json::Value = serde_json::from_slice(&body).unwrap(); + if value["method"] == "initialize" { + return Response::builder() + .header(HEADER_CONNECTION_ID, "callback-test") + .body(Body::from(r#"{"jsonrpc":"2.0","id":0,"result":{}}"#)) + .unwrap(); + } + if value.get("method").is_some() { + start.notify_one(); + received.notified().await; + finished.notify_one(); + } else { + received.notify_one(); + } + Response::builder().status(202).body(Body::empty()).unwrap() + } + }) + .get(move || { + let start = start_callback.clone(); + async move { + let events = futures::stream::once(async move { + start.notified().await; + Ok::<_, std::convert::Infallible>( + "data: {\"jsonrpc\":\"2.0\",\"id\":77,\"method\":\"callback\"}\n\n", + ) + }) + .chain(futures::stream::pending()); + Response::builder() + .header("Content-Type", "text/event-stream") + .body(Body::from_stream(events)) + .unwrap() + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}/acp", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let limits = HttpClientLimits { + max_request_posts: 1, + max_response_posts: 1, + ..Default::default() + }; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(limits) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + assert!( + tokio::time::timeout(std::time::Duration::from_secs(3), channel.rx.next()) + .await + .unwrap() + .is_some() + ); + channel + .tx + .try_send(TransportFrame::parse_json( + r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#, + )) + .unwrap(); + let callback = tokio::time::timeout(std::time::Duration::from_secs(3), channel.rx.next()) + .await + .unwrap() + .unwrap(); + assert!(String::from_utf8_lossy(callback.as_bytes()).contains("callback")); + channel + .tx + .try_send(TransportFrame::parse_json( + r#"{"jsonrpc":"2.0","id":77,"result":{}}"#, + )) + .unwrap(); + tokio::time::timeout( + std::time::Duration::from_secs(3), + request_finished.notified(), + ) + .await + .unwrap(); + driver.abort(); + assert!(driver.await.unwrap_err().is_cancelled()); + assert!(channel.tx.try_send(initialize()).is_err()); + server.abort(); + } + + fn initialize() -> TransportFrame { + TransportFrame::parse_json(r#"{"jsonrpc":"2.0","id":0,"method":"initialize","params":{}}"#) + } + + #[tokio::test] + async fn duplicate_batch_entries_are_rejected_before_post() { + let (url, count, server) = fixture().await; + let limits = HttpClientLimits { + max_pending_requests: 1, + ..Default::default() + }; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(limits) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + assert!( + tokio::time::timeout(std::time::Duration::from_secs(3), channel.rx.next()) + .await + .unwrap() + .is_some() + ); + channel.tx.try_send(TransportFrame::parse_json(r#"[{"jsonrpc":"2.0","id":1,"method":"ping"},{"jsonrpc":"2.0","id":1,"method":"ping"}]"#)).unwrap(); + let error = tokio::time::timeout(std::time::Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .unwrap_err(); + assert!(error.to_string().contains("outstanding RPC")); + assert_eq!(count.load(Ordering::SeqCst), 1); + assert!(channel.tx.try_send(initialize()).is_err()); + server.abort(); + } + + #[tokio::test] + async fn http_acceptance_does_not_release_outstanding_rpc_slot() { + let (url, count, server) = fixture().await; + let limits = HttpClientLimits { + max_pending_requests: 1, + ..Default::default() + }; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(limits) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + assert!( + tokio::time::timeout(std::time::Duration::from_secs(3), channel.rx.next()) + .await + .unwrap() + .is_some() + ); + let request = || TransportFrame::parse_json(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#); + channel.tx.try_send(request()).unwrap(); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while count.load(Ordering::SeqCst) != 2 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + channel.tx.try_send(request()).unwrap(); + assert!( + tokio::time::timeout(std::time::Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .is_err() + ); + assert_eq!(count.load(Ordering::SeqCst), 2); + server.abort(); + } +} + +struct Pending { + id: RequestId, + method: String, + _lease: Arc, +} +struct Post { + frame: ChargedFrame, + sessions: Vec, + session_header: Option, + lease: Arc, + response_workspace: Arc, +} +#[derive(Default)] +struct Lane { + queue: VecDeque, + active: bool, +} +impl Lane { + fn count(&self) -> usize { + self.queue.len() + usize::from(self.active) + } +} +type PostFuture = BoxFuture<'static, Result>; + +fn messages(frame: &TransportFrame) -> Result, AcpError> { + match frame { + TransportFrame::Single(message) => Ok(vec![message]), + TransportFrame::Batch(batch) => batch + .entries() + .map(|entry| match entry { + TransportBatchEntry::Message(message) => Ok(message), + TransportBatchEntry::Malformed { .. } => Err(failure("malformed batch entry")), + }) + .collect(), + TransportFrame::Malformed { .. } => Err(failure("malformed frame")), + } +} + +struct Streams { + registered: HashMap, Arc>, + ready: HashSet>, + futures: FuturesUnordered, +} +impl Streams { + fn new() -> Self { + Self { + registered: HashMap::new(), + ready: HashSet::new(), + futures: FuturesUnordered::new(), + } + } + fn open( + &mut self, + client: &HttpClient, + connection: &str, + key: Option, + limits: &HttpClientLimits, + budget: &Budget, + ) -> Result<(), AcpError> { + if self.registered.contains_key(&key) { + return Ok(()); + } + if self.registered.len() >= limits.max_sse_streams { + return Err(failure("SSE stream limit exhausted")); + } + // Account for the registry, readiness set and live stream's key copies. + let key_bytes = key + .as_ref() + .map_or(0, String::len) + .checked_mul(4) + .and_then(|n| n.checked_add(connection.len())) + .ok_or_else(|| failure("session metadata arithmetic overflow"))?; + let metadata = budget.reserve(key_bytes)?; + let future = open_sse(client, connection, key.clone(), limits, budget)?; + self.registered.insert(key, metadata); + self.futures.push(future); + Ok(()) + } +} + +fn launch( + lane: &mut Lane, + response_lane: bool, + streams: &Streams, + client: &HttpClient, + connection: &str, + cap: usize, + posts: &mut FuturesUnordered, +) { + if lane.active || !streams.ready.contains(&None) { + return; + } + let Some(post) = lane.queue.front() else { + return; + }; + if post.sessions.iter().any(|id| { + !streams + .ready + .iter() + .any(|key| key.as_deref() == Some(id.as_str())) + }) { + return; + } + let post = lane.queue.pop_front().expect("front checked"); + let mut request = client + .http + .post(client.endpoint.clone()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .header(HEADER_CONNECTION_ID, connection); + if let Some(session) = &post.session_header { + request = request.header(HEADER_SESSION_ID, session); + } + lane.active = true; + let expected_url = client.endpoint.clone(); + posts.push( + async move { + // The charged core frame stays alive through reqwest's body handoff. + // Waiting until response headers is conservative; this is NOT peer ACK. + let request = request.body(non_replayable_body(post.frame.as_bytes().to_vec())); + let result = request.send().await.map_err(|e| failure(e.to_string())); + drop(post.frame); + let response = result?; + if response.url() != &expected_url { + return Err(failure("POST redirect is not supported")); + } + let status = response.status(); + drop(capped_body(response, cap).await?); + drop(post.response_workspace); + drop(post.lease); + if !status.is_success() { + return Err(failure(format!("POST HTTP {status}"))); + } + Ok(response_lane) + } + .boxed(), + ); +} + +async fn run_bounded( + client: HttpClient, + limits: HttpClientLimits, + channel: BoundedChannel, +) -> Result<(), AcpError> { + struct Terminate { + sender: agent_client_protocol::BoundedSender, + armed: bool, + } + impl Drop for Terminate { + fn drop(&mut self) { + if self.armed { + drop(self.sender.fail("bounded HTTP transport ended")); + } + } + } + let mut terminal = Terminate { + sender: channel.tx.clone(), + armed: true, + }; + let failed = terminal.sender.failure(); + let running = run_bounded_inner(client, limits, channel).boxed(); + match futures::future::select(failed, running).await { + futures::future::Either::Left((error, _)) => Err(error), + futures::future::Either::Right((result, _)) => { + // Preserve already-admitted inbound responses on graceful EOF. + terminal.armed = result.is_err(); + result + } + } +} + +async fn run_bounded_inner( + client: HttpClient, + limits: HttpClientLimits, + channel: BoundedChannel, +) -> Result<(), AcpError> { + enum Event { + Outgoing(Option), + Sse(Box>), + Post(Result), + } + let BoundedChannel { + tx: incoming, + rx: mut outgoing, + } = channel; + let budget = Budget::new(&limits); + let Some(first) = outgoing.next().await else { + return Ok(()); + }; + let frame = first.decode(); + if !matches!(&frame, TransportFrame::Single(message) if is_initialize_request(message)) { + return Err(failure("first frame must be a single initialize request")); + } + let init_bytes = first + .as_bytes() + .len() + .checked_mul(2) + .and_then(|n| n.checked_add(limits.max_response_bytes)) + .ok_or_else(|| failure("initialize budget arithmetic overflow"))?; + let initialization = budget.reserve(init_bytes)?; + let response = client + .http + .post(client.endpoint.clone()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .body(non_replayable_body(first.as_bytes().to_vec())) + .send() + .await + .map_err(|e| failure(e.to_string()))?; + drop(first); + let expected_id = match &frame { + TransportFrame::Single(RawJsonRpcMessage::Request(request)) => request.id.clone(), + _ => unreachable!("initialize shape checked"), + }; + drop(frame); + if response.url() != &client.endpoint { + return Err(failure("initialize redirect is not supported")); + } + let status = response.status(); + let _connection_metadata = response + .headers() + .get(HEADER_CONNECTION_ID) + .map(|header| budget.reserve(header.as_bytes().len())) + .transpose()?; + let connection = response + .headers() + .get(HEADER_CONNECTION_ID) + .map(|header| { + if header.as_bytes().len() > limits.channel.max_frame_bytes { + return Err(failure("connection header exceeds frame limit")); + } + header + .to_str() + .map(String::from) + .map_err(|_| failure("invalid connection header")) + }) + .transpose()?; + let body = capped_body(response, limits.max_response_bytes).await?; + if !status.is_success() { + return Err(failure(format!("initialize HTTP {status}"))); + } + let text = std::str::from_utf8(&body).map_err(|_| failure("invalid initialize UTF-8"))?; + let reply = TransportFrame::parse_json(text); + let rejected = match &reply { + TransportFrame::Single(RawJsonRpcMessage::Response(RpcResponse::Error { .. })) => true, + TransportFrame::Single(RawJsonRpcMessage::Response(_)) => false, + _ => return Err(failure("initialize response must be one JSON-RPC response")), + }; + if !matches!(&reply, TransportFrame::Single(message) if message.response_id() == Some(&expected_id)) + { + return Err(failure("initialize response ID mismatch")); + } + incoming.try_send(reply)?; + drop(expected_id); + drop(body); + drop(initialization); + if rejected { + return Ok(()); + } + let connection = connection.ok_or_else(|| failure("missing connection ID"))?; + let mut streams = Streams::new(); + streams.open(&client, &connection, None, &limits, &budget)?; + let mut pending = VecDeque::::new(); + let mut ordered = Lane::default(); + let mut responses = Lane::default(); + let mut posts = FuturesUnordered::::new(); + let mut closed = false; + loop { + launch( + &mut ordered, + false, + &streams, + &client, + &connection, + limits.max_response_bytes, + &mut posts, + ); + launch( + &mut responses, + true, + &streams, + &client, + &connection, + limits.max_response_bytes, + &mut posts, + ); + if closed && ordered.count() == 0 && responses.count() == 0 && pending.is_empty() { + return Ok(()); + } + let event = { + let next_outgoing = async { + if closed { + futures::future::pending().await + } else { + outgoing.next().await + } + } + .fuse(); + let next_sse = async { + match streams.futures.next().await { + Some(event) => event, + None => futures::future::pending().await, + } + } + .fuse(); + let next_post = async { + match posts.next().await { + Some(event) => event, + None => futures::future::pending().await, + } + } + .fuse(); + pin_mut!(next_outgoing, next_sse, next_post); + futures::select! { + frame = next_outgoing => Event::Outgoing(frame), + event = next_sse => Event::Sse(Box::new(event)), + event = next_post => Event::Post(event), + } + }; + match event { + Event::Outgoing(None) => closed = true, + Event::Outgoing(Some(charged)) => { + let frame = charged.decode(); + let is_response = is_response_only_frame(&frame); + let lane = if is_response { + &mut responses + } else { + &mut ordered + }; + let cap = if is_response { + limits.max_response_posts + } else { + limits.max_request_posts + }; + if lane.count() >= cap { + return Err(failure("queued plus active POST limit exhausted")); + } + let entries = messages(&frame)?; + let requests = entries + .iter() + .filter(|message| matches!(message, RawJsonRpcMessage::Request(_))) + .count(); + if requests > limits.max_pending_requests.saturating_sub(pending.len()) { + return Err(failure("outstanding RPC entry limit exhausted")); + } + let bookkeeping = FrameBookkeeping::for_frame(&frame).map_err(failure)?; + // Body, pending ID/method, session list and header copies + // can coexist; reserve four body-sized wire allowances. + let retained = charged + .as_bytes() + .len() + .checked_mul(4) + .and_then(|n| n.checked_add(connection.len())) + .ok_or_else(|| failure("POST metadata arithmetic overflow"))?; + let lease = budget.reserve(retained)?; + let workspace = budget.reserve(limits.max_response_bytes)?; + let session_header = match &frame { + TransportFrame::Single(message) => { + validated_session_id(message).map_err(failure)? + } + _ => None, + }; + // Every entry, including duplicate and null IDs, consumes a slot. + // The body-sized lease remains until all entries complete, not + // merely until the HTTP POST is accepted. + for message in entries { + if let RawJsonRpcMessage::Request(request) = message { + pending.push_back(Pending { + id: request.id.clone(), + method: request.method.to_string(), + _lease: lease.clone(), + }); + } + } + for session in &bookkeeping.session_ids { + streams.open( + &client, + &connection, + Some(session.clone()), + &limits, + &budget, + )?; + } + lane.queue.push_back(Post { + frame: charged, + sessions: bookkeeping.session_ids, + session_header, + lease, + response_workspace: workspace, + }); + } + Event::Post(result) => { + if result? { + responses.active = false; + } else { + ordered.active = false; + } + } + Event::Sse(result) => match (*result)? { + SseEvent::Ready(stream) => { + streams.ready.insert(stream.key.clone()); + streams.futures.push(stream.next()); + } + SseEvent::Frame(stream, bytes) => { + if bytes.len() > limits.channel.max_frame_bytes { + return Err(failure("SSE frame exceeds core frame limit")); + } + let text = + std::str::from_utf8(&bytes).map_err(|_| failure("invalid SSE UTF-8"))?; + let frame = TransportFrame::parse_json(text); + for message in messages(&frame)? { + let RawJsonRpcMessage::Response(response) = message else { + continue; + }; + let Some(id) = message.response_id() else { + continue; + }; + if let Some(index) = pending.iter().position(|entry| &entry.id == id) { + let entry = pending.remove(index).expect("position checked"); + if is_session_opening_method(&entry.method) + && let RpcResponse::Result { result, .. } = response + && let Some(session) = + result.get("sessionId").and_then(|v| v.as_str()) + { + streams.open( + &client, + &connection, + Some(session.to_owned()), + &limits, + &budget, + )?; + } + } + } + incoming.try_send(frame)?; + streams.futures.push(stream.next()); + } + }, + } + } +} + +struct Sse { + key: Option, + response: reqwest::Response, + parser: Parser, + chunk: Vec, + offset: usize, + max_chunk: usize, + _workspace: Arc, +} +enum SseEvent { + Ready(Sse), + Frame(Sse, Vec), +} +type SseFuture = BoxFuture<'static, Result>; + +fn open_sse( + client: &HttpClient, + connection: &str, + key: Option, + limits: &HttpClientLimits, + budget: &Budget, +) -> Result { + let workspace = budget.reserve(limits.workspace()?)?; + let mut request = client + .http + .get(client.endpoint.clone()) + .header("Accept", "text/event-stream") + .header(HEADER_CONNECTION_ID, connection); + if let Some(key) = &key { + request = request.header(HEADER_SESSION_ID, key); + } + let limits = limits.clone(); + Ok(async move { + let response = request.send().await.map_err(|e| failure(e.to_string()))?; + if !response.status().is_success() { + // Error text is deliberately not collected, so an endless error body + // cannot allocate or delay teardown. + return Err(failure(format!("SSE HTTP {}", response.status()))); + } + Ok(SseEvent::Ready(Sse { + key, + response, + parser: Parser::new(&limits), + chunk: Vec::new(), + offset: 0, + max_chunk: limits.max_sse_chunk_bytes, + _workspace: workspace, + })) + } + .boxed()) +} +impl Sse { + fn next(mut self) -> SseFuture { + async move { + let mut work = 0; + loop { + if self.offset == self.chunk.len() { + self.chunk.clear(); + self.offset = 0; + let chunk = self + .response + .chunk() + .await + .map_err(|e| failure(e.to_string()))? + .ok_or_else(|| failure("SSE ended; resume/replay is disabled"))?; + if chunk.len() > self.max_chunk { + return Err(failure("SSE chunk limit exhausted")); + } + self.chunk.extend_from_slice(&chunk); + } + while self.offset < self.chunk.len() { + let byte = self.chunk[self.offset]; + self.offset += 1; + if let Some(data) = self.parser.byte(byte)? + && !data.is_empty() + { + return Ok(SseEvent::Frame(self, data)); + } + work += 1; + if work == 8192 { + // Cooperative yielding even for a continuously ready + // comment-only stream, without a spawned observer task. + let mut yielded = false; + futures::future::poll_fn(|cx| { + if yielded { + Poll::Ready(()) + } else { + yielded = true; + cx.waker().wake_by_ref(); + Poll::Pending + } + }) + .await; + work = 0; + } + } + } + } + .boxed() + } +} diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 525b8e7c..fcac0df1 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -183,7 +183,9 @@ pub(crate) async fn handle_post( StatusCode::ACCEPTED.into_response() } -fn initial_initialize_request(frame: &TransportFrame) -> Option<(&RequestId, Option)> { +pub(crate) fn initial_initialize_request( + frame: &TransportFrame, +) -> Option<(&RequestId, Option)> { fn initialize_id(message: &RawJsonRpcMessage) -> Option<&RequestId> { if !is_initialize_request(message) { return None; @@ -213,7 +215,10 @@ fn initial_initialize_request(frame: &TransportFrame) -> Option<(&RequestId, Opt } } -fn initialize_response_failed(frame: &TransportFrame, initialize_id: &RequestId) -> Option { +pub(crate) fn initialize_response_failed( + frame: &TransportFrame, + initialize_id: &RequestId, +) -> Option { fn response_failed(message: &RawJsonRpcMessage, initialize_id: &RequestId) -> Option { match message { RawJsonRpcMessage::Response(RpcResponse::Result { id, .. }) if id == initialize_id => { @@ -238,7 +243,7 @@ fn initialize_response_failed(frame: &TransportFrame, initialize_id: &RequestId) } } -fn prepare_message_route( +pub(crate) fn prepare_message_route( message: &mut RawJsonRpcMessage, session_id: Option<&str>, ) -> Result, &'static str> { @@ -260,7 +265,7 @@ fn prepare_message_route( }) } -fn collect_route( +pub(crate) fn collect_route( message: &RawJsonRpcMessage, route: Option, session_routes: &mut Vec, diff --git a/src/agent-client-protocol-http/src/lib.rs b/src/agent-client-protocol-http/src/lib.rs index 854c12a5..3e8f0ab8 100644 --- a/src/agent-client-protocol-http/src/lib.rs +++ b/src/agent-client-protocol-http/src/lib.rs @@ -14,6 +14,9 @@ mod server; mod websocket_server; #[cfg(feature = "client")] -pub use client::{HttpClient, HttpClientError}; +pub use client::{BoundedHttpClient, HttpClient, HttpClientError, HttpClientLimits}; #[cfg(feature = "server")] -pub use server::{AcpHttpServer, CorsOptions, ServerOptions}; +pub use server::{ + AcpHttpServer, BoundedAcpHttpServer, CorsOptions, ServerLimits, ServerLimitsError, + ServerOptions, +}; diff --git a/src/agent-client-protocol-http/src/server.rs b/src/agent-client-protocol-http/src/server.rs index 695b1699..0ef8ba19 100644 --- a/src/agent-client-protocol-http/src/server.rs +++ b/src/agent-client-protocol-http/src/server.rs @@ -13,6 +13,10 @@ use tower_http::cors::{AllowOrigin, CorsLayer}; use crate::connection::ConnectionRegistry; +#[path = "bounded_server.rs"] +mod bounded; +pub use bounded::{BoundedAcpHttpServer, ServerLimits, ServerLimitsError}; + #[derive(Debug, Clone)] pub struct ServerOptions { pub path: String, @@ -101,6 +105,19 @@ impl std::fmt::Debug for AcpHttpServer { } impl AcpHttpServer { + /// Opt into finite HTTP/SSE admission using a bounded core transport factory. + /// WebSocket upgrades are rejected on this path; legacy constructors are unchanged. + pub fn new_bounded( + factory: F, + limits: ServerLimits, + ) -> Result + where + F: Fn() -> C + Send + Sync + 'static, + C: ConnectTo, + { + BoundedAcpHttpServer::new(factory, limits) + } + pub fn new(factory: F) -> Self where F: Fn() -> C + Send + Sync + 'static, diff --git a/src/agent-client-protocol-http/tests/bounded_interop.rs b/src/agent-client-protocol-http/tests/bounded_interop.rs new file mode 100644 index 00000000..663f0d9a --- /dev/null +++ b/src/agent-client-protocol-http/tests/bounded_interop.rs @@ -0,0 +1,329 @@ +#![cfg(all(feature = "client", feature = "server"))] + +//! Public-API interoperability over real loopback HTTP, including bounded core actors. +use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, +}; +use std::time::Duration; + +use agent_client_protocol::schema::{ + ProtocolVersion, + v1::{ + ContentBlock, ContentChunk, InitializeRequest, InitializeResponse, NewSessionRequest, + NewSessionResponse, PromptRequest, PromptResponse, SessionNotification, SessionUpdate, + StopReason, TextContent, + }, +}; +use agent_client_protocol::{Agent, ChannelLimits, Client, ConnectTo}; +use agent_client_protocol_http::{AcpHttpServer, HttpClient, HttpClientLimits, ServerLimits}; +use axum::{ + Router, + body::Body, + http::{Method, StatusCode}, + middleware, +}; +use tokio::{ + net::TcpListener, + sync::{Notify, Semaphore, oneshot}, + task::JoinHandle, + time::timeout, +}; + +const DEADLINE: Duration = Duration::from_secs(30); +const CHUNK_BYTES: usize = 64 * 1024; + +fn channel_limits() -> ChannelLimits { + ChannelLimits { + max_frame_bytes: 128 * 1024, + max_buffered_bytes: 1024 * 1024, + ..ChannelLimits::default() + } +} + +fn client_limits() -> HttpClientLimits { + HttpClientLimits { + channel: channel_limits(), + max_buffered_bytes: 2 * 1024 * 1024, + max_response_bytes: 128 * 1024, + max_sse_line_bytes: 128 * 1024, + max_sse_event_bytes: 128 * 1024, + max_sse_chunk_bytes: 128 * 1024, + ..HttpClientLimits::default() + } +} + +fn server_limits() -> ServerLimits { + ServerLimits { + channel_limits: channel_limits(), + max_frame_bytes: 128 * 1024, + max_egress_bytes: 256 * 1024, + max_egress_frames: 4, + ..ServerLimits::default() + } +} + +struct Observed { + bytes: AtomicUsize, + notifications: AtomicUsize, + mutations: AtomicUsize, + mutated: Notify, + consumed: Semaphore, +} + +impl Default for Observed { + fn default() -> Self { + Self { + bytes: AtomicUsize::new(0), + notifications: AtomicUsize::new(0), + mutations: AtomicUsize::new(0), + mutated: Notify::new(), + consumed: Semaphore::new(0), + } + } +} + +fn agent(observed: Arc, chunks: usize) -> impl ConnectTo { + Agent + .builder() + .on_receive_request( + async |request: InitializeRequest, responder, _cx| { + responder.respond(InitializeResponse::new(request.protocol_version)) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |_request: NewSessionRequest, responder, _cx| { + responder.respond(NewSessionResponse::new("interop-session")) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: PromptRequest, responder, cx| { + observed.mutations.fetch_add(1, Ordering::SeqCst); + observed.mutated.notify_one(); + for _ in 0..chunks { + cx.send_notification(SessionNotification::new( + request.session_id.clone(), + SessionUpdate::AgentMessageChunk(ContentChunk::new(ContentBlock::Text( + TextContent::new("x".repeat(CHUNK_BYTES)), + ))), + ))?; + // Test-only application consumption handshake. This is deliberately + // not an HTTP body poll or a claimed transport acknowledgment. + observed.consumed.acquire().await.unwrap().forget(); + } + responder.respond(PromptResponse::new(StopReason::EndTurn)) + }, + agent_client_protocol::on_receive_request!(), + ) +} + +struct NetworkServer { + url: String, + shutdown: Option>, + task: Option>>, +} + +impl NetworkServer { + async fn start(router: Router) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}", listener.local_addr().unwrap()); + let (shutdown, stopped) = oneshot::channel(); + let task = tokio::spawn(async move { + axum::serve(listener, router) + .with_graceful_shutdown(async { + let _ = stopped.await; + }) + .await + }); + Self { + url, + shutdown: Some(shutdown), + task: Some(task), + } + } + + async fn finish(mut self) { + let _ = self.shutdown.take().unwrap().send(()); + timeout(DEADLINE, self.task.take().unwrap()) + .await + .expect("HTTP server did not shut down gracefully") + .unwrap() + .unwrap(); + } +} + +impl Drop for NetworkServer { + fn drop(&mut self) { + if let Some(shutdown) = self.shutdown.take() { + let _ = shutdown.send(()); + } + if let Some(task) = self.task.take() { + task.abort(); + } + } +} + +async fn conversation( + transport: impl ConnectTo, + observed: Arc, +) -> agent_client_protocol::Result<()> { + Client + .builder() + .on_receive_notification( + async move |notification: SessionNotification, _cx| { + assert_eq!(notification.session_id.to_string(), "interop-session"); + if let SessionUpdate::AgentMessageChunk(chunk) = notification.update { + let ContentBlock::Text(text) = chunk.content else { + panic!("expected text"); + }; + assert_eq!(text.text.len(), CHUNK_BYTES); + observed.bytes.fetch_add(text.text.len(), Ordering::SeqCst); + observed.notifications.fetch_add(1, Ordering::SeqCst); + observed.consumed.add_permits(1); + } + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_with(transport, async |cx| { + let initialized = cx + .send_request(InitializeRequest::new(ProtocolVersion::V1)) + .block_task() + .await?; + assert_eq!(initialized.protocol_version, ProtocolVersion::V1); + let session = cx + .send_request(NewSessionRequest::new(std::env::current_dir().unwrap())) + .block_task() + .await?; + let response = cx + .send_request(PromptRequest::new( + session.session_id, + vec![ContentBlock::Text(TextContent::new("mutate once"))], + )) + .block_task() + .await?; + assert_eq!(response.stop_reason, StopReason::EndTurn); + Ok(()) + }) + .await +} + +async fn interoperates(bounded_server: bool, bounded_client: bool, chunks: usize) { + let observed = Arc::new(Observed::default()); + let factory_observed = observed.clone(); + let factory = move || agent(factory_observed.clone(), chunks); + let router = if bounded_server { + AcpHttpServer::new_bounded(factory, server_limits()) + .unwrap() + .into_router() + } else { + AcpHttpServer::new(factory).into_router() + }; + let server = NetworkServer::start(router).await; + let client = HttpClient::new(&server.url).unwrap(); + let result = if bounded_client { + timeout( + DEADLINE, + conversation( + client.with_limits(client_limits()).unwrap(), + observed.clone(), + ), + ) + .await + } else { + timeout(DEADLINE, conversation(client, observed.clone())).await + }; + result + .expect("typed HTTP conversation or graceful client teardown stalled") + .expect("typed HTTP conversation failed"); + assert_eq!(observed.mutations.load(Ordering::SeqCst), 1); + assert_eq!(observed.notifications.load(Ordering::SeqCst), chunks); + assert_eq!(observed.bytes.load(Ordering::SeqCst), chunks * CHUNK_BYTES); + server.finish().await; +} + +#[tokio::test] +async fn bounded_client_and_server_use_typed_core_and_shutdown_gracefully() { + interoperates(true, true, 2).await; +} + +#[tokio::test] +async fn bounded_server_interoperates_with_standard_http_client() { + interoperates(true, false, 2).await; +} + +#[tokio::test] +async fn bounded_client_interoperates_with_standard_http_server() { + interoperates(false, true, 2).await; +} + +#[tokio::test] +async fn healthy_session_exceeds_eight_mib_with_bounded_outstanding_work() { + // Nine MiB over one session, while consumption releases each chunk before + // production resumes. HTTP egress is 256 KiB and client accounting is 2 MiB: + // these must be outstanding-work limits, never lifetime traffic limits. + interoperates(true, true, 144).await; +} + +#[tokio::test] +async fn lost_http_response_after_accepted_mutation_is_not_resubmitted() { + let observed = Arc::new(Observed::default()); + let factory_observed = observed.clone(); + let attempts = Arc::new(AtomicUsize::new(0)); + let middleware_attempts = attempts.clone(); + let middleware_observed = observed.clone(); + let router = + AcpHttpServer::new_bounded(move || agent(factory_observed.clone(), 0), server_limits()) + .unwrap() + .into_router() + .layer(middleware::from_fn( + move |request: axum::extract::Request, next: middleware::Next| { + let attempts = middleware_attempts.clone(); + let observed = middleware_observed.clone(); + async move { + let mutating = request.method() == Method::POST + && request.headers().contains_key("acp-session-id"); + let mut response = next.run(request).await; + if mutating { + attempts.fetch_add(1, Ordering::SeqCst); + assert_eq!(response.status(), StatusCode::ACCEPTED); + // Prove the side effect happened before destroying the HTTP result. + observed.mutated.notified().await; + *response.body_mut() = + Body::from_stream(futures::stream::iter([Err::, _>( + std::io::Error::new( + std::io::ErrorKind::ConnectionReset, + "injected response loss after accepted mutation", + ), + )])); + } + response + } + }, + )); + let server = NetworkServer::start(router).await; + let client = HttpClient::new(&server.url) + .unwrap() + .with_limits(client_limits()) + .unwrap(); + let result = timeout(DEADLINE, conversation(client, observed.clone())) + .await + .expect("uncertain POST did not terminate"); + assert!( + result.is_err(), + "lost accepted HTTP response must report transport failure" + ); + assert_eq!( + observed.mutations.load(Ordering::SeqCst), + 1, + "mutation was executed more than once" + ); + assert_eq!( + attempts.load(Ordering::SeqCst), + 1, + "accepted POST was retried" + ); + server.finish().await; +} diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index a49de91a..a21abce2 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -4,6 +4,10 @@ ### Added +- Add opt-in `BoundedChannel`, charged serialized frames, `ChannelLimits`, and + bounded `ConnectTo` extraction with producer-side protocol, pending-request, + task, and dynamic-handler admission. Legacy `Channel` fields remain unchanged. + - *(unstable-v2)* Expose pending session injection through the `unstable_session_inject` feature, including typed JSON-RPC dispatch and `V2Session` helpers to inject, replace, and revoke typed content. diff --git a/src/agent-client-protocol/README.md b/src/agent-client-protocol/README.md index d146e5d9..19cc4cb9 100644 --- a/src/agent-client-protocol/README.md +++ b/src/agent-client-protocol/README.md @@ -112,3 +112,19 @@ See the [crate documentation](https://docs.rs/agent-client-protocol) for: This project does not require a Contributor License Agreement (CLA). Instead, contributions are accepted under the following terms: > By contributing to this project, you agree that your contributions will be licensed under the [Apache License, Version 2.0](https://www.apache.org/licenses/LICENSE-2.0). You affirm that you have the legal right to submit your work, that you are not including code you do not have rights to, and that you understand contributions are made without requiring a Contributor License Agreement (CLA). + +## Opt-in bounded transports + +`BoundedChannel::duplex(ChannelLimits)` provides fail-fast, charged frame +admission without changing the compatibility-only unbounded `Channel` API. +Connect a protocol `Builder` directly to a bounded endpoint, or implement +`ConnectTo::into_transport_and_future` with `TransportChannel::Bounded`. +For component factories, `into_bounded_channel_and_future(limits)` supplies a +bounded endpoint directly; `DynConnectTo` preserves both interfaces. + +Outgoing protocol producers, pending requests, tasks, and dynamic handlers use +finite admission. `ChargedFrame` retains its reservation after dequeue; a +transport must keep it until body handoff/drop, which is **not** a peer ACK. +Limits measure reserved encoded bytes and work counts, not exact heap usage or +application conversion-hook allocations. See the book's Transport Architecture +chapter for defaults, ownership, terminal overload, and compatibility scope. diff --git a/src/agent-client-protocol/src/bounded.rs b/src/agent-client-protocol/src/bounded.rs new file mode 100644 index 00000000..99873778 --- /dev/null +++ b/src/agent-client-protocol/src/bounded.rs @@ -0,0 +1,721 @@ +//! Opt-in admission-bounded frame channels. Legacy [`crate::Channel`] is unchanged. +use crate::{Channel, ConnectTo, Error, Role, TransportFrame}; +use futures::{ + FutureExt, Stream, StreamExt, + channel::{mpsc, oneshot}, + future::{BoxFuture, Shared}, +}; +use std::{ + io::Write, + pin::Pin, + sync::{Arc, Mutex}, + task::{Context, Poll}, +}; + +/// Finite per-direction admission limits. Byte limits measure encoded JSON, +/// not allocator overhead or application-owned values/future captures. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ChannelLimits { + /// Maximum encoded size of one frame. + pub max_frame_bytes: usize, + /// Maximum reserved encoded bytes, including queued and handed-off frames. + pub max_buffered_bytes: usize, + /// Maximum queued, preparing, or handed-off frames. + pub max_buffered_frames: usize, + /// Maximum requests awaiting a reply. + pub max_pending_requests: usize, + /// Maximum queued/running tasks and registered dynamic handlers (including + /// the transport driver). + pub max_tasks: usize, +} +impl Default for ChannelLimits { + fn default() -> Self { + Self { + max_frame_bytes: 1024 * 1024, + max_buffered_bytes: 16 * 1024 * 1024, + max_buffered_frames: 256, + max_pending_requests: 256, + max_tasks: 256, + } + } +} +impl ChannelLimits { + /// Reject zero limits and a byte budget smaller than one maximum frame. + pub fn validate(self) -> Result { + if self.max_frame_bytes == 0 + || self.max_buffered_bytes < self.max_frame_bytes + || self.max_buffered_frames == 0 + || self.max_pending_requests == 0 + || self.max_tasks == 0 + { + return Err(limit_error("invalid channel limits")); + } + Ok(self) + } +} +fn limit_error(message: &str) -> Error { + crate::util::internal_error(message) +} +type Failure = Shared>; +struct Terminal { + error: Mutex>, + signal: Mutex>>, + failure: Failure, +} +impl Terminal { + fn new() -> Arc { + let (tx, rx) = oneshot::channel(); + Arc::new(Self { + error: Mutex::new(None), + signal: Mutex::new(Some(tx)), + failure: rx + .map(|r| r.unwrap_or_else(|_| limit_error("bounded channel closed"))) + .boxed() + .shared(), + }) + } + fn fail(&self, message: &str) -> Error { + let mut error = self.error.lock().expect("terminal mutex poisoned"); + let error = error.get_or_insert_with(|| limit_error(message)).clone(); + if let Some(tx) = self.signal.lock().expect("terminal signal poisoned").take() { + drop(tx.send(error.clone())); + } + error + } + fn check(&self) -> Result<(), Error> { + match &*self.error.lock().expect("terminal mutex poisoned") { + Some(e) => Err(e.clone()), + None => Ok(()), + } + } +} +#[derive(Default)] +struct Usage { + frames: usize, + bytes: usize, + tasks: usize, + admission_closed: bool, +} +#[derive(Clone)] +pub(crate) struct Budget { + limits: ChannelLimits, + usage: Arc>, + terminal: Arc, +} +impl std::fmt::Debug for Budget { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Budget") + .field("limits", &self.limits) + .finish_non_exhaustive() + } +} +impl Budget { + pub(crate) fn limits(&self) -> ChannelLimits { + self.limits + } + pub(crate) fn fail(&self, message: &str) -> Error { + self.terminal.fail(message) + } + pub(crate) fn check(&self) -> Result<(), Error> { + self.terminal.check() + } + pub(crate) fn check_admission(&self) -> Result<(), Error> { + self.check()?; + if self + .usage + .lock() + .expect("budget mutex poisoned") + .admission_closed + { + return Err(limit_error("bounded channel admission closed")); + } + Ok(()) + } + fn close_admission(&self) { + self.usage + .lock() + .expect("budget mutex poisoned") + .admission_closed = true; + } + pub(crate) fn failure(&self) -> BoxFuture<'static, Error> { + self.terminal.failure.clone().boxed() + } + pub(crate) fn reserve(&self) -> Result { + self.check()?; + let mut used = self.usage.lock().expect("budget mutex poisoned"); + if used.admission_closed { + return Err(limit_error("bounded channel admission closed")); + } + let bytes = self.limits.max_frame_bytes; + if used.frames >= self.limits.max_buffered_frames + || bytes > self.limits.max_buffered_bytes - used.bytes + { + return Err(self.fail("bounded channel frame/byte admission exhausted")); + } + used.frames += 1; + used.bytes += bytes; + Ok(Charge(Arc::new(Permit { + budget: self.clone(), + bytes, + task: false, + }))) + } + pub(crate) fn task(&self) -> Result { + self.check()?; + let mut used = self.usage.lock().expect("budget mutex poisoned"); + if used.admission_closed { + return Err(limit_error("bounded channel admission closed")); + } + if used.tasks >= self.limits.max_tasks { + return Err(self.fail("bounded channel task admission exhausted")); + } + used.tasks += 1; + Ok(Charge(Arc::new(Permit { + budget: self.clone(), + bytes: 0, + task: true, + }))) + } +} +struct Permit { + budget: Budget, + bytes: usize, + task: bool, +} +impl Drop for Permit { + fn drop(&mut self) { + let mut used = self.budget.usage.lock().expect("budget mutex poisoned"); + if self.task { + used.tasks -= 1; + } else { + used.frames -= 1; + used.bytes -= self.bytes; + } + } +} +#[derive(Clone)] +pub(crate) struct Charge(Arc); +impl std::fmt::Debug for Charge { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Charge").finish_non_exhaustive() + } +} + +/// Serialized frame with admission ownership. Receiving or forwarding it does +/// not release its reservation. Retain it until body handoff/drop, which is NOT +/// an acknowledgment of socket flush, peer parsing, or protocol dispatch. +#[derive(Debug)] +pub struct ChargedFrame { + bytes: Box<[u8]>, + pub(crate) charges: Vec, +} +impl ChargedFrame { + /// Borrow the charged JSON bytes without releasing admission. + #[must_use] + pub fn as_bytes(&self) -> &[u8] { + &self.bytes + } + /// Decode while retaining admission. Decoded JSON has additional structural + /// overhead; consumers must not accumulate uncharged decoded copies. + #[must_use] + pub fn decode(&self) -> TransportFrame { + TransportFrame::parse_json( + std::str::from_utf8(&self.bytes).expect("serialized frame is UTF-8"), + ) + } +} + +/// Cloneable fail-fast producer. There is no waiting-producer queue. +#[derive(Clone)] +pub struct BoundedSender { + tx: mpsc::UnboundedSender, + pub(crate) budget: Budget, +} +impl std::fmt::Debug for BoundedSender { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BoundedSender") + .field("budget", &self.budget) + .finish_non_exhaustive() + } +} +impl BoundedSender { + /// Validated per-direction limits. + #[must_use] + pub fn limits(&self) -> ChannelLimits { + self.budget.limits + } + /// Resolve when any admission/serialization failure makes this channel terminal. + #[must_use] + pub fn failure(&self) -> BoxFuture<'static, Error> { + self.budget.failure() + } + /// Fail all escaped producers and wake the connection driver. + #[allow( + clippy::must_use_candidate, + reason = "The primary effect is terminating the channel; callers may optionally propagate the error." + )] + pub fn fail(&self, reason: &str) -> Error { + self.budget.fail(reason) + } + /// Gracefully close this direction for every sender clone. + /// + /// New sends fail without making the channel terminal. Already enqueued + /// frames retain their charges and remain available to the receiver, which + /// observes EOF after draining them. This does not signal [`Self::failure`] + /// or close the opposite direction. + pub fn close_channel(&self) { + self.budget.close_admission(); + self.tx.close_channel(); + } + /// Admit and serialize without allocating an unbounded intermediate String. + pub fn try_send(&self, frame: TransportFrame) -> Result<(), Error> { + let charge = self.budget.reserve()?; + self.send_charged(frame, vec![charge]) + } + /// Copy already serialized wire text after admission and size validation. + /// Like TransportFrame::parse_json, malformed wire text is preserved. + pub fn try_send_serialized(&self, json: &str) -> Result<(), Error> { + let charge = self.budget.reserve()?; + if json.len() > self.limits().max_frame_bytes { + return Err(self.fail("bounded channel frame too large")); + } + self.enqueue(ChargedFrame { + bytes: json.as_bytes().into(), + charges: vec![charge], + }) + } + /// Forward admission ownership without releasing it at an internal handoff. + /// Cross-budget forwarding acquires destination admission before releasing + /// the original reservation. Same-budget forwarding does not double charge. + pub fn try_forward(&self, mut frame: ChargedFrame) -> Result<(), Error> { + self.budget.check_admission()?; + if frame.bytes.len() > self.limits().max_frame_bytes { + return Err(self.fail("bounded channel forwarded frame too large")); + } + if !frame + .charges + .iter() + .any(|c| Arc::ptr_eq(&c.0.budget.usage, &self.budget.usage)) + { + let charge = self.budget.reserve()?; + frame.charges = vec![charge]; + } + self.enqueue(frame) + } + fn enqueue(&self, frame: ChargedFrame) -> Result<(), Error> { + self.budget.check_admission()?; + // Serialize closure and enqueue so a racing close cannot turn a normal + // rejected send into terminal failure and discard admitted frames. + let result = { + let used = self.budget.usage.lock().expect("budget mutex poisoned"); + if used.admission_closed { + return Err(limit_error("bounded channel admission closed")); + } + self.tx.unbounded_send(frame) + }; + result.map_err(|_| self.fail("bounded channel receiver closed")) + } + pub(crate) fn send_charged( + &self, + frame: TransportFrame, + charges: Vec, + ) -> Result<(), Error> { + let mut charges = charges; + if !charges + .iter() + .any(|c| Arc::ptr_eq(&c.0.budget.usage, &self.budget.usage)) + { + charges.push(self.budget.reserve()?); + } + self.budget.check_admission()?; + let bytes = serialize_frame(&frame, self.limits().max_frame_bytes) + .map_err(|_| self.fail("bounded channel frame serialization exceeded limit"))?; + self.enqueue(ChargedFrame { bytes, charges }) + } +} +/// Single consumer of serialized, charged frames. +#[derive(Debug)] +pub struct BoundedReceiver { + rx: mpsc::UnboundedReceiver, + budget: Budget, + failure: Failure, +} +impl Stream for BoundedReceiver { + type Item = ChargedFrame; + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if Pin::new(&mut self.failure).poll(cx).is_ready() { + self.rx.close(); + while let Poll::Ready(Some(_)) = Pin::new(&mut self.rx).poll_next(cx) {} + return Poll::Ready(None); + } + Pin::new(&mut self.rx).poll_next(cx) + } +} +/// Opt-in bounded frame endpoint; no concrete unbounded sender is exposed. +#[derive(Debug)] +pub struct BoundedChannel { + /// Fail-fast producer for frames sent to the peer. + pub tx: BoundedSender, + /// Charged frames received from the peer. + pub rx: BoundedReceiver, +} +impl BoundedChannel { + /// Construct two endpoints with independent directional budgets and shared + /// terminal state. Limits must be nonzero and fit one maximum-sized frame. + pub fn duplex(limits: ChannelLimits) -> Result<(Self, Self), Error> { + let limits = limits.validate()?; + let terminal = Terminal::new(); + let a = Budget { + limits, + usage: Arc::default(), + terminal: terminal.clone(), + }; + let b = Budget { + limits, + usage: Arc::default(), + terminal, + }; + let (atx, brx) = mpsc::unbounded(); + let (btx, arx) = mpsc::unbounded(); + Ok(( + Self { + tx: BoundedSender { + tx: atx, + budget: a.clone(), + }, + rx: BoundedReceiver { + rx: arx, + failure: b.terminal.failure.clone(), + budget: b.clone(), + }, + }, + Self { + tx: BoundedSender { tx: btx, budget: b }, + rx: BoundedReceiver { + rx: brx, + failure: a.terminal.failure.clone(), + budget: a, + }, + }, + )) + } + /// Validated per-direction limits. + #[must_use] + pub fn limits(&self) -> ChannelLimits { + self.tx.limits() + } +} +/// Additive transport boundary. The legacy variant has no admission guarantee. +#[derive(Debug)] +pub enum TransportChannel { + /// Compatibility-only, unbounded endpoint. + Legacy(Channel), + /// Opt-in admission-bounded endpoint. + Bounded(BoundedChannel), +} +impl ConnectTo for BoundedChannel { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + async fn copy(mut rx: BoundedReceiver, tx: BoundedSender) -> Result<(), Error> { + while let Some(frame) = rx.next().await { + tx.try_forward(frame)?; + } + rx.budget.check() + } + let (other, future) = client.into_bounded_channel_and_future(self.limits())?; + futures::try_join!(copy(self.rx, other.tx), copy(other.rx, self.tx), future)?; + Ok(()) + } + fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), Error>>) { + drop(self); + let (channel, peer) = Channel::duplex(); + drop(peer); + ( + channel, + async { + Err(limit_error( + "bounded channel cannot be extracted as legacy Channel", + )) + } + .boxed(), + ) + } + fn into_transport_and_future( + self, + ) -> (TransportChannel, BoxFuture<'static, Result<(), Error>>) { + (TransportChannel::Bounded(self), async { Ok(()) }.boxed()) + } + fn into_bounded_channel_and_future( + self, + limits: ChannelLimits, + ) -> Result<(BoundedChannel, BoxFuture<'static, Result<(), Error>>), Error> { + let limits = limits.validate()?; + if self.limits() != limits { + return Err(self + .tx + .fail("bounded endpoint limits differ from requested limits")); + } + Ok((self, async { Ok(()) }.boxed())) + } +} + +pub(crate) struct CappedWriter { + bytes: Vec, + limit: usize, +} +impl CappedWriter { + pub(crate) fn new(limit: usize) -> Self { + Self { + bytes: Vec::new(), + limit, + } + } + pub(crate) fn finish(self) -> Box<[u8]> { + self.bytes.into_boxed_slice() + } +} +impl Write for CappedWriter { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + if buf.len() > self.limit - self.bytes.len() { + return Err(std::io::Error::other("encoded JSON limit exceeded")); + } + self.bytes + .try_reserve_exact(buf.len()) + .map_err(std::io::Error::other)?; + self.bytes.extend_from_slice(buf); + Ok(buf.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} +pub(crate) fn serialize( + value: &T, + limit: usize, +) -> Result, Error> { + let mut out = CappedWriter::new(limit); + serde_json::to_writer(&mut out, value)?; + Ok(out.finish()) +} +fn serialize_frame(frame: &TransportFrame, limit: usize) -> Result, Error> { + match frame { + TransportFrame::Single(value) => serialize(value, limit), + TransportFrame::Batch(value) => serialize(value, limit), + TransportFrame::Malformed { raw, .. } => { + if raw.len() > limit { + return Err(limit_error("frame too large")); + } + Ok(raw.as_bytes().into()) + } + } +} + +#[derive(Debug)] +pub(crate) struct IncomingFrame { + pub(crate) frame: TransportFrame, + pub(crate) charges: Vec, +} +impl From for IncomingFrame { + fn from(frame: TransportFrame) -> Self { + Self { + frame, + charges: vec![], + } + } +} +impl From for IncomingFrame { + fn from(frame: ChargedFrame) -> Self { + Self { + frame: frame.decode(), + charges: frame.charges, + } + } +} +pub(crate) enum TransportSender { + Legacy(mpsc::UnboundedSender), + Bounded(BoundedSender), +} +impl From> for TransportSender { + fn from(tx: mpsc::UnboundedSender) -> Self { + Self::Legacy(tx) + } +} +impl TransportSender { + pub(crate) fn send(&self, frame: TransportFrame, charges: Vec) -> Result<(), Error> { + match self { + Self::Legacy(tx) => tx + .unbounded_send(frame) + .map_err(crate::util::internal_error), + Self::Bounded(tx) => tx.send_charged(frame, charges), + } + } +} +impl TransportChannel { + pub(crate) fn split( + self, + ) -> ( + futures::stream::BoxStream<'static, IncomingFrame>, + TransportSender, + Option, + ) { + match self { + Self::Legacy(channel) => ( + channel.rx.map(IncomingFrame::from).boxed(), + TransportSender::Legacy(channel.tx), + None, + ), + Self::Bounded(channel) => { + let budget = channel.tx.budget.clone(); + ( + channel.rx.map(IncomingFrame::from).boxed(), + TransportSender::Bounded(channel.tx), + Some(budget), + ) + } + } + } +} + +pub(crate) struct DriverLifetime { + budget: Budget, + graceful: bool, +} +impl DriverLifetime { + pub(crate) fn new(budget: Budget) -> Self { + Self { + budget, + graceful: false, + } + } + pub(crate) fn finish_gracefully(&mut self) { + self.budget.close_admission(); + self.graceful = true; + } +} +impl Drop for DriverLifetime { + fn drop(&mut self) { + if !self.graceful { + self.budget.fail("bounded protocol driver stopped"); + } + } +} + +#[cfg(test)] +impl Budget { + pub(crate) fn snapshot(&self) -> (usize, usize, usize) { + let used = self.usage.lock().unwrap(); + (used.frames, used.bytes, used.tasks) + } +} + +#[cfg(test)] +mod tests { + use super::*; + fn limits() -> ChannelLimits { + ChannelLimits { + max_frame_bytes: 64, + max_buffered_bytes: 128, + max_buffered_frames: 2, + max_pending_requests: 2, + max_tasks: 2, + } + } + + #[test] + fn byte_and_frame_admission_persist_after_dequeue_and_forward() { + let (a, mut b) = BoundedChannel::duplex(limits()).unwrap(); + a.tx.try_send_serialized("{}").unwrap(); + let frame = b.rx.next().now_or_never().unwrap().unwrap(); + assert_eq!(a.tx.budget.snapshot(), (1, 64, 0)); + a.tx.try_forward(frame).unwrap(); + assert_eq!(a.tx.budget.snapshot(), (1, 64, 0)); + let frame = b.rx.next().now_or_never().unwrap().unwrap(); + drop(frame); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + a.tx.try_send_serialized(&"x".repeat(64)).unwrap(); + assert!(a.tx.try_send_serialized(&"x".repeat(65)).is_err()); + assert!(a.tx.failure().now_or_never().is_some()); + assert!(b.rx.next().now_or_never().unwrap().is_none()); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + } + + #[test] + fn graceful_close_rejects_clones_but_drains_charged_frames_without_failure() { + let (a, mut b) = BoundedChannel::duplex(limits()).unwrap(); + let clone = a.tx.clone(); + a.tx.try_send_serialized("{}").unwrap(); + a.tx.try_send_serialized("[]").unwrap(); + a.tx.close_channel(); + a.tx.close_channel(); + assert!(clone.try_send_serialized("rejected").is_err()); + assert!(clone.try_send(TransportFrame::parse_json("{}")).is_err()); + assert!(a.tx.failure().now_or_never().is_none()); + assert_eq!(a.tx.budget.snapshot(), (2, 128, 0)); + let first = b.rx.next().now_or_never().unwrap().unwrap(); + let second = b.rx.next().now_or_never().unwrap().unwrap(); + assert_eq!(first.as_bytes(), b"{}"); + assert_eq!(second.as_bytes(), b"[]"); + assert!(b.rx.next().now_or_never().unwrap().is_none()); + assert_eq!(a.tx.budget.snapshot(), (2, 128, 0)); + drop((first, second)); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + b.tx.try_send_serialized("{}").unwrap(); + assert!(a.tx.failure().now_or_never().is_none()); + } + + #[test] + fn producer_flood_before_poll_is_terminal_and_drains_charges() { + let (a, mut b) = BoundedChannel::duplex(limits()).unwrap(); + let escaped = a.tx.clone(); + a.tx.try_send_serialized("{}").unwrap(); + a.tx.try_send_serialized("{}").unwrap(); + assert!(escaped.try_send_serialized("{}").is_err()); + assert_eq!(a.tx.budget.snapshot(), (2, 128, 0)); + assert!(b.rx.next().now_or_never().unwrap().is_none()); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + assert!(escaped.try_send_serialized("{}").is_err()); + } + + #[test] + fn cross_budget_forward_reserves_destination_and_releases_source() { + let (a, mut b) = BoundedChannel::duplex(limits()).unwrap(); + let (c, mut d) = BoundedChannel::duplex(limits()).unwrap(); + a.tx.try_send_serialized("{}").unwrap(); + let frame = b.rx.next().now_or_never().unwrap().unwrap(); + c.tx.try_forward(frame).unwrap(); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + assert_eq!(c.tx.budget.snapshot(), (1, 64, 0)); + drop(d.rx.next().now_or_never().unwrap().unwrap()); + assert_eq!(c.tx.budget.snapshot(), (0, 0, 0)); + } + + #[test] + fn dropped_receiver_and_serialization_errors_release_reservations() { + let (a, b) = BoundedChannel::duplex(limits()).unwrap(); + drop(b); + assert!(a.tx.try_send_serialized("{}").is_err()); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + let (a, _b) = BoundedChannel::duplex(limits()).unwrap(); + let frame = TransportFrame::Single( + crate::RawJsonRpcMessage::notification( + "large".into(), + serde_json::json!({"payload": "x".repeat(65)}), + ) + .unwrap(), + ); + assert!(a.tx.try_send(frame).is_err()); + assert_eq!(a.tx.budget.snapshot(), (0, 0, 0)); + } + + #[test] + fn dyn_transport_preserves_bounded_variant_and_legacy_extraction_fails() { + let (a, _b) = BoundedChannel::duplex(limits()).unwrap(); + let erased = crate::DynConnectTo::::new(a); + let (channel, _) = erased.into_transport_and_future(); + assert!(matches!(channel, TransportChannel::Bounded(_))); + let (a, _b) = BoundedChannel::duplex(limits()).unwrap(); + let (mut channel, future) = + >::into_channel_and_future(a); + assert!(future.now_or_never().unwrap().is_err()); + assert!(channel.rx.next().now_or_never().unwrap().is_none()); + } +} diff --git a/src/agent-client-protocol/src/component.rs b/src/agent-client-protocol/src/component.rs index ba8b9297..13c6fe1f 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -129,6 +129,29 @@ pub trait ConnectTo: Send + 'static { client: impl ConnectTo, ) -> impl Future> + Send; + /// Extract the transport without erasing opt-in bounded admission. + fn into_transport_and_future(self) -> (crate::TransportChannel, BoxFuture<'static, Result<()>>) + where + Self: Sized, + { + let (channel, future) = self.into_channel_and_future(); + (crate::TransportChannel::Legacy(channel), future) + } + + /// Connect using an opt-in bounded endpoint, without a legacy adapter pump. + /// Components that consume only legacy channels fail closed when extracting + /// the supplied bounded endpoint. + fn into_bounded_channel_and_future( + self, + limits: crate::ChannelLimits, + ) -> Result<(crate::BoundedChannel, BoxFuture<'static, Result<()>>)> + where + Self: Sized, + { + let (channel, peer) = crate::BoundedChannel::duplex(limits)?; + Ok((channel, Box::pin(self.connect_to(peer)))) + } + /// Convert this component into a channel endpoint and connection future. /// /// The returned [`Channel`] is the canonical frame-aware boundary. It carries @@ -171,6 +194,14 @@ trait ErasedConnectTo: Send { client: Box>, ) -> BoxFuture<'static, Result<()>>; + fn into_transport_and_future_erased( + self: Box, + ) -> (crate::TransportChannel, BoxFuture<'static, Result<()>>); + fn into_bounded_channel_and_future_erased( + self: Box, + limits: crate::ChannelLimits, + ) -> Result<(crate::BoundedChannel, BoxFuture<'static, Result<()>>)>; + fn into_channel_and_future_erased(self: Box) -> (Channel, BoxFuture<'static, Result<()>>); } @@ -195,6 +226,18 @@ impl, R: Role> ErasedConnectTo for C { }) } + fn into_transport_and_future_erased( + self: Box, + ) -> (crate::TransportChannel, BoxFuture<'static, Result<()>>) { + (*self).into_transport_and_future() + } + fn into_bounded_channel_and_future_erased( + self: Box, + limits: crate::ChannelLimits, + ) -> Result<(crate::BoundedChannel, BoxFuture<'static, Result<()>>)> { + (*self).into_bounded_channel_and_future(limits) + } + fn into_channel_and_future_erased( self: Box, ) -> (Channel, BoxFuture<'static, Result<()>>) { @@ -251,6 +294,18 @@ impl ConnectTo for DynConnectTo { .await } + fn into_transport_and_future( + self, + ) -> (crate::TransportChannel, BoxFuture<'static, Result<()>>) { + self.inner.into_transport_and_future_erased() + } + fn into_bounded_channel_and_future( + self, + limits: crate::ChannelLimits, + ) -> Result<(crate::BoundedChannel, BoxFuture<'static, Result<()>>)> { + self.inner.into_bounded_channel_and_future_erased(limits) + } + fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>) { self.inner.into_channel_and_future_erased() } diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 5068cea4..e9033811 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -362,14 +362,12 @@ impl Serialize for RawJsonRpcMessage { S: serde::Serializer, { match self { - Self::Request(request) => { - VersionedJsonRpcMessage::wrap(request.clone()).serialize(serializer) - } + Self::Request(request) => VersionedJsonRpcMessage::wrap(request).serialize(serializer), Self::Notification(notification) => { - VersionedJsonRpcMessage::wrap(notification.clone()).serialize(serializer) + VersionedJsonRpcMessage::wrap(notification).serialize(serializer) } Self::Response(response) => { - VersionedJsonRpcMessage::wrap(response.clone()).serialize(serializer) + VersionedJsonRpcMessage::wrap(response).serialize(serializer) } } } @@ -1906,7 +1904,14 @@ impl< // Convert transport into server - this returns a channel for us to use // and a future that runs the transport. let transport_component = crate::DynConnectTo::new(transport); - let (transport_channel, transport_future) = transport_component.into_channel_and_future(); + let (transport_channel, transport_future) = transport_component.into_transport_and_future(); + let (transport_incoming_rx, transport_outgoing_tx, budget) = transport_channel.split(); + pending_replies + .inner + .lock() + .expect("pending replies mutex poisoned") + .budget + .clone_from(&budget); let (transport_completion_tx, transport_completion_rx) = oneshot::channel(); let transport_completion = transport_completion_rx .map(|result| { @@ -1919,7 +1924,7 @@ impl< .boxed() .shared(); - let connection = ConnectionTo::new( + let mut connection = ConnectionTo::new( me.counterpart(), outgoing_tx, new_task_tx, @@ -1928,23 +1933,23 @@ impl< pending_replies.registrar(), protocol_mode, ); + connection.message_tx.set_budget(budget.clone()); + connection.task_tx.budget.clone_from(&budget); + connection.dynamic_handler_tx.budget.clone_from(&budget); + let admission_failure = budget.as_ref().map(crate::bounded::Budget::failure); let spawn_result = connection.spawn(async move { let result = transport_future.await; drop(transport_completion_tx.send(result.clone())); result }); - // Destructure the channel endpoints - let Channel { - rx: transport_incoming_rx, - tx: transport_outgoing_tx, - } = transport_channel; - let protocol_compat = ProtocolCompat::new(protocol_mode); + let bounded_lifetime = budget.map(crate::bounded::DriverLifetime::new); let future = crate::util::instrument_with_connection_name(name, { let connection = connection.clone(); async move { + let mut bounded_lifetime = bounded_lifetime; let () = spawn_result?; let background = async { @@ -1984,12 +1989,25 @@ impl< .await }; - run_until_connection_close( + let driver = run_until_connection_close( background, main_fn(connection.clone()), connection.incoming_closed.clone(), - ) - .await + ); + let result = if let Some(failure) = admission_failure { + match future::select(failure, Box::pin(driver)).await { + Either::Left((error, _)) => Err(error), + Either::Right((result, _)) => result, + } + } else { + driver.await + }; + if result.is_ok() + && let Some(lifetime) = &mut bounded_lifetime + { + lifetime.finish_gracefully(); + } + result } }); @@ -2097,6 +2115,8 @@ pub(crate) struct ResponsePayload { /// blocking consumers, and responses routed later do not hold the dispatch /// loop. pub(crate) ack_tx: Option>, + /// Admission retained through the response oneshot and callback consumption. + charges: Vec, } type ResponseRouteHook = @@ -2141,6 +2161,7 @@ impl std::fmt::Debug for ResponsePayload { f.debug_struct("ResponsePayload") .field("result", &self.result) .field("ack_tx", &self.ack_tx.as_ref().map(|_| "...")) + .field("charges", &self.charges) .finish() } } @@ -2177,6 +2198,7 @@ impl PendingReply { .send(ResponsePayload { result: Err(error), ack_tx: None, + charges: vec![], }) .is_err() { @@ -2194,6 +2216,7 @@ impl PendingReply { struct PendingRepliesInner { incoming_closed: bool, replies: HashMap, + budget: Option, } #[derive(Clone, Default)] @@ -2274,6 +2297,15 @@ impl PendingRepliesRegistrar { let mut inner = inner.lock().expect("pending replies mutex poisoned"); if inner.incoming_closed { Err(reply) + } else if let Some(budget) = &inner.budget + && (budget.check().is_err() + || (!inner.replies.contains_key(&id) + && inner.replies.len() >= budget.limits().max_pending_requests)) + { + let error = budget.fail("bounded pending request admission exhausted"); + drop(inner); + reply.fail(error); + return false; } else { Ok(inner.replies.insert(id, reply)) } @@ -2804,6 +2836,7 @@ fn response_receipt_teardown_error() -> crate::Error { struct CompletedResponseFrame { frame: TransportFrame, receipts: Vec, + charges: Vec, } /// Messages send to be serialized over the transport. @@ -2823,6 +2856,23 @@ impl std::fmt::Debug for ResponseDestination { } impl ResponseDestination { + fn retain_incoming(&self, charges: &[crate::bounded::Charge]) { + match self { + Self::Individual(slot) => slot + .state + .incoming_charges + .lock() + .expect("response charge mutex poisoned") + .extend_from_slice(charges), + Self::Batch(slot) => slot + .state + .lock() + .expect("batch response accumulator mutex poisoned") + .charges + .extend_from_slice(charges), + } + } + fn individual() -> Self { Self::Individual(IndividualResponseSlot::default()) } @@ -2836,6 +2886,7 @@ impl ResponseDestination { receipts: (0..slot_count).map(|_| None).collect(), dispatch_complete: false, emitted: false, + charges: vec![], })); ( @@ -2898,7 +2949,13 @@ impl ResponseDestination { #[derive(Clone, Debug, Default)] struct IndividualResponseSlot { - completed: Arc, + state: Arc, +} + +#[derive(Debug, Default)] +struct IndividualResponseState { + completed: AtomicBool, + incoming_charges: Mutex>, } impl IndividualResponseSlot { @@ -2907,7 +2964,7 @@ impl IndividualResponseSlot { response: RawJsonRpcMessage, receipt: Option, ) -> Option { - if self.completed.swap(true, Ordering::AcqRel) { + if self.state.completed.swap(true, Ordering::AcqRel) { tracing::warn!("Ignoring duplicate completion of JSON-RPC request"); return None; } @@ -2915,12 +2972,23 @@ impl IndividualResponseSlot { Some(CompletedResponseFrame { frame: TransportFrame::Single(response), receipts: receipt.into_iter().collect(), + charges: std::mem::take( + &mut *self + .state + .incoming_charges + .lock() + .expect("response charge mutex poisoned"), + ), }) } } fn batch_response_frame( - (responses, receipts): (Vec, Vec), + (responses, receipts, charges): ( + Vec, + Vec, + Vec, + ), ) -> CompletedResponseFrame { CompletedResponseFrame { frame: TransportFrame::Batch( @@ -2928,6 +2996,7 @@ fn batch_response_frame( .expect("a completed JSON-RPC response batch is non-empty"), ), receipts, + charges, } } @@ -2974,7 +3043,11 @@ fn promote_abandoned_response(state: &mut BatchResponseState, index: usize) { fn take_completed_batch( state: &mut BatchResponseState, -) -> Option<(Vec, Vec)> { +) -> Option<( + Vec, + Vec, + Vec, +)> { if !state.dispatch_complete || state.remaining != 0 || state.emitted { return None; } @@ -2990,7 +3063,7 @@ fn take_completed_batch( }) .collect(); let receipts = state.receipts.iter_mut().filter_map(Option::take).collect(); - Some((responses, receipts)) + Some((responses, receipts, std::mem::take(&mut state.charges))) } #[derive(Clone)] @@ -3019,7 +3092,11 @@ impl BatchResponseSlot { fn finish_handler_attempt( self, - ) -> Option<(Vec, Vec)> { + ) -> Option<( + Vec, + Vec, + Vec, + )> { let mut state = self .state .lock() @@ -3037,7 +3114,11 @@ impl BatchResponseSlot { self, response: RawJsonRpcMessage, receipt: Option, - ) -> Option<(Vec, Vec)> { + ) -> Option<( + Vec, + Vec, + Vec, + )> { let mut state = self .state .lock() @@ -3071,7 +3152,11 @@ impl BatchResponseSlot { fn abandon( self, fallback: RawJsonRpcMessage, - ) -> Option<(Vec, Vec)> { + ) -> Option<( + Vec, + Vec, + Vec, + )> { let mut state = self .state .lock() @@ -3105,6 +3190,7 @@ struct BatchResponseState { receipts: Vec>, dispatch_complete: bool, emitted: bool, + charges: Vec, } #[derive(Clone, Debug)] @@ -3136,18 +3222,29 @@ impl Drop for ResponderHandlerAttempt { struct ResponseReplyTarget { id: RequestId, method: String, - sender: Arc>>>, + state: Arc>, ordering: ResponseOrdering, dispatch: ResponseDispatch, } +struct ResponseReplyState { + sender: Option>, + charges: Vec, +} + impl ResponseReplyTarget { - fn route(self, result: Result) { - let sender = self - .sender + fn retain_incoming(&self, charges: &[crate::bounded::Charge]) { + self.state .lock() .expect("response reply mutex poisoned") - .take(); + .charges + .extend_from_slice(charges); + } + fn route(self, result: Result) { + let (sender, charges) = { + let mut state = self.state.lock().expect("response reply mutex poisoned"); + (state.sender.take(), std::mem::take(&mut state.charges)) + }; let Some(sender) = sender else { tracing::debug!( method = %self.method, @@ -3158,7 +3255,14 @@ impl ResponseReplyTarget { }; let ack_tx = self.dispatch.acknowledgment(&self.ordering); - if sender.send(ResponsePayload { result, ack_tx }).is_err() { + if sender + .send(ResponsePayload { + result, + ack_tx, + charges, + }) + .is_err() + { tracing::debug!( method = %self.method, id = ?self.id, @@ -3225,6 +3329,10 @@ impl HandlerErrorTarget { #[derive(Debug)] enum OutgoingMessage { + Admitted { + message: Box, + charge: crate::bounded::Charge, + }, /// Close the outgoing application queue and acknowledge after every /// already-accepted message has entered the raw transport queue. CloseAfterDraining { done: oneshot::Sender<()> }, @@ -3293,6 +3401,50 @@ enum OutgoingMessage { }, } +impl OutgoingMessage { + fn normalize(&mut self, budget: &crate::bounded::Budget) -> Result<(), crate::Error> { + fn normalized( + v: &impl Serialize, + limit: usize, + ) -> Result { + let bytes = crate::bounded::serialize(v, limit)?; + Ok(serde_json::from_slice(&bytes)?) + } + let limit = budget.limits().max_frame_bytes; + match self { + Self::Request { + id, + method, + untyped, + .. + } => { + let values = normalized(&(&*id, &*method, &*untyped), limit)?; + (*id, *method, *untyped) = values; + } + Self::Notification { untyped } => *untyped = normalized(untyped, limit)?, + Self::Response { + id, + method, + response, + .. + } => { + let values = normalized(&(&*id, &*method, &*response), limit)?; + (*id, *method, *response) = values; + } + Self::AbandonedBatchResponse { id, method, .. } => { + let values = normalized(&(&*id, &*method), limit)?; + (*id, *method) = values; + } + Self::UncorrelatedErrorResponse { error, .. } => *error = normalized(error, limit)?, + Self::CloseAfterDraining { .. } + | Self::BatchDispatchComplete { .. } + | Self::BatchHandlerAttemptComplete { .. } => {} + Self::Admitted { .. } => unreachable!("nested admission"), + } + Ok(()) + } +} + /// Return type from JrHandler; indicates whether the request was handled or not. #[must_use] #[derive(Debug)] @@ -3577,7 +3729,7 @@ pub struct ConnectionTo { counterpart: Counterpart, message_tx: OutgoingMessageTx, task_tx: TaskTx, - dynamic_handler_tx: mpsc::UnboundedSender>, + dynamic_handler_tx: dynamic_handler::DynamicHandlerTx, transport_completion: SharedTransportCompletion, pending_replies: PendingRepliesRegistrar, #[cfg_attr( @@ -3738,9 +3890,9 @@ impl ConnectionTo { ) -> Self { Self { counterpart, - message_tx, - task_tx, - dynamic_handler_tx, + message_tx: message_tx.into(), + task_tx: task_tx.into(), + dynamic_handler_tx: dynamic_handler_tx.into(), transport_completion, pending_replies, protocol_mode, @@ -4212,13 +4364,17 @@ impl ConnectionTo { } let role_id = peer.role_id(); let remote_style = self.counterpart.remote_style(peer); - let cancellation = + let mut cancellation = SentRequestCancellation::new(self.message_tx.clone(), remote_style, id.clone()); + if self.message_tx.budget().is_some() { + cancellation.pending = Some(self.pending_replies.clone()); + } if self.is_incoming_closing() { cancellation.disarm(); drop(response_tx.send(ResponsePayload { result: Err(incoming_transport_closed_error(&method)), ack_tx: None, + charges: vec![], })); return SentRequest::new( id, @@ -4231,8 +4387,12 @@ impl ConnectionTo { .map(move |json| ::from_value(&method, json)); } - match request.to_untyped_message() { - Ok(untyped) => { + match self.message_tx.reserve().and_then(|charge| { + request + .to_untyped_message() + .map(|untyped| (untyped, charge)) + }) { + Ok((untyped, charge)) => { // Register before enqueueing so incoming EOF can fail every // observable request before close callbacks begin. The // outgoing actor checks that the registration still exists @@ -4258,12 +4418,10 @@ impl ConnectionTo { readiness, }; - if let Err(error) = self.message_tx.unbounded_send(message) { + if let Err(error) = self.message_tx.send_with_charge(message, charge) { cancellation.disarm(); - let OutgoingMessage::Request { id, method, .. } = error.into_inner() else { - unreachable!(); - }; + drop(error); if let Some(pending_reply) = self.pending_replies.remove(&id) { if self.is_incoming_closing() { @@ -4287,6 +4445,7 @@ impl ConnectionTo { "failed to create untyped request for `{method}`: {err}" ))), ack_tx: None, + charges: vec![], }) .unwrap(); } @@ -4351,16 +4510,17 @@ impl ConnectionTo { original_method = notification.method(), "send_notification_to" ); + let charge = self.message_tx.reserve()?; let transformed = remote_style.transform_outgoing_message(notification)?; tracing::debug!( transformed_method = %transformed.method, "send_notification_to transformed" ); - send_raw_message( - &self.message_tx, + self.message_tx.send_with_charge( OutgoingMessage::Notification { untyped: transformed, }, + charge, ) } @@ -4413,6 +4573,11 @@ impl ConnectionTo { &self, handler: impl HandleDispatchFrom + 'static, ) -> Result, crate::Error> { + let lifetime_charge = self + .message_tx + .budget() + .map(crate::bounded::Budget::task) + .transpose()?; let uuid = Uuid::new_v4(); let active = Arc::new(AtomicBool::new(true)); self.dynamic_handler_tx @@ -4421,6 +4586,7 @@ impl ConnectionTo { Box::new(GuardedDynamicHandler { active: active.clone(), handler, + _charge: lifetime_charge, }), )) .map_err(crate::util::internal_error)?; @@ -4464,6 +4630,7 @@ impl ConnectionTo { struct GuardedDynamicHandler { active: Arc, handler: Handler, + _charge: Option, } impl HandleDispatchFrom for GuardedDynamicHandler @@ -4602,6 +4769,7 @@ pub struct Responder { dyn FnOnce( Result, Option, + Option, ) -> Result<(), crate::Error> + Send, >, @@ -4686,9 +4854,8 @@ impl Responder { cancellation, destination, send_fn: Box::new( - move |response: Result, receipt| { - send_raw_message( - &message_tx, + move |response: Result, receipt, charge| { + message_tx.send_with_charge( OutgoingMessage::Response { id: id_clone, method: method_clone, @@ -4696,6 +4863,7 @@ impl Responder { destination: send_destination, receipt, }, + charge, ) }, ), @@ -4774,9 +4942,9 @@ impl Responder { id: self.id, cancellation: self.cancellation, destination: self.destination, - send_fn: Box::new(move |input: Result, receipt| { + send_fn: Box::new(move |input: Result, receipt, charge| { let t_value = wrap_fn(&method, input); - (self.send_fn)(t_value, receipt) + (self.send_fn)(t_value, receipt, charge) }), drop_guard: self.drop_guard, } @@ -4788,8 +4956,9 @@ impl Responder { response: Result, ) -> Result<(), crate::Error> { tracing::debug!(id = ?self.id, "respond called"); + let charge = self.drop_guard.message_tx.reserve()?; self.drop_guard.disarm(); - (self.send_fn)(response, None) + (self.send_fn)(response, None, charge) } /// Respond to the JSON-RPC request with either a value (`Ok`) or an error (`Err`) @@ -4817,9 +4986,10 @@ impl Responder { response: Result, ) -> Result { tracing::debug!(id = ?self.id, "tracked respond called"); + let charge = self.drop_guard.message_tx.reserve()?; let (sender, receipt) = ResponseReceiptSender::channel(); self.drop_guard.disarm(); - (self.send_fn)(response, Some(sender))?; + (self.send_fn)(response, Some(sender), charge)?; Ok(receipt) } @@ -4921,7 +5091,10 @@ impl ResponseRouter { let reply_target = ResponseReplyTarget { id: id.clone(), method: method.clone(), - sender: Arc::new(Mutex::new(Some(sender))), + state: Arc::new(Mutex::new(ResponseReplyState { + sender: Some(sender), + charges: vec![], + })), ordering, dispatch, }; @@ -5661,6 +5834,7 @@ struct SentRequestCancellation { remote_style: crate::role::RemoteStyle, request_id: RequestId, disarm: SentRequestCancellationDisarm, + pending: Option, } impl SentRequestCancellation { @@ -5674,6 +5848,7 @@ impl SentRequestCancellation { remote_style, request_id, disarm: SentRequestCancellationDisarm::new(), + pending: None, } } @@ -5692,11 +5867,13 @@ impl SentRequestCancellation { // Build the notification lazily: most requests are never cancelled, // so this avoids serializing a notification per outgoing request. + let charge = self.message_tx.reserve()?; let untyped = self.remote_style.transform_outgoing_message( crate::schema::v1::CancelRequestNotification::new(self.request_id.clone()), )?; - send_raw_message(&self.message_tx, OutgoingMessage::Notification { untyped }) + self.message_tx + .send_with_charge(OutgoingMessage::Notification { untyped }, charge) } } @@ -5705,6 +5882,9 @@ impl Drop for SentRequestCancellation { if let Err(error) = self.send() { tracing::debug!(?error, "failed to auto-cancel dropped request"); } + if let Some(pending) = &self.pending { + drop(pending.remove(&self.request_id)); + } } } @@ -5784,7 +5964,7 @@ impl SentRequest { fn new( id: RequestId, method: String, - task_tx: mpsc::UnboundedSender, + task_tx: impl Into, response_rx: oneshot::Receiver, cancellation: SentRequestCancellation, response_ordering: ResponseOrdering, @@ -5793,7 +5973,7 @@ impl SentRequest { id, method, response_rx, - task_tx, + task_tx: task_tx.into(), to_result: Box::new(Ok), cancellation, response_ordering, @@ -6038,7 +6218,11 @@ impl SentRequest { .await; match response { - Ok(ResponsePayload { result, ack_tx }) => { + Ok(ResponsePayload { + result, + ack_tx, + charges: _charges, + }) => { // Convert the result using to_result for Ok values let typed_result = match result { Ok(json_value) => to_result(json_value), @@ -6142,6 +6326,7 @@ impl SentRequest { Ok(ResponsePayload { result: Ok(json_value), ack_tx, + charges: _charges, }) => { // Blocking consumers ack before converting or returning the // value, so dispatch can continue while the caller processes it. @@ -6156,6 +6341,7 @@ impl SentRequest { Ok(ResponsePayload { result: Err(err), ack_tx, + charges: _charges, }) => { if let Some(tx) = ack_tx { let _ = tx.send(()); @@ -6704,6 +6890,244 @@ impl ConnectTo for Channel { mod tests { use super::*; + fn bounded_test_connection( + limits: crate::ChannelLimits, + ) -> ( + ConnectionTo, + impl Future>, + crate::BoundedChannel, + ) { + let (channel, peer) = crate::BoundedChannel::duplex(limits).unwrap(); + let (cx, future) = crate::UntypedRole + .builder() + .into_connection_and_future(channel, async |_| future::pending().await); + (cx, future, peer) + } + + fn small_bounded_limits() -> crate::ChannelLimits { + crate::ChannelLimits { + max_frame_bytes: 1024, + max_buffered_bytes: 4096, + max_buffered_frames: 4, + max_pending_requests: 2, + max_tasks: 3, + } + } + + #[test] + fn bounded_successful_driver_preserves_admitted_frames_for_transport_drain() { + let (channel, mut peer) = crate::BoundedChannel::duplex(small_bounded_limits()).unwrap(); + let shutdown = channel.tx.clone(); + let (cx, driver) = crate::UntypedRole + .builder() + .into_connection_and_future(channel, async |_| Ok(())); + cx.send_notification(UntypedMessage::new("notice", serde_json::json!({})).unwrap()) + .unwrap(); + futures::executor::block_on(driver).unwrap(); + shutdown.close_channel(); + assert!( + cx.send_notification(UntypedMessage::new("late", serde_json::json!({})).unwrap()) + .is_err() + ); + assert!(shutdown.failure().now_or_never().is_none()); + let frame = peer.rx.next().now_or_never().unwrap().unwrap(); + assert!(matches!(frame.decode(), TransportFrame::Single(_))); + assert!(peer.rx.next().now_or_never().unwrap().is_none()); + drop(frame); + assert_eq!(shutdown.budget.snapshot(), (0, 0, 0)); + } + + #[test] + fn bounded_accepted_invalid_notification_fails_driver_and_signals_terminal() { + let (cx, driver, peer) = bounded_test_connection(small_bounded_limits()); + cx.send_notification(UntypedMessage { + method: "invalid".into(), + params: serde_json::json!(42), + }) + .unwrap(); + assert!(futures::executor::block_on(driver).is_err()); + assert!(peer.tx.failure().now_or_never().is_some()); + assert_eq!(cx.message_tx.budget().unwrap().snapshot(), (0, 0, 0)); + } + + #[test] + fn legacy_invalid_notification_still_does_not_terminate_driver() { + let (channel, _peer) = Channel::duplex(); + let (cx, driver) = crate::UntypedRole + .builder() + .into_connection_and_future(channel, async |_| { + future::pending::>().await + }); + cx.send_notification(UntypedMessage { + method: "invalid".into(), + params: serde_json::json!(42), + }) + .unwrap(); + let mut driver = Box::pin(driver); + assert!(driver.as_mut().now_or_never().is_none()); + } + + #[cfg(feature = "unstable_protocol_v2")] + #[test] + fn bounded_accepted_notification_before_initialize_fails_terminally() { + let (channel, peer) = crate::BoundedChannel::duplex(small_bounded_limits()).unwrap(); + let (cx, driver) = Client.v2().into_connection_and_future(channel, async |_| { + future::pending::>().await + }); + cx.send_notification(UntypedMessage::new("session/update", serde_json::json!({})).unwrap()) + .unwrap(); + assert!(futures::executor::block_on(driver).is_err()); + assert!(peer.tx.failure().now_or_never().is_some()); + } + + #[test] + fn outgoing_sender_stays_pointer_sized() { + assert_eq!( + std::mem::size_of::(), + std::mem::size_of::() + ); + } + + #[test] + fn bounded_protocol_producers_admit_before_driver_poll() { + let (cx, driver, _peer) = bounded_test_connection(small_bounded_limits()); + let budget = cx.message_tx.budget().cloned().unwrap(); + for _ in 0..4 { + cx.send_notification(UntypedMessage::new("notice", serde_json::json!({})).unwrap()) + .unwrap(); + } + assert_eq!(budget.snapshot(), (4, 4096, 1)); + assert!( + cx.send_notification(UntypedMessage::new("notice", serde_json::json!({})).unwrap()) + .is_err() + ); + assert!(futures::executor::block_on(driver).is_err()); + assert_eq!(budget.snapshot(), (0, 0, 0)); + assert!( + cx.send_notification(UntypedMessage::new("notice", serde_json::json!({})).unwrap()) + .is_err() + ); + } + + #[test] + fn bounded_unpolled_driver_drop_releases_tasks_frames_and_pending() { + let (cx, driver, _peer) = bounded_test_connection(small_bounded_limits()); + let budget = cx.message_tx.budget().cloned().unwrap(); + let request = + cx.send_request(UntypedMessage::new("request", serde_json::json!({})).unwrap()); + cx.spawn(async { future::pending().await }).unwrap(); + assert_eq!(budget.snapshot(), (1, 1024, 2)); + drop(driver); + assert_eq!(budget.snapshot(), (0, 0, 0)); + assert!(futures::executor::block_on(request.block_task()).is_err()); + assert!(cx.spawn(async { Ok(()) }).is_err()); + } + + #[test] + fn bounded_pending_and_task_limits_fail_fast() { + let (cx, driver, _peer) = bounded_test_connection(small_bounded_limits()); + let first = cx.send_request(UntypedMessage::new("request", serde_json::json!({})).unwrap()); + let second = + cx.send_request(UntypedMessage::new("request", serde_json::json!({})).unwrap()); + let third = cx.send_request(UntypedMessage::new("request", serde_json::json!({})).unwrap()); + assert!(futures::executor::block_on(third.block_task()).is_err()); + drop((first, second, driver)); + let (cx, driver, _peer) = bounded_test_connection(small_bounded_limits()); + cx.spawn(async { future::pending().await }).unwrap(); + cx.spawn(async { future::pending().await }).unwrap(); + assert!(cx.spawn(async { future::pending().await }).is_err()); + drop(driver); + assert_eq!(cx.message_tx.budget().unwrap().snapshot(), (0, 0, 0)); + } + + #[test] + fn bounded_unconsumed_response_retains_incoming_admission() { + let (cx, driver, mut peer) = bounded_test_connection(small_bounded_limits()); + let request = + cx.send_request(UntypedMessage::new("request", serde_json::json!({})).unwrap()); + let id = request.id.clone(); + let mut driver = Box::pin(driver); + assert!(driver.as_mut().now_or_never().is_none()); + drop(peer.rx.next().now_or_never().unwrap().unwrap()); + peer.tx + .try_send(TransportFrame::Single(RawJsonRpcMessage::response( + id, + Ok(serde_json::json!({"answer": 42})), + ))) + .unwrap(); + assert!(driver.as_mut().now_or_never().is_none()); + assert_eq!(peer.tx.budget.snapshot(), (1, 1024, 0)); + assert!( + cx.pending_replies + .inner + .upgrade() + .unwrap() + .lock() + .unwrap() + .replies + .is_empty() + ); + let response = futures::executor::block_on(request.block_task()).unwrap(); + assert_eq!(response, serde_json::json!({"answer": 42})); + assert_eq!(peer.tx.budget.snapshot(), (0, 0, 0)); + drop(driver); + } + + #[test] + fn bounded_batch_producer_charges_survive_accumulation_and_handoff() { + let (channel, mut peer) = crate::BoundedChannel::duplex(small_bounded_limits()).unwrap(); + let budget = channel.tx.budget.clone(); + let (tx, rx) = mpsc::unbounded(); + let mut tx = OutgoingMessageTx::from(tx); + tx.set_budget(Some(budget.clone())); + let (mut destinations, completion) = ResponseDestination::batch(2); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + rx, + PendingReplies::default(), + crate::bounded::TransportSender::Bounded(channel.tx), + ProtocolCompat::new(ProtocolMode::disabled()), + )); + tx.unbounded_send(OutgoingMessage::Response { + id: RequestId::Str("first".into()), + method: "request".into(), + response: Ok(serde_json::json!({"first": true})), + destination: destinations.next().unwrap(), + receipt: None, + }) + .unwrap(); + assert!(actor.as_mut().now_or_never().is_none()); + assert_eq!(budget.snapshot(), (1, 1024, 0)); + tx.unbounded_send(OutgoingMessage::Response { + id: RequestId::Str("second".into()), + method: "request".into(), + response: Ok(serde_json::json!({"second": true})), + destination: destinations.next().unwrap(), + receipt: None, + }) + .unwrap(); + tx.unbounded_send(OutgoingMessage::BatchDispatchComplete { completion }) + .unwrap(); + assert!(actor.as_mut().now_or_never().is_none()); + let frame = peer.rx.next().now_or_never().unwrap().unwrap(); + assert!(matches!(frame.decode(), TransportFrame::Batch(_))); + assert_eq!(budget.snapshot(), (2, 2048, 0)); + drop(frame); + assert_eq!(budget.snapshot(), (0, 0, 0)); + drop((actor, tx, destinations)); + } + + #[test] + fn bounded_sent_request_drop_removes_pending_registration() { + let (cx, driver, _peer) = bounded_test_connection(small_bounded_limits()); + let request = + cx.send_request(UntypedMessage::new("request", serde_json::json!({})).unwrap()); + let registry = cx.pending_replies.inner.upgrade().unwrap(); + assert_eq!(registry.lock().unwrap().replies.len(), 1); + drop(request); + assert!(registry.lock().unwrap().replies.is_empty()); + drop(driver); + } + #[cfg(feature = "unstable_protocol_v2")] fn connection_with_task_receiver() -> ( ConnectionTo, diff --git a/src/agent-client-protocol/src/jsonrpc/dynamic_handler.rs b/src/agent-client-protocol/src/jsonrpc/dynamic_handler.rs index ff88b0b8..4765ae88 100644 --- a/src/agent-client-protocol/src/jsonrpc/dynamic_handler.rs +++ b/src/agent-client-protocol/src/jsonrpc/dynamic_handler.rs @@ -37,6 +37,7 @@ impl> DynHandleDispatchFro /// Messages used to add/remove dynamic handlers pub(crate) enum DynamicHandlerMessage { + Admitted(Box, crate::bounded::Charge), AddDynamicHandler(Uuid, Box>), RemoveDynamicHandler(Uuid), /// Marks the end of the registrations queued by an ordered response callback. @@ -50,6 +51,7 @@ pub(crate) enum DynamicHandlerMessage { impl std::fmt::Debug for DynamicHandlerMessage { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { + Self::Admitted(message, _) => message.fmt(f), Self::AddDynamicHandler(arg0, arg1) => f .debug_tuple("AddDynamicHandler") .field(arg0) @@ -64,3 +66,39 @@ impl std::fmt::Debug for DynamicHandlerMessage { } } } + +impl DynamicHandlerMessage { + pub(super) fn unpack(self) -> (Self, Option) { + match self { + Self::Admitted(message, charge) => (*message, Some(charge)), + message => (message, None), + } + } +} +#[derive(Clone, Debug)] +pub(super) struct DynamicHandlerTx { + tx: futures::channel::mpsc::UnboundedSender>, + pub(super) budget: Option, +} +impl From>> + for DynamicHandlerTx +{ + fn from(tx: futures::channel::mpsc::UnboundedSender>) -> Self { + Self { tx, budget: None } + } +} +impl DynamicHandlerTx { + pub(super) fn unbounded_send( + &self, + message: DynamicHandlerMessage, + ) -> Result<(), crate::Error> { + let message = if let Some(budget) = &self.budget { + DynamicHandlerMessage::Admitted(Box::new(message), budget.reserve()?) + } else { + message + }; + self.tx + .unbounded_send(message) + .map_err(crate::util::internal_error) + } +} diff --git a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs index a4ab7a3d..45d28bdc 100644 --- a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs @@ -59,7 +59,7 @@ impl IncomingHandlers { pub(super) async fn incoming_protocol_actor( counterpart: Counterpart, connection: &ConnectionTo, - transport_rx: mpsc::UnboundedReceiver, + transport_rx: impl futures::Stream> + Unpin, dynamic_handler_rx: mpsc::UnboundedReceiver>, pending_replies: PendingReplies, handlers: IncomingHandlers< @@ -77,7 +77,7 @@ pub(super) async fn incoming_protocol_actor( // transport EOF as an explicit event so the other, connection-internal // streams cannot hide it. let transport_with_close = futures::StreamExt::chain( - transport_rx.map(IncomingProtocolMsg::Transport), + transport_rx.map(|frame| IncomingProtocolMsg::Transport(frame.into())), stream::iter([IncomingProtocolMsg::TransportClosed]), ); let mut my_rx = @@ -86,6 +86,7 @@ pub(super) async fn incoming_protocol_actor( let mut dynamic_handlers: FxHashMap>> = FxHashMap::default(); let mut pending_messages: Vec = vec![]; + let mut pending_frame_charges = vec![]; let request_cancellations = super::RequestCancellationRegistry::new(); let mut on_close = Some(on_close); @@ -100,6 +101,10 @@ pub(super) async fn incoming_protocol_actor( }; message }; + let (message_result, _control_charge) = message_result.unpack(); + if pending_messages.is_empty() { + pending_frame_charges.clear(); + } tracing::trace!(message = ?message_result, actor = "incoming_protocol_actor"); match message_result { IncomingProtocolMsg::TransportClosed => { @@ -127,8 +132,15 @@ pub(super) async fn incoming_protocol_actor( .await?; } - IncomingProtocolMsg::Transport(frame) => { - let (entries, batch_completion) = frame_entries(frame); + IncomingProtocolMsg::Transport(incoming) => { + let charges = incoming.charges; + let (entries, batch_completion) = frame_entries(incoming.frame); + let pending_before = pending_messages.len(); + for (_, destination) in &entries { + if let Some(destination) = destination { + destination.retain_incoming(&charges); + } + } for (message, destination) in entries { match message { Ok(RawJsonRpcMessage::Request(request)) => { @@ -217,6 +229,9 @@ pub(super) async fn incoming_protocol_actor( .incoming_response(&pending_reply.method, result); let (dispatch, response_dispatch) = dispatch_from_response(id, pending_reply, result); + if let Dispatch::Response(_, router) = &dispatch { + router.reply_target.retain_incoming(&charges); + } dispatch_dispatch( counterpart.clone(), connection, @@ -245,6 +260,7 @@ pub(super) async fn incoming_protocol_actor( "dynamic-handler stream closed before its barrier", )); }; + let (message, _control_charge) = message.unpack(); match message { IncomingProtocolMsg::DynamicHandler( DynamicHandlerMessage::Barrier, @@ -287,6 +303,9 @@ pub(super) async fn incoming_protocol_actor( } } } + if pending_messages.len() > pending_before { + pending_frame_charges.extend(charges); + } if let Some(completion) = batch_completion { send_raw_message( &connection.message_tx, @@ -305,7 +324,9 @@ async fn handle_dynamic_handler_message( dynamic_handlers: &mut FxHashMap>>, pending_messages: &mut Vec, ) -> Result<(), crate::Error> { + let (message, _charge) = message.unpack(); match message { + DynamicHandlerMessage::Admitted(..) => unreachable!("nested admission"), DynamicHandlerMessage::AddDynamicHandler(uuid, mut handler) => { // Before adding the new handler, give it a chance to process // any pending messages. @@ -357,11 +378,23 @@ async fn handle_dynamic_handler_message( #[derive(Debug)] enum IncomingProtocolMsg { - Transport(TransportFrame), + Transport(crate::bounded::IncomingFrame), TransportClosed, DynamicHandler(DynamicHandlerMessage), } +impl IncomingProtocolMsg { + fn unpack(self) -> (Self, Option) { + match self { + Self::DynamicHandler(message) => { + let (message, charge) = message.unpack(); + (Self::DynamicHandler(message), charge) + } + message => (message, None), + } + } +} + fn frame_entries( frame: TransportFrame, ) -> ( diff --git a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs index 0d115c36..cf41e8ef 100644 --- a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs @@ -11,7 +11,67 @@ use crate::jsonrpc::{ }; use crate::schema::v1::RequestId; -pub type OutgoingMessageTx = mpsc::UnboundedSender; +#[derive(Clone, Debug)] +pub struct OutgoingMessageTx { + state: Arc, +} +#[derive(Debug)] +struct OutgoingMessageState { + tx: mpsc::UnboundedSender, + budget: Option, +} +impl From> for OutgoingMessageTx { + fn from(tx: mpsc::UnboundedSender) -> Self { + Self { + state: Arc::new(OutgoingMessageState { tx, budget: None }), + } + } +} +impl OutgoingMessageTx { + pub(super) fn budget(&self) -> Option<&crate::bounded::Budget> { + self.state.budget.as_ref() + } + pub(super) fn set_budget(&mut self, budget: Option) { + Arc::get_mut(&mut self.state) + .expect("configure sender before cloning") + .budget = budget; + } + pub(super) fn reserve(&self) -> Result, crate::Error> { + self.state + .budget + .as_ref() + .map(crate::bounded::Budget::reserve) + .transpose() + } + pub(super) fn unbounded_send(&self, message: OutgoingMessage) -> Result<(), crate::Error> { + self.send_with_charge(message, self.reserve()?) + } + pub(super) fn send_with_charge( + &self, + mut message: OutgoingMessage, + charge: Option, + ) -> Result<(), crate::Error> { + if let Some(budget) = &self.state.budget { + budget.check_admission()?; + let charge = charge.expect("bounded producers reserve before conversion"); + message + .normalize(budget) + .map_err(|_| budget.fail("outgoing protocol message exceeds bounded limit"))?; + self.state + .tx + .unbounded_send(OutgoingMessage::Admitted { + message: Box::new(message), + charge, + }) + .map_err(|_| budget.fail("outgoing protocol actor closed")) + } else { + self.state + .tx + .unbounded_send(message) + .map_err(crate::util::internal_error) + } + } +} pub(crate) fn send_raw_message( tx: &OutgoingMessageTx, @@ -43,10 +103,10 @@ impl Drop for ResponseReceiptRegistry { } fn enqueue_completed_response( - transport_tx: &mpsc::UnboundedSender, + transport_tx: &crate::bounded::TransportSender, completed: CompletedResponseFrame, ) -> Result<(), crate::Error> { - match transport_tx.unbounded_send(completed.frame) { + match transport_tx.send(completed.frame, completed.charges) { Ok(()) => { for receipt in completed.receipts { receipt.resolve(Ok(())); @@ -73,17 +133,38 @@ fn enqueue_completed_response( pub(super) async fn outgoing_protocol_actor( mut outgoing_rx: mpsc::UnboundedReceiver, pending_replies: PendingReplies, - transport_tx: mpsc::UnboundedSender, + transport_tx: impl Into, protocol_compat: ProtocolCompat, ) -> Result<(), crate::Error> { + let transport_tx = transport_tx.into(); + let bounded = matches!(&transport_tx, crate::bounded::TransportSender::Bounded(_)); let mut drain_waiters = Vec::new(); let mut receipt_registry = ResponseReceiptRegistry::default(); while let Some(message) = outgoing_rx.next().await { + let (message, mut charges) = match message { + OutgoingMessage::Admitted { message, charge } => (*message, vec![charge]), + message => (message, vec![]), + }; + match &message { + OutgoingMessage::Response { destination, .. } + | OutgoingMessage::AbandonedBatchResponse { destination, .. } + | OutgoingMessage::UncorrelatedErrorResponse { destination, .. } => { + if let super::ResponseDestination::Batch(slot) = destination { + slot.state + .lock() + .expect("batch response accumulator mutex poisoned") + .charges + .append(&mut charges); + } + } + _ => {} + } tracing::debug!(?message, "outgoing_protocol_actor"); // Create the message to be sent over the transport let (json_rpc_message, destination, receipt) = match message { + OutgoingMessage::Admitted { .. } => unreachable!("nested admission"), OutgoingMessage::CloseAfterDraining { done } => { // Reject later sends while preserving every message that was // already accepted into this receiver's buffer. @@ -178,7 +259,10 @@ pub(super) async fn outgoing_protocol_actor( continue; } - if let Err(error) = transport_tx.unbounded_send(TransportFrame::Single(request)) { + if let Err(error) = transport_tx.send( + TransportFrame::Single(request), + std::mem::take(&mut charges), + ) { let error = crate::Error::into_internal_error(error); if let Some(pending_reply) = pending_replies.remove(&id) { pending_reply.fail(error.clone()); @@ -191,6 +275,9 @@ pub(super) async fn outgoing_protocol_actor( let messages = match protocol_compat.outgoing_notification(untyped) { Ok(messages) => messages, Err(error) => { + if bounded { + return Err(error); + } tracing::warn!( ?error, "Dropping outgoing notification after preparation failed" @@ -203,6 +290,9 @@ pub(super) async fn outgoing_protocol_actor( let message = match untyped.into_raw_jsonrpc_message(None) { Ok(message) => message, Err(error) => { + if bounded { + return Err(error); + } tracing::warn!( ?error, "Dropping outgoing notification after serialization failed" @@ -211,7 +301,10 @@ pub(super) async fn outgoing_protocol_actor( } }; transport_tx - .unbounded_send(TransportFrame::Single(message)) + .send( + TransportFrame::Single(message), + std::mem::take(&mut charges), + ) .map_err(crate::Error::into_internal_error)?; } continue; @@ -254,7 +347,8 @@ pub(super) async fn outgoing_protocol_actor( if let Some(receipt) = receipt.as_ref() { receipt_registry.register(receipt); } - if let Some(frame) = destination.complete(json_rpc_message, receipt) { + if let Some(mut frame) = destination.complete(json_rpc_message, receipt) { + frame.charges.append(&mut charges); enqueue_completed_response(&transport_tx, frame)?; } } @@ -304,9 +398,10 @@ mod tests { Ok(serde_json::Value::Null), )), receipts: vec![first_sender, second_sender], + charges: vec![], }; - assert!(enqueue_completed_response(&transport_tx, completed).is_err()); + assert!(enqueue_completed_response(&transport_tx.into(), completed).is_err()); let (first_result, second_result) = block_on(join(first_receipt, second_receipt)); assert!(first_result.is_err()); assert!(second_result.is_err()); diff --git a/src/agent-client-protocol/src/jsonrpc/task_actor.rs b/src/agent-client-protocol/src/jsonrpc/task_actor.rs index 92a05a6d..81d102eb 100644 --- a/src/agent-client-protocol/src/jsonrpc/task_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/task_actor.rs @@ -6,7 +6,16 @@ use crate::ConnectionTo; use crate::role::Role; use crate::util::process_stream_concurrently; -pub type TaskTx = mpsc::UnboundedSender; +#[derive(Clone, Debug)] +pub struct TaskTx { + tx: mpsc::UnboundedSender, + pub(super) budget: Option, +} +impl From> for TaskTx { + fn from(tx: mpsc::UnboundedSender) -> Self { + Self { tx, budget: None } + } +} #[must_use] pub(crate) struct Task { @@ -39,8 +48,17 @@ impl Task { } } - pub fn spawn(self, task_tx: &TaskTx) -> Result<(), crate::Error> { + pub fn spawn(mut self, task_tx: &TaskTx) -> Result<(), crate::Error> { + if let Some(budget) = &task_tx.budget { + let charge = budget.task()?; + self.future = async move { + let _charge = charge; + self.future.await + } + .boxed(); + } task_tx + .tx .unbounded_send(self) .map_err(crate::util::internal_error)?; Ok(()) diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index b94cddf8..ae18923e 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -243,3 +243,8 @@ macro_rules! on_receive_dispatch { |f: &mut _, dispatch, cx| Box::pin(f(dispatch, cx)) }; } + +mod bounded; +pub use bounded::{ + BoundedChannel, BoundedReceiver, BoundedSender, ChannelLimits, ChargedFrame, TransportChannel, +}; From c85cc5da56d6d55ba651bdd336e525efc153af76 Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 29 Sep 2026 01:30:15 +0100 Subject: [PATCH 5/9] fix(deps): update rustls for RUSTSEC-2026-0285 --- Cargo.lock | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 1df89af3..91c485b0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2234,9 +2234,9 @@ dependencies = [ [[package]] name = "rustls" -version = "0.23.43" +version = "0.23.45" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0283386ce02abc0151e1761d08802dfe86c173b0b494af5cbc086574e453da06" +checksum = "0d41d731c7d2f962d1ccc364cec258de3c0e93b38c2fb3ba97ac74513048d634" dependencies = [ "aws-lc-rs", "once_cell", From 62bde00c7fe0508e3d0ffaad4a4974d6ae7e7a45 Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 29 Sep 2026 02:02:14 +0100 Subject: [PATCH 6/9] fix(http): await bounded connection teardown --- .../src/client_limits.rs | 327 +++++++++++++++--- 1 file changed, 283 insertions(+), 44 deletions(-) diff --git a/src/agent-client-protocol-http/src/client_limits.rs b/src/agent-client-protocol-http/src/client_limits.rs index 2085472a..0194a55b 100644 --- a/src/agent-client-protocol-http/src/client_limits.rs +++ b/src/agent-client-protocol-http/src/client_limits.rs @@ -100,9 +100,12 @@ impl HttpClientLimits { /// [`HttpClient::with_limits`]; legacy channel extraction fails closed. /// /// There are no background SSE tasks, observer mailboxes, or hidden unbounded -/// bridges. The transport future directly polls all streams and POSTs. Dropping -/// it releases local resources; it does not wait for an HTTP DELETE or promise -/// peer cleanup. No uncertain or accepted POST is ever retried. +/// bridges. The transport future directly polls all streams and POSTs. On exit, +/// it awaits one best-effort HTTP DELETE, bounded to five seconds, for an admitted +/// connection ID. Dropping it releases local streams and POSTs and, if teardown +/// has not started, schedules that DELETE when a Tokio runtime is available. +/// Cancellation can interrupt DELETE and never guarantees peer cleanup. No +/// uncertain or accepted POST is ever retried. pub struct BoundedHttpClient { client: HttpClient, limits: HttpClientLimits, @@ -135,7 +138,8 @@ impl BoundedHttpClient { } /// Extract the charged channel without a legacy adapter. Poll the returned - /// future to drive HTTP; dropping it cancels all local HTTP work. + /// future to drive HTTP; dropping it cancels streams and POSTs. Connection + /// cleanup is best effort, as described on [`BoundedHttpClient`]. #[must_use] pub fn into_bounded_channel_and_future( self, @@ -578,6 +582,182 @@ mod tests { (format!("http://{address}/acp"), count, server) } + // Loopback HTTP boundary; DELETE headers are withheld until explicitly released. + async fn teardown_fixture( + initialize_body: &'static str, + ) -> ( + String, + Arc, + Arc, + tokio::task::JoinHandle<()>, + ) { + use axum::{Router, body::Body, response::Response, routing::post}; + let started = Arc::new(tokio::sync::Notify::new()); + let release = Arc::new(tokio::sync::Semaphore::new(0)); + let delete_started = started.clone(); + let delete_release = release.clone(); + let app = Router::new().route( + "/acp", + post(move || async move { + Response::builder() + .header(HEADER_CONNECTION_ID, "teardown-test") + .body(Body::from(initialize_body)) + .unwrap() + }) + .get(|| async { + Response::builder() + .header("Content-Type", "text/event-stream") + .body(Body::from_stream(futures::stream::pending::< + Result, + >())) + .unwrap() + }) + .delete(move |headers: axum::http::HeaderMap| { + let started = delete_started.clone(); + let release = delete_release.clone(); + async move { + assert_eq!(headers.get(HEADER_CONNECTION_ID).unwrap(), "teardown-test"); + started.notify_one(); + release.acquire().await.unwrap().forget(); + // Headers complete DELETE; an infinite body must not delay it. + Response::builder() + .body(Body::from_stream(futures::stream::pending::< + Result, + >())) + .unwrap() + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + (format!("http://{address}/acp"), started, release, server) + } + + const INITIALIZED: &str = r#"{"jsonrpc":"2.0","id":0,"result":{}}"#; + + #[tokio::test] + async fn graceful_eof_waits_for_delete_headers_not_body() { + let (url, started, release, server) = teardown_fixture(INITIALIZED).await; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(HttpClientLimits::default()) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + assert!(channel.rx.next().await.is_some()); + channel.tx.close_channel(); + tokio::time::timeout(std::time::Duration::from_secs(3), started.notified()) + .await + .unwrap(); + assert!(!driver.is_finished(), "EOF must await DELETE headers"); + release.add_permits(1); + tokio::time::timeout(std::time::Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .unwrap(); + server.abort(); + } + + #[tokio::test] + async fn malformed_initialize_closes_admission_and_awaits_delete() { + let (url, started, release, server) = teardown_fixture("not JSON").await; + let (channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(HttpClientLimits::default()) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + tokio::time::timeout(std::time::Duration::from_secs(3), started.notified()) + .await + .unwrap(); + assert!(channel.tx.try_send(initialize()).is_err()); + assert!(!driver.is_finished()); + release.add_permits(1); + let error = tokio::time::timeout(std::time::Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .unwrap_err(); + assert!(error.to_string().contains("initialize response")); + server.abort(); + } + + #[tokio::test] + async fn cancellation_closes_locally_without_waiting_for_remote_delete() { + let (url, started, release, server) = teardown_fixture(INITIALIZED).await; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(HttpClientLimits::default()) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + assert!(channel.rx.next().await.is_some()); + driver.abort(); + assert!(driver.await.unwrap_err().is_cancelled()); + assert!(channel.tx.try_send(initialize()).is_err()); + tokio::time::timeout(std::time::Duration::from_secs(3), started.notified()) + .await + .unwrap(); + // Local cancellation completed while the server still withholds DELETE + // completion. Best-effort dispatch is not a remote-deletion guarantee. + assert_eq!(release.available_permits(), 0); + release.add_permits(1); + server.abort(); + } + + #[tokio::test] + async fn cancellation_during_graceful_delete_preserves_admitted_response() { + let (url, started, _release, server) = teardown_fixture(INITIALIZED).await; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(HttpClientLimits::default()) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + channel.tx.close_channel(); + let driver = tokio::spawn(transport); + tokio::time::timeout(std::time::Duration::from_secs(3), started.notified()) + .await + .unwrap(); + driver.abort(); + assert!(driver.await.unwrap_err().is_cancelled()); + let response = channel + .rx + .next() + .await + .expect("admitted response survives graceful cleanup cancellation"); + assert_eq!(response.decode(), TransportFrame::parse_json(INITIALIZED)); + server.abort(); + } + + #[tokio::test] + async fn delete_without_response_headers_has_finite_deadline() { + let (url, started, _release, server) = teardown_fixture(INITIALIZED).await; + let (mut channel, transport) = HttpClient::with_endpoint(url) + .unwrap() + .with_limits(HttpClientLimits::default()) + .unwrap() + .into_bounded_channel_and_future(); + channel.tx.try_send(initialize()).unwrap(); + let driver = tokio::spawn(transport); + assert!(channel.rx.next().await.is_some()); + channel.tx.close_channel(); + tokio::time::timeout(std::time::Duration::from_secs(3), started.notified()) + .await + .unwrap(); + tokio::time::timeout(std::time::Duration::from_secs(10), driver) + .await + .unwrap() + .unwrap() + .unwrap(); + server.abort(); + } + #[tokio::test] async fn non_replayable_posts_ignore_redirect_and_retry_policies() { use axum::{Router, body::Body, response::Response, routing::post}; @@ -1002,6 +1182,60 @@ fn launch( ); } +// Single-owner teardown: no registry or task per POST. The connection metadata +// remains charged even if cancellation transfers cleanup to a background task. +struct ConnectionCleanup { + client: HttpClient, + connection: Option<(String, Arc)>, +} + +impl ConnectionCleanup { + async fn close(&mut self) { + if let Some(connection) = self.connection.take() { + Self::send_close( + self.client.http.clone(), + self.client.endpoint.clone(), + connection, + ) + .await; + } + } + + async fn send_close( + http: reqwest::Client, + endpoint: url::Url, + (connection, lease): (String, Arc), + ) { + // Do not consume the response body: it may be arbitrarily large or never + // end. The request deadline also bounds a peer that never sends headers. + if let Err(error) = http + .delete(endpoint) + .header(HEADER_CONNECTION_ID, connection) + .timeout(std::time::Duration::from_secs(5)) + .send() + .await + { + debug!("bounded HTTP DELETE failed (ignored): {error}"); + } + drop(lease); + } +} + +impl Drop for ConnectionCleanup { + fn drop(&mut self) { + let Some(connection) = self.connection.take() else { + return; + }; + if let Ok(runtime) = tokio::runtime::Handle::try_current() { + drop(runtime.spawn(Self::send_close( + self.client.http.clone(), + self.client.endpoint.clone(), + connection, + ))); + } + } +} + async fn run_bounded( client: HttpClient, limits: HttpClientLimits, @@ -1022,20 +1256,34 @@ async fn run_bounded( sender: channel.tx.clone(), armed: true, }; + let mut cleanup = ConnectionCleanup { + client, + connection: None, + }; let failed = terminal.sender.failure(); - let running = run_bounded_inner(client, limits, channel).boxed(); - match futures::future::select(failed, running).await { - futures::future::Either::Left((error, _)) => Err(error), - futures::future::Either::Right((result, _)) => { - // Preserve already-admitted inbound responses on graceful EOF. - terminal.armed = result.is_err(); + let running = run_bounded_inner(&mut cleanup, limits, channel).boxed(); + let result = match futures::future::select(failed, running).await { + futures::future::Either::Left((error, running)) => { + drop(running); + Err(error) + } + futures::future::Either::Right((result, failed)) => { + drop(failed); result } + }; + // Preserve already-admitted inbound responses on graceful EOF, including + // cancellation during teardown. Errors close local admission before DELETE. + terminal.armed = result.is_err(); + if result.is_err() { + drop(terminal.sender.fail("bounded HTTP transport ended")); } + cleanup.close().await; + result } async fn run_bounded_inner( - client: HttpClient, + cleanup: &mut ConnectionCleanup, limits: HttpClientLimits, channel: BoundedChannel, ) -> Result<(), AcpError> { @@ -1044,6 +1292,7 @@ async fn run_bounded_inner( Sse(Box>), Post(Result), } + let client = &cleanup.client; let BoundedChannel { tx: incoming, rx: mut outgoing, @@ -1082,24 +1331,17 @@ async fn run_bounded_inner( return Err(failure("initialize redirect is not supported")); } let status = response.status(); - let _connection_metadata = response - .headers() - .get(HEADER_CONNECTION_ID) - .map(|header| budget.reserve(header.as_bytes().len())) - .transpose()?; - let connection = response - .headers() - .get(HEADER_CONNECTION_ID) - .map(|header| { - if header.as_bytes().len() > limits.channel.max_frame_bytes { - return Err(failure("connection header exceeds frame limit")); - } - header - .to_str() - .map(String::from) - .map_err(|_| failure("invalid connection header")) - }) - .transpose()?; + if let Some(header) = response.headers().get(HEADER_CONNECTION_ID) { + if header.as_bytes().len() > limits.channel.max_frame_bytes { + return Err(failure("connection header exceeds frame limit")); + } + let connection = header + .to_str() + .map_err(|_| failure("invalid connection header"))?; + let lease = budget.reserve(header.as_bytes().len())?; + // Capture before any body await so failed initialization is cleaned up. + cleanup.connection = Some((connection.to_owned(), lease)); + } let body = capped_body(response, limits.max_response_bytes).await?; if !status.is_success() { return Err(failure(format!("initialize HTTP {status}"))); @@ -1122,9 +1364,12 @@ async fn run_bounded_inner( if rejected { return Ok(()); } - let connection = connection.ok_or_else(|| failure("missing connection ID"))?; + let (connection, _) = cleanup + .connection + .as_ref() + .ok_or_else(|| failure("missing connection ID"))?; let mut streams = Streams::new(); - streams.open(&client, &connection, None, &limits, &budget)?; + streams.open(client, connection, None, &limits, &budget)?; let mut pending = VecDeque::::new(); let mut ordered = Lane::default(); let mut responses = Lane::default(); @@ -1135,8 +1380,8 @@ async fn run_bounded_inner( &mut ordered, false, &streams, - &client, - &connection, + client, + connection, limits.max_response_bytes, &mut posts, ); @@ -1144,8 +1389,8 @@ async fn run_bounded_inner( &mut responses, true, &streams, - &client, - &connection, + client, + connection, limits.max_response_bytes, &mut posts, ); @@ -1238,13 +1483,7 @@ async fn run_bounded_inner( } } for session in &bookkeeping.session_ids { - streams.open( - &client, - &connection, - Some(session.clone()), - &limits, - &budget, - )?; + streams.open(client, connection, Some(session.clone()), &limits, &budget)?; } lane.queue.push_back(Post { frame: charged, @@ -1288,8 +1527,8 @@ async fn run_bounded_inner( result.get("sessionId").and_then(|v| v.as_str()) { streams.open( - &client, - &connection, + client, + connection, Some(session.to_owned()), &limits, &budget, From 1dedc0148b6c6e4c7fa497636b99b70c8fc76da7 Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 29 Sep 2026 02:02:55 +0100 Subject: [PATCH 7/9] test(http): compare retained response wire encoding --- src/agent-client-protocol-http/src/client_limits.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/agent-client-protocol-http/src/client_limits.rs b/src/agent-client-protocol-http/src/client_limits.rs index 0194a55b..d829ee43 100644 --- a/src/agent-client-protocol-http/src/client_limits.rs +++ b/src/agent-client-protocol-http/src/client_limits.rs @@ -731,7 +731,7 @@ mod tests { .next() .await .expect("admitted response survives graceful cleanup cancellation"); - assert_eq!(response.decode(), TransportFrame::parse_json(INITIALIZED)); + assert_eq!(response.decode().to_json().unwrap(), INITIALIZED); server.abort(); } From d5c4ff74166f5f0117500184e544239774a659c6 Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 29 Sep 2026 02:48:44 +0100 Subject: [PATCH 8/9] fix(http): add allocation-free graceful bounded server delete --- md/http-transport.md | 17 + src/agent-client-protocol-http/README.md | 8 + .../src/bounded_server.rs | 230 +++++++-- .../src/graceful_delete_tests.rs | 439 ++++++++++++++++++ 4 files changed, 654 insertions(+), 40 deletions(-) create mode 100644 src/agent-client-protocol-http/src/graceful_delete_tests.rs diff --git a/md/http-transport.md b/md/http-transport.md index d1bd3ad7..df3470ae 100644 --- a/md/http-transport.md +++ b/md/http-transport.md @@ -208,6 +208,23 @@ connection so callbacks cannot remain indefinitely blocked behind saturated request bodies. Cancellation/drop releases local reservations; client teardown does not guarantee a remote DELETE completed. +By default DELETE aborts the server connection. To drain accepted work instead, +call `.with_graceful_delete(std::time::Duration::from_secs(30))` on the bounded +server before `into_router()`. DELETE atomically seals inbound HTTP admission +without reserving a POST body or core frame. Later POSTs return 410 before body +admission; already-reading POSTs recheck the seal before enqueue. The existing +connection task continues to process accepted frames even if the DELETE waiter +is canceled. A 202 response requires successful component-future completion +and clean output EOF, not merely an empty queue or a terminal-failure EOF. +This does not acknowledge consumption by the remote HTTP peer. + +Each DELETE waits at most the configured duration. A 503 reports a deadline or +failure; a deadline leaves the connection closing and accepted work running, +never reopens admission, and retains its capacity slot until work finishes. +Repeated DELETE joins the same drain while the connection exists; after actual +completion removes it, subsequent DELETE returns 404. Existing subscribed SSE +streams can consume queued output through EOF after clean completion. + Encoded-byte budgets are **not hard peak-heap limits**. Parsed JSON and bounded serialization scratch add overhead; application conversion hooks can allocate intermediate values before capped normalization. Allocator capacity, arbitrary diff --git a/src/agent-client-protocol-http/README.md b/src/agent-client-protocol-http/README.md index b4934c1c..d7898c56 100644 --- a/src/agent-client-protocol-http/README.md +++ b/src/agent-client-protocol-http/README.md @@ -25,6 +25,14 @@ Use `AcpHttpServer::new_bounded(factory, ServerLimits::default())?` and bounded transport path. Existing constructors and `ServerOptions` literals keep their compatibility behavior; legacy transports remain unbounded. +The bounded server also offers `.with_graceful_delete(Duration)` before +`into_router()`. It seals POST admission without consuming body/frame capacity, +then waits for actual component completion and clean output EOF. DELETE returns +202 only for clean completion; timeout/failure returns 503. Timeout or canceled +DELETE waiters do not abort accepted work or reopen admission. Concurrent DELETE +waiters join the same drain; after completion removes the connection, later +DELETE returns 404. Without this option, DELETE remains abortive. + The bounded path integrates directly with the core `BoundedChannel` and producer admission. Finite defaults constrain serialized bytes, frame counts, pending work, POSTs, sessions, and streams. Exhaustion fails explicitly rather diff --git a/src/agent-client-protocol-http/src/bounded_server.rs b/src/agent-client-protocol-http/src/bounded_server.rs index 8c006b85..0ddb1587 100644 --- a/src/agent-client-protocol-http/src/bounded_server.rs +++ b/src/agent-client-protocol-http/src/bounded_server.rs @@ -3,6 +3,7 @@ use std::{ collections::{HashMap, VecDeque}, convert::Infallible, sync::{Arc, Mutex, Weak}, + time::Duration, }; use agent_client_protocol::{ @@ -17,7 +18,7 @@ use axum::{ response::{IntoResponse, Response}, routing::{delete, get, post}, }; -use futures::{StreamExt, future::BoxFuture}; +use futures::{FutureExt, StreamExt, future::BoxFuture}; use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore, oneshot}; use super::ServerOptions; @@ -172,6 +173,7 @@ impl BoundedAcpHttpServer { posts: Arc::new(Semaphore::new(limits.max_in_flight_posts)), bodies: Arc::new(Semaphore::new(limits.max_body_bytes)), limits, + graceful_delete: None, }), options: ServerOptions::default(), }) @@ -183,6 +185,24 @@ impl BoundedAcpHttpServer { self } + /// Opt into graceful DELETE with a bounded wait per HTTP request. + /// + /// DELETE seals inbound admission without allocating a body or core frame. It + /// returns 202 only after the component future succeeds and its output reaches + /// clean EOF. This is not acknowledgment that an HTTP peer consumed the output. + /// A deadline or terminal failure returns 503. A deadline never reopens admission + /// or aborts accepted work; canceling the DELETE request only drops its waiter. + /// Concurrent/repeated DELETE requests join the connection-owned drain while it + /// exists. Once completion removes the connection, subsequent requests return 404. + /// POST requests racing with the seal are rejected with 410 before enqueue. + /// Without this option, DELETE retains its legacy abortive behavior. + #[must_use] + pub fn with_graceful_delete(mut self, timeout: Duration) -> Self { + // The server has not exposed the registry before into_router consumes it. + Arc::get_mut(&mut self.state).unwrap().graceful_delete = Some(timeout); + self + } + pub fn into_router(self) -> Router { let mut router = Router::new() .route(&self.options.path, post(handle_post)) @@ -200,6 +220,7 @@ impl BoundedAcpHttpServer { } struct Registry { + graceful_delete: Option, factory: Arc, limits: ServerLimits, connections: Mutex>>, @@ -232,6 +253,8 @@ struct Connection { struct ConnectionState { tx: Option, closed: bool, + draining: bool, + drain_result: Option, task: Option, pending: VecDeque<(RequestId, ResponseRoute)>, streams: HashMap, Mailbox>, @@ -252,25 +275,97 @@ struct Envelope { impl Connection { fn close(&self) { - { + self.close_if_open(false); + } + + // A POST that raced with DELETE must not turn a harmless admission rejection + // into terminal core failure. All destructors and wakeups run after unlocking. + fn close_if_open(&self, only_open: bool) -> bool { + let retired = { let mut state = self.inner.lock().unwrap(); - state.closed = true; - if let Some(tx) = state.tx.take() { - tx.fail("HTTP connection terminated"); - } - state.pending.clear(); - state.streams.clear(); - if let Some(task) = state.task.take() { - task.abort(); + if only_open && (state.draining || state.closed) { + return false; } + state.closed = true; + state.drain_result.get_or_insert(false); + ( + state.tx.take(), + state.task.take(), + std::mem::take(&mut state.pending), + std::mem::take(&mut state.streams), + ) + }; + if let Some(tx) = &retired.0 { + tx.fail("HTTP connection terminated"); } + if let Some(task) = &retired.1 { + task.abort(); + } + drop(retired); self.wake.notify_waiters(); + true + } + + fn remove(&self) { + if let Some(registry) = self.registry.upgrade() { + let removed = registry.connections.lock().unwrap().remove(&self.id); + drop(removed); + } } fn terminate(&self) { self.close(); - if let Some(registry) = self.registry.upgrade() { - registry.connections.lock().unwrap().remove(&self.id); + self.remove(); + } + + fn terminate_if_open(&self) { + if self.close_if_open(true) { + self.remove(); + } + } + + fn begin_drain(&self) { + let tx = { + let mut state = self.inner.lock().unwrap(); + if state.closed || state.draining { + return; + } + // This is the HTTP acceptance linearization point, shared with enqueue. + state.draining = true; + state.tx.take() + }; + // No await separates seal publication and the core close. Cancellation cannot + // strand this operation; close preserves all already accepted frame charges. + if let Some(tx) = tx { + tx.close_channel(); + } + } + + fn finish_drain(&self) -> bool { + let finished = { + let mut state = self.inner.lock().unwrap(); + if state.closed || !state.draining { + return false; + } + state.closed = true; + state.drain_result = Some(true); + state.task.take() + }; + drop(finished); + self.wake.notify_waiters(); + self.remove(); + true + } + + async fn drain_result(&self) -> bool { + loop { + let notified = self.wake.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + if let Some(result) = self.inner.lock().unwrap().drain_result { + return result; + } + notified.await; } } @@ -486,6 +581,21 @@ async fn post_inner( } let connection_id = header_value(request.headers(), HEADER_CONNECTION_ID)?; let session_id = header_value(request.headers(), HEADER_SESSION_ID)?; + // Reject closing connections before reserving global POST/body capacity. A + // second check at enqueue handles bodies already being read when DELETE seals. + let existing = if let Some(id) = &connection_id { + let connection = registry.connections.lock().unwrap().get(id).cloned(); + if let Some(connection) = &connection { + let state = connection.inner.lock().unwrap(); + if state.draining || state.closed { + return Err(StatusCode::GONE); + } + } + connection + } else { + None + }; + // Body shape is unknown until read: rather than let slow request bodies starve // callback responses indefinitely, saturation terminates the addressed connection. // The rejected POST itself has not been accepted and is never resubmitted here. @@ -493,7 +603,7 @@ async fn post_inner( if let Some(id) = &connection_id { let connection = registry.connections.lock().unwrap().get(id).cloned(); if let Some(connection) = connection { - connection.terminate(); + connection.terminate_if_open(); } } StatusCode::TOO_MANY_REQUESTS @@ -576,18 +686,15 @@ async fn post_inner( let encoded = encode_frame(&frame, registry.limits.max_frame_bytes)?; drop(frame); let mut initialization = None; - let connection = if let Some(connection_id) = connection_id { - registry - .connections - .lock() - .unwrap() - .get(&connection_id) - .cloned() - .ok_or(StatusCode::NOT_FOUND)? + let connection = if connection_id.is_some() { + existing.ok_or(StatusCode::NOT_FOUND)? } else { let (BoundedChannel { tx, mut rx }, agent) = (registry.factory)(registry.limits.channel_limits) .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; + let terminal_owner = tx.clone(); + let failure = tx.failure(); + let graceful = registry.graceful_delete.is_some(); let mut streams = HashMap::new(); streams.insert(None, Mailbox::default()); let connection = Arc::new(Connection { @@ -597,6 +704,8 @@ async fn post_inner( inner: Mutex::new(ConnectionState { tx: Some(tx), closed: false, + draining: false, + drain_result: None, task: None, pending: VecDeque::new(), streams, @@ -622,51 +731,80 @@ async fn post_inner( let task_connection = connection.clone(); // Capture cleanup before spawning: dropping even a never-polled task cleans up. let task = tokio::spawn(async move { - let _cleanup = cleanup; + let mut cleanup = cleanup; + // Keep the terminal signal owner alive: dropping every channel endpoint + // cancels its signal, which is not evidence of a core failure. + let _terminal_owner = terminal_owner; + let completion_connection = task_connection.clone(); let router = async move { let mut init_tx = Some(init_tx); while let Some(charged) = rx.next().await { if charged.as_bytes().len() > task_connection.limits.max_frame_bytes { - break; + return Err(()); } let frame = charged.decode(); if !check_batch(&frame, task_connection.limits.max_batch_entries) { - break; + return Err(()); } let Ok(envelope) = task_connection.envelope(charged) else { - break; + return Err(()); }; if let Some(failed) = initialize_response_failed(&frame, &init_id) && let Some(sender) = init_tx.take() { task_connection.complete_initial_routes(&frame); if sender.send((envelope, failed)).is_err() { - break; + return Err(()); } continue; } if task_connection.route(envelope, &frame).is_err() { - break; + return Err(()); } } + Ok::<(), ()>(()) }; - // A bare channel's driver is immediately successful. Success must not - // discard frames still owned by the channel or an escaped producer. - let agent = async move { - if agent.await.is_ok() { - futures::future::pending::<()>().await; + if graceful { + // EOF alone is insufficient: core failure also produces EOF. Poll + // failure first and again after both real futures have completed. + let done = async { + futures::try_join!(router, async { agent.await.map_err(|_| ()) }).map(|_| ()) + }; + tokio::pin!(failure); + let result = tokio::select! { + biased; + _ = &mut failure => Err(()), + result = done => result, + }; + if result.is_ok() + && failure.now_or_never().is_none() + && completion_connection.finish_drain() + { + cleanup.armed = false; } - }; - futures::pin_mut!(agent, router); - let _ = futures::future::select(router, agent).await; + } else { + // A bare channel's driver is immediately successful. Legacy mode + // keeps routing until EOF instead of dropping escaped producers. + let agent = async move { + if agent.await.is_ok() { + futures::future::pending::<()>().await; + } + }; + futures::pin_mut!(agent, router); + let _ = futures::future::select(router, agent).await; + } }); - { + let abort = { let mut state = connection.inner.lock().unwrap(); if state.closed { - task.abort(); + true } else { state.task = Some(task.abort_handle()); + false } + }; + if abort { + task.abort(); } connection }; @@ -678,7 +816,7 @@ async fn post_inner( // One lock makes route reservation, duplicate-ID ordering and enqueue transactional. // There is no await/cancellation point between bookkeeping and core acceptance. let mut state = connection.inner.lock().unwrap(); - if state.closed { + if state.closed || state.draining { return Err(StatusCode::GONE); } if state @@ -868,14 +1006,26 @@ async fn handle_delete(State(registry): State>, request: Request StatusCode::ACCEPTED, + Ok(false) | Err(_) => StatusCode::SERVICE_UNAVAILABLE, + } + .into_response(); + } + connection.terminate(); StatusCode::ACCEPTED.into_response() } +#[cfg(test)] +#[path = "graceful_delete_tests.rs"] +mod graceful_delete_tests; + #[cfg(test)] mod tests { use super::*; diff --git a/src/agent-client-protocol-http/src/graceful_delete_tests.rs b/src/agent-client-protocol-http/src/graceful_delete_tests.rs new file mode 100644 index 00000000..1f00eeb5 --- /dev/null +++ b/src/agent-client-protocol-http/src/graceful_delete_tests.rs @@ -0,0 +1,439 @@ +//! End-to-end graceful deletion tests through the real bounded ConnectTo boundary. +use super::*; +use agent_client_protocol::{Error, TransportChannel}; +use futures::Future; +use std::time::Duration; +use tokio::time::timeout; + +const WAIT: Duration = Duration::from_secs(5); +const NOTICE: &str = r#"{"jsonrpc":"2.0","method":"notice"}"#; + +#[derive(Clone, Copy)] +enum Boundary { + Normal, + EarlyOutboundEof, + CoreFailure, +} + +struct BlockedAdapter { + boundary: Boundary, + read: oneshot::Receiver<()>, + drained: oneshot::Sender>>, + finish: oneshot::Receiver, +} + +impl ConnectTo for BlockedAdapter { + async fn connect_to( + self, + client: impl ConnectTo, + ) -> agent_client_protocol::Result<()> { + let (transport, driver) = client.into_transport_and_future(); + let TransportChannel::Bounded(mut channel) = transport else { + panic!("graceful server must preserve bounded transport"); + }; + driver.await?; + let init = channel.rx.next().await.expect("initialize frame"); + assert!(matches!(init.decode(), TransportFrame::Single(_))); + drop(init); + channel + .tx + .try_send_serialized(r#"{"jsonrpc":"2.0","id":1,"result":{}}"#)?; + self.read.await.expect("release adapter reader"); + let mut received = Vec::new(); + while let Some(frame) = channel.rx.next().await { + received.push(frame.as_bytes().to_vec()); + } + match self.boundary { + Boundary::Normal => {} + Boundary::EarlyOutboundEof => channel.tx.close_channel(), + Boundary::CoreFailure => { + // Fail the actual shared core terminal, but return adapter success: + // clean EOF and Ok(()) must not hide this independent failure. + channel.tx.fail("intentional shared core terminal failure"); + self.drained.send(received).unwrap(); + return Ok(()); + } + } + self.drained.send(received).unwrap(); + // EOF alone is not adapter completion. Keep the outbound side alive until + // the test explicitly permits the real component future to complete. + let succeed = self.finish.await.expect("release adapter completion"); + drop(channel); + if succeed { + Ok(()) + } else { + Err(Error::internal_error().data("intentional adapter failure")) + } + } +} + +struct Harness { + state: Arc, + id: String, + read: oneshot::Sender<()>, + drained: oneshot::Receiver>>, + finish: oneshot::Sender, +} + +async fn checked(future: F) -> F::Output { + timeout(WAIT, future).await.expect("test made no progress") +} + +fn post(body: impl Into, id: Option<&str>) -> Request { + let mut request = Request::builder() + .method("POST") + .header(header::CONTENT_TYPE, JSON_MIME_TYPE); + if let Some(id) = id { + request = request.header(HEADER_CONNECTION_ID, id); + } + request.body(body.into()).unwrap() +} + +fn delete_request(id: &str) -> Request { + Request::builder() + .method("DELETE") + .header(HEADER_CONNECTION_ID, id) + .body(Body::empty()) + .unwrap() +} + +async fn setup(limits: ServerLimits, grace: Duration) -> Harness { + setup_many(limits, grace, &[Boundary::Normal]) + .await + .pop() + .unwrap() +} + +async fn setup_many( + limits: ServerLimits, + grace: Duration, + boundaries: &[Boundary], +) -> Vec { + let mut controls = Vec::new(); + let mut adapters = VecDeque::new(); + for &boundary in boundaries { + let (read, read_rx) = oneshot::channel(); + let (drained_tx, drained) = oneshot::channel(); + let (finish, finish_rx) = oneshot::channel(); + adapters.push_back(BlockedAdapter { + boundary, + read: read_rx, + drained: drained_tx, + finish: finish_rx, + }); + controls.push((read, drained, finish)); + } + let adapters = Mutex::new(adapters); + let state = BoundedAcpHttpServer::new( + move || adapters.lock().unwrap().pop_front().unwrap(), + limits, + ) + .unwrap() + .with_graceful_delete(grace) + .state; + let mut harnesses = Vec::new(); + for (read, drained, finish) in controls { + let id = initialize_connection(state.clone()).await; + harnesses.push(Harness { + state: state.clone(), + id, + read, + drained, + finish, + }); + } + harnesses +} + +async fn initialize_connection(state: Arc) -> String { + let response = checked(handle_post( + State(state.clone()), + post( + r#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}"#, + None, + ), + )) + .await; + assert_eq!(response.status(), StatusCode::OK); + let id = response.headers()[HEADER_CONNECTION_ID] + .to_str() + .unwrap() + .to_owned(); + // Consuming the initialization body disarms the server's cleanup guard. + checked(axum::body::to_bytes(response.into_body(), 4096)) + .await + .unwrap(); + id +} + +async fn enqueue(h: &Harness, body: &str) { + assert_eq!( + checked(handle_post( + State(h.state.clone()), + post(body.to_owned(), Some(&h.id)) + )) + .await + .status(), + StatusCode::ACCEPTED, + ); +} + +async fn assert_closing_before_body_poll(state: Arc, id: &str) { + let body = Body::from_stream(futures::stream::poll_fn(|_| { + panic!("a POST to a closing connection must not poll its body"); + #[allow(unreachable_code)] + std::task::Poll::Ready(None::>) + })); + assert_eq!( + checked(handle_post(State(state), post(body, Some(id)))) + .await + .status(), + StatusCode::GONE, + ); +} + +async fn full_core_drains(limits: ChannelLimits, frames: Vec) { + let h = setup( + ServerLimits { + channel_limits: limits, + ..ServerLimits::default() + }, + WAIT, + ) + .await; + for frame in &frames { + enqueue(&h, frame).await; + } + let mut delete = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut delete).is_pending()); + assert_closing_before_body_poll(h.state.clone(), &h.id).await; + h.read.send(()).unwrap(); + let received = checked(h.drained).await.unwrap(); + assert_eq!( + received, + frames + .iter() + .map(|s| s.as_bytes().to_vec()) + .collect::>() + ); + assert!( + futures::poll!(&mut delete).is_pending(), + "EOF must not substitute for adapter completion" + ); + h.finish.send(true).unwrap(); + assert_eq!(checked(delete).await.status(), StatusCode::ACCEPTED); +} + +#[tokio::test] +async fn graceful_delete_drains_a_full_core_frame_queue() { + full_core_drains( + ChannelLimits { + max_buffered_frames: 2, + ..ChannelLimits::default() + }, + vec![NOTICE.to_owned(), NOTICE.to_owned()], + ) + .await; +} + +#[tokio::test] +async fn graceful_delete_drains_a_full_core_byte_budget() { + let prefix = r#"{"jsonrpc":"2.0","method":"notice","params":""#; + let frame = format!("{prefix}{}\"}}", "x".repeat(256 - prefix.len() - 2)); + assert_eq!(frame.len(), 256); + full_core_drains( + ChannelLimits { + max_frame_bytes: 256, + max_buffered_bytes: 256, + max_buffered_frames: 8, + ..ChannelLimits::default() + }, + vec![frame], + ) + .await; +} + +#[tokio::test] +async fn graceful_delete_bypasses_eight_mib_of_slow_body_reservations() { + let mut connections = setup_many( + ServerLimits { + max_frame_bytes: 1024 * 1024, + max_body_bytes: 8 * 1024 * 1024, + max_in_flight_posts: 8, + ..ServerLimits::default() + }, + WAIT, + &[Boundary::Normal; 9], + ) + .await; + let h = connections.remove(0); + enqueue(&h, NOTICE).await; + let mut slow_posts = Vec::new(); + for other in &connections { + let (polled, first_poll) = oneshot::channel(); + let (release, released) = oneshot::channel(); + let body = Body::from_stream(futures::stream::once(async move { + polled.send(()).unwrap(); + released.await.unwrap(); + let mut bytes = NOTICE.as_bytes().to_vec(); + bytes.resize(1024 * 1024, b' '); + Ok::<_, Infallible>(axum::body::Bytes::from(bytes)) + })); + let task = tokio::spawn(handle_post( + State(h.state.clone()), + post(body, Some(&other.id)), + )); + checked(first_poll).await.unwrap(); + slow_posts.push((release, task)); + } + assert_eq!(h.state.bodies.available_permits(), 0); + assert_eq!(h.state.posts.available_permits(), 0); + let mut delete = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut delete).is_pending()); + assert_closing_before_body_poll(h.state.clone(), &h.id).await; + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + assert!(futures::poll!(&mut delete).is_pending()); + h.finish.send(true).unwrap(); + assert_eq!(checked(delete).await.status(), StatusCode::ACCEPTED); + for (release, task) in slow_posts { + release.send(()).unwrap(); + assert_eq!(checked(task).await.unwrap().status(), StatusCode::ACCEPTED); + } + assert_eq!(h.state.bodies.available_permits(), 8 * 1024 * 1024); + for other in connections { + let mut delete = Box::pin(handle_delete(State(other.state), delete_request(&other.id))); + assert!(futures::poll!(&mut delete).is_pending()); + other.read.send(()).unwrap(); + assert_eq!( + checked(other.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + other.finish.send(true).unwrap(); + assert_eq!(checked(delete).await.status(), StatusCode::ACCEPTED); + } +} + +#[tokio::test] +async fn cancelled_delete_waiter_does_not_cancel_drain_and_repeat_joins() { + let h = setup(ServerLimits::default(), WAIT).await; + enqueue(&h, NOTICE).await; + let mut first = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut first).is_pending()); + drop(first); + assert_closing_before_body_poll(h.state.clone(), &h.id).await; + let mut repeated = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut repeated).is_pending()); + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + assert!(futures::poll!(&mut repeated).is_pending()); + h.finish.send(true).unwrap(); + assert_eq!(checked(repeated).await.status(), StatusCode::ACCEPTED); +} + +#[tokio::test] +async fn graceful_delete_timeout_keeps_connection_closing() { + let h = setup(ServerLimits::default(), Duration::ZERO).await; + enqueue(&h, NOTICE).await; + assert_eq!( + checked(handle_delete(State(h.state.clone()), delete_request(&h.id))) + .await + .status(), + StatusCode::SERVICE_UNAVAILABLE + ); + assert_closing_before_body_poll(h.state.clone(), &h.id).await; + assert_eq!( + checked(handle_delete(State(h.state.clone()), delete_request(&h.id))) + .await + .status(), + StatusCode::SERVICE_UNAVAILABLE + ); + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + h.finish.send(true).unwrap(); +} + +#[tokio::test] +async fn genuine_adapter_failure_never_reports_graceful_success() { + let h = setup(ServerLimits::default(), WAIT).await; + enqueue(&h, NOTICE).await; + let mut delete = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut delete).is_pending()); + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + h.finish.send(false).unwrap(); + assert_ne!(checked(delete).await.status(), StatusCode::ACCEPTED); + assert_ne!( + checked(handle_delete(State(h.state), delete_request(&h.id))) + .await + .status(), + StatusCode::ACCEPTED + ); +} + +#[tokio::test] +async fn shared_core_terminal_failure_is_not_clean_eof() { + let h = setup_many(ServerLimits::default(), WAIT, &[Boundary::CoreFailure]) + .await + .pop() + .unwrap(); + enqueue(&h, NOTICE).await; + let mut delete = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut delete).is_pending()); + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + assert_eq!( + checked(delete).await.status(), + StatusCode::SERVICE_UNAVAILABLE + ); +} + +#[tokio::test] +async fn outbound_eof_does_not_substitute_for_adapter_completion() { + let h = setup_many( + ServerLimits::default(), + Duration::ZERO, + &[Boundary::EarlyOutboundEof], + ) + .await + .pop() + .unwrap(); + enqueue(&h, NOTICE).await; + assert_eq!( + checked(handle_delete(State(h.state.clone()), delete_request(&h.id))) + .await + .status(), + StatusCode::SERVICE_UNAVAILABLE + ); + h.read.send(()).unwrap(); + // This acknowledgment follows closing the actual bounded outbound direction. + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + // Give the woken router a turn to observe EOF; the adapter remains gated. + tokio::task::yield_now().await; + assert_eq!( + checked(handle_delete(State(h.state.clone()), delete_request(&h.id))) + .await + .status(), + StatusCode::SERVICE_UNAVAILABLE + ); + assert_closing_before_body_poll(h.state.clone(), &h.id).await; + h.finish.send(true).unwrap(); +} From 8cd41abf948a1d281eb6e91fdc04864cbf21f12d Mon Sep 17 00:00:00 2001 From: daniel Date: Tue, 29 Sep 2026 02:58:12 +0100 Subject: [PATCH 9/9] test(http): cover seal races and retained graceful output --- .../src/graceful_delete_tests.rs | 71 ++++++++++++++++++- 1 file changed, 70 insertions(+), 1 deletion(-) diff --git a/src/agent-client-protocol-http/src/graceful_delete_tests.rs b/src/agent-client-protocol-http/src/graceful_delete_tests.rs index 1f00eeb5..66bf7d5c 100644 --- a/src/agent-client-protocol-http/src/graceful_delete_tests.rs +++ b/src/agent-client-protocol-http/src/graceful_delete_tests.rs @@ -6,12 +6,13 @@ use std::time::Duration; use tokio::time::timeout; const WAIT: Duration = Duration::from_secs(5); -const NOTICE: &str = r#"{"jsonrpc":"2.0","method":"notice"}"#; +const NOTICE: &str = r#"{"jsonrpc":"2.0","method":"notice","params":null}"#; #[derive(Clone, Copy)] enum Boundary { Normal, EarlyOutboundEof, + EmitOutput, CoreFailure, } @@ -46,6 +47,7 @@ impl ConnectTo for BlockedAdapter { match self.boundary { Boundary::Normal => {} Boundary::EarlyOutboundEof => channel.tx.close_channel(), + Boundary::EmitOutput => channel.tx.try_send_serialized(NOTICE)?, Boundary::CoreFailure => { // Fail the actual shared core terminal, but return adapter success: // clean EOF and Ok(()) must not hide this independent failure. @@ -437,3 +439,70 @@ async fn outbound_eof_does_not_substitute_for_adapter_completion() { assert_closing_before_body_poll(h.state.clone(), &h.id).await; h.finish.send(true).unwrap(); } + +#[tokio::test] +async fn already_reading_post_crossing_delete_seal_does_not_enqueue_or_abort_drain() { + let h = setup(ServerLimits::default(), WAIT).await; + enqueue(&h, NOTICE).await; + let (polled, first_poll) = oneshot::channel(); + let (release, released) = oneshot::channel(); + let body = Body::from_stream(futures::stream::once(async move { + polled.send(()).unwrap(); + released.await.unwrap(); + Ok::<_, Infallible>(axum::body::Bytes::from_static( + br#"{"jsonrpc":"2.0","method":"late","params":null}"#, + )) + })); + let post = tokio::spawn(handle_post(State(h.state.clone()), post(body, Some(&h.id)))); + checked(first_poll).await.unwrap(); + let mut delete = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut delete).is_pending()); + // The body was admitted before DELETE, but is not a core-accepted frame. + // Finishing its read after the seal must neither enqueue it nor fail the core. + release.send(()).unwrap(); + assert_eq!(checked(post).await.unwrap().status(), StatusCode::GONE); + assert!(futures::poll!(&mut delete).is_pending()); + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + assert!(futures::poll!(&mut delete).is_pending()); + h.finish.send(true).unwrap(); + assert_eq!(checked(delete).await.status(), StatusCode::ACCEPTED); +} + +#[tokio::test] +async fn existing_sse_body_retains_queued_output_after_clean_delete() { + let h = setup_many(ServerLimits::default(), WAIT, &[Boundary::EmitOutput]) + .await + .pop() + .unwrap(); + let response = checked(handle_get( + State(h.state.clone()), + Request::builder() + .method("GET") + .header(HEADER_CONNECTION_ID, &h.id) + .header(header::ACCEPT, "text/event-stream") + .body(Body::empty()) + .unwrap(), + )) + .await; + assert_eq!(response.status(), StatusCode::OK); + // Subscribe through the real handler, but do not consume any SSE bytes yet. + let body = response.into_body(); + enqueue(&h, NOTICE).await; + let mut delete = Box::pin(handle_delete(State(h.state.clone()), delete_request(&h.id))); + assert!(futures::poll!(&mut delete).is_pending()); + h.read.send(()).unwrap(); + assert_eq!( + checked(h.drained).await.unwrap(), + vec![NOTICE.as_bytes().to_vec()] + ); + h.finish.send(true).unwrap(); + assert_eq!(checked(delete).await.status(), StatusCode::ACCEPTED); + // Completion must preserve already-routed output for existing subscribers, + // then terminate their streams after the final queued event. + let bytes = checked(axum::body::to_bytes(body, 4096)).await.unwrap(); + assert_eq!(bytes.as_ref(), format!("data: {NOTICE}\n\n").as_bytes()); +}