From 764a6e0b086ac1ae0606b5758c26c639678578d2 Mon Sep 17 00:00:00 2001 From: meh Date: Mon, 31 Aug 2026 09:05:37 +0700 Subject: [PATCH] perf: signal proxy drain completion --- Cargo.lock | 6 +- Cargo.toml | 6 +- crates/proxy/src/lib.rs | 4 +- crates/proxy/src/runtime.rs | 145 +++++++++++++++++++++++++---- crates/proxy/src/web.rs | 4 +- crates/proxy/tests/local_socket.rs | 24 ++--- 6 files changed, 147 insertions(+), 42 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index cb6814b..96069b6 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4,7 +4,7 @@ version = 4 [[package]] name = "agency-proxy" -version = "0.1.8" +version = "0.1.9" dependencies = [ "agency-proxy-client", "agency-proxy-protocol", @@ -25,7 +25,7 @@ dependencies = [ [[package]] name = "agency-proxy-client" -version = "0.1.8" +version = "0.1.9" dependencies = [ "agency-proxy-protocol", "endpoint-libs", @@ -37,7 +37,7 @@ dependencies = [ [[package]] name = "agency-proxy-protocol" -version = "0.1.8" +version = "0.1.9" dependencies = [ "serde", "serde_json", diff --git a/Cargo.toml b/Cargo.toml index a8513d4..3709412 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,7 +3,7 @@ resolver = "3" members = ["crates/client", "crates/protocol", "crates/proxy"] [workspace.package] -version = "0.1.8" +version = "0.1.9" edition = "2024" license = "Apache-2.0" publish = true @@ -11,8 +11,8 @@ homepage = "https://github.com/pathscale/agencyproxy" repository = "https://github.com/pathscale/agencyproxy" [workspace.dependencies] -agency-proxy-protocol = { version = "^0.1.8", path = "crates/protocol" } -agency-proxy-client = { version = "^0.1.8", path = "crates/client" } +agency-proxy-protocol = { version = "^0.1.9", path = "crates/protocol" } +agency-proxy-client = { version = "^0.1.9", path = "crates/client" } async-trait = "0.1" clap = { version = "4", features = ["derive"] } endpoint-libs = { version = "^3", default-features = false, features = ["framed-transport"] } diff --git a/crates/proxy/src/lib.rs b/crates/proxy/src/lib.rs index a8882f3..9c14d97 100644 --- a/crates/proxy/src/lib.rs +++ b/crates/proxy/src/lib.rs @@ -257,9 +257,7 @@ async fn handle_connection( if mode == ShutdownMode::Terminate { registry.cancel_all().await; } - while registry.active_count().await > 0 { - tokio::time::sleep(std::time::Duration::from_millis(25)).await; - } + registry.wait_until_idle().await; let _ = request_shutdown.send(true); }); return Ok(()); diff --git a/crates/proxy/src/runtime.rs b/crates/proxy/src/runtime.rs index 216b517..1ae4b96 100644 --- a/crates/proxy/src/runtime.rs +++ b/crates/proxy/src/runtime.rs @@ -8,7 +8,7 @@ use std::{ sync::Arc, }; use thiserror::Error; -use tokio::sync::{RwLock, broadcast, oneshot}; +use tokio::sync::{RwLock, broadcast, oneshot, watch}; const MAX_REPLAY_EVENTS: usize = 2_048; @@ -27,18 +27,59 @@ struct LiveRun { events: broadcast::Sender, control: agent_abstraction::RunControl, cancel: Option>, + completed: watch::Sender, } #[derive(Clone, Debug)] pub struct RuntimeRegistry { runs: Arc>>, + activity: RunActivity, executor: tokio::runtime::Handle, } +#[derive(Clone, Debug)] +struct RunActivity { + active: watch::Sender, +} + +impl Default for RunActivity { + fn default() -> Self { + let (active, _) = watch::channel(0); + Self { active } + } +} + +impl RunActivity { + fn started(&self) { + self.active.send_modify(|active| *active += 1); + } + + fn finished(&self) { + self.active.send_modify(|active| { + debug_assert!(*active > 0, "a run finished without being active"); + *active = active.saturating_sub(1); + }); + } + + fn count(&self) -> usize { + *self.active.borrow() + } + + async fn wait_until_idle(&self) { + let mut active = self.active.subscribe(); + while *active.borrow_and_update() > 0 { + if active.changed().await.is_err() { + break; + } + } + } +} + impl Default for RuntimeRegistry { fn default() -> Self { Self { runs: Arc::default(), + activity: RunActivity::default(), // The registry owns provider tasks across connection transports. // Capturing the daemon runtime here prevents a WebSocket shard or // disconnected client runtime from becoming their accidental owner. @@ -176,6 +217,7 @@ impl RuntimeRegistry { let control = run.control(); let (events, _) = broadcast::channel(256); let (cancel, mut cancelled) = oneshot::channel(); + let (completed, _) = watch::channel(false); let snapshot = RunSnapshot { run_id: run_id.clone(), state: RunState::Starting, @@ -201,8 +243,10 @@ impl RuntimeRegistry { events, control, cancel: Some(cancel), + completed, }, ); + self.activity.started(); } let registry = self.clone(); @@ -261,20 +305,12 @@ impl RuntimeRegistry { } pub async fn active_count(&self) -> usize { - self.runs - .read() - .await - .values() - .filter(|run| { - matches!( - run.snapshot.state, - RunState::Starting - | RunState::Running - | RunState::WaitingApproval - | RunState::Finishing - ) - }) - .count() + self.activity.count() + } + + /// Wait until every run accepted by this registry reaches a terminal state. + pub async fn wait_until_idle(&self) { + self.activity.wait_until_idle().await; } pub async fn attach(&self, run_id: &RunId, after: u64) -> Result { @@ -339,10 +375,17 @@ impl RuntimeRegistry { } pub async fn cancel(&self, run_id: &RunId) -> Result<(), RuntimeError> { - let mut runs = self.runs.write().await; - let run = runs.get_mut(run_id).ok_or(RuntimeError::NotFound)?; - let cancel = run.cancel.take().ok_or(RuntimeError::Conflict)?; + // Subscribe before delivering cancellation. A fast provider can emit + // its terminal event in the same scheduling turn; installing the + // receiver first makes the acknowledgement race-free. + let (cancel, completed) = { + let mut runs = self.runs.write().await; + let run = runs.get_mut(run_id).ok_or(RuntimeError::NotFound)?; + let cancel = run.cancel.take().ok_or(RuntimeError::Conflict)?; + (cancel, run.completed.subscribe()) + }; let _ = cancel.send(()); + wait_until_completed(completed).await; Ok(()) } @@ -503,9 +546,11 @@ impl RuntimeRegistry { return; }; run.snapshot.latest_sequence += 1; + let was_active = is_active_state(&run.snapshot.state); if let Some(state) = state { run.snapshot.state = state; } + let finished = was_active && !is_active_state(&run.snapshot.state); if let Some(session) = session { run.snapshot.provider_session_id = Some(session); } @@ -516,9 +561,28 @@ impl RuntimeRegistry { }; push_journal(&mut run.journal, &mut run.replay_floor, event.clone()); let _ = run.events.send(event); + if finished { + run.completed.send_replace(true); + self.activity.finished(); + } } } +async fn wait_until_completed(mut completed: watch::Receiver) { + while !*completed.borrow_and_update() { + if completed.changed().await.is_err() { + break; + } + } +} + +fn is_active_state(state: &RunState) -> bool { + matches!( + state, + RunState::Starting | RunState::Running | RunState::WaitingApproval | RunState::Finishing + ) +} + fn push_journal( journal: &mut VecDeque, replay_floor: &mut u64, @@ -636,4 +700,49 @@ mod tests { assert_eq!(replay_floor, 3); assert_eq!(journal.front().map(|event| event.sequence), Some(4)); } + + #[tokio::test] + async fn activity_waits_for_the_last_run_without_polling() { + let activity = RunActivity::default(); + activity.started(); + activity.started(); + + let waiter = tokio::spawn({ + let activity = activity.clone(); + async move { activity.wait_until_idle().await } + }); + tokio::task::yield_now().await; + assert!(!waiter.is_finished()); + + activity.finished(); + tokio::task::yield_now().await; + assert!(!waiter.is_finished()); + + activity.finished(); + waiter.await.unwrap(); + assert_eq!(activity.count(), 0); + } + + #[tokio::test] + async fn activity_wait_returns_immediately_when_idle() { + RunActivity::default().wait_until_idle().await; + } + + #[tokio::test] + async fn completion_wait_observes_a_signal_without_polling() { + let (completed, receiver) = watch::channel(false); + let waiter = tokio::spawn(wait_until_completed(receiver)); + tokio::task::yield_now().await; + assert!(!waiter.is_finished()); + + completed.send_replace(true); + waiter.await.unwrap(); + } + + #[tokio::test] + async fn completion_wait_returns_immediately_after_the_signal() { + let (completed, receiver) = watch::channel(false); + completed.send_replace(true); + wait_until_completed(receiver).await; + } } diff --git a/crates/proxy/src/web.rs b/crates/proxy/src/web.rs index b307b5f..daa8d21 100644 --- a/crates/proxy/src/web.rs +++ b/crates/proxy/src/web.rs @@ -125,9 +125,7 @@ pub async fn serve_websocket( // endpoint-libs closes admission before returning from `listen` on TERM. // Keep the process alive until every provider already accepted by this // registry settles, matching the Unix transport's drain restart contract. - while drain_registry.active_count().await > 0 { - tokio::time::sleep(std::time::Duration::from_millis(25)).await; - } + drain_registry.wait_until_idle().await; if active > 0 { eprintln!("AgencyProxy WebSocket runs drained"); } diff --git a/crates/proxy/tests/local_socket.rs b/crates/proxy/tests/local_socket.rs index 0ff9e36..374d3c7 100644 --- a/crates/proxy/tests/local_socket.rs +++ b/crates/proxy/tests/local_socket.rs @@ -143,7 +143,7 @@ printf '%s\n' '{"type":"result","subtype":"success","is_error":false,"result":"d } #[tokio::test] -async fn cancel_stops_an_uncooperative_provider_and_publishes_a_terminal_event() { +async fn cancel_acknowledges_only_after_an_uncooperative_provider_is_terminal() { let dir = tempdir().expect("temp dir should exist"); let binary = dir.path().join("stubborn-claude"); std::fs::write(&binary, "#!/bin/sh\ntrap '' TERM\nsleep 30\n") @@ -211,6 +211,15 @@ async fn cancel_stops_an_uncooperative_provider_and_publishes_a_terminal_event() ServerResponse::Accepted ); + assert!(matches!( + client + .request(ClientMessage::ListRuns) + .await + .expect("list should answer immediately after cancel"), + ServerResponse::Runs { runs } + if runs.iter().any(|run| run.run_id == run_id && run.state == agency_proxy_protocol::RunState::Canceled) + )); + let terminal = tokio::time::timeout(Duration::from_secs(2), async { loop { if let ServerFrame::Event { @@ -227,15 +236,6 @@ async fn cancel_stops_an_uncooperative_provider_and_publishes_a_terminal_event() .await .expect("cancel should publish a terminal event promptly"); assert_eq!(terminal, "the run was canceled"); - - assert!(matches!( - client - .request(ClientMessage::ListRuns) - .await - .expect("list should answer"), - ServerResponse::Runs { runs } - if runs.iter().any(|run| run.run_id == run_id && run.state == agency_proxy_protocol::RunState::Canceled) - )); task.abort(); } @@ -325,7 +325,7 @@ printf '%s\n' '{"type":"result","subtype":"success","is_error":false,"result":"d "draining closes admission before acknowledging" ); std::fs::write(&release, "release").expect("provider release should write"); - tokio::time::timeout(Duration::from_secs(2), task) + tokio::time::timeout(Duration::from_secs(5), task) .await .expect("server should stop after the run drains") .expect("server task should join") @@ -532,7 +532,7 @@ printf '%s\n' '{"type":"result","subtype":"success","is_error":false,"result":"s let mut saw_finished = false; let mut latest = 0; for _ in 0..8 { - let frame = tokio::time::timeout(Duration::from_secs(2), events.recv()) + let frame = tokio::time::timeout(Duration::from_secs(5), events.recv()) .await .expect("proxy should replay promptly") .expect("replay frame should exist");