Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

6 changes: 3 additions & 3 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,16 +3,16 @@ 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
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"] }
Expand Down
4 changes: 1 addition & 3 deletions crates/proxy/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(());
Expand Down
145 changes: 127 additions & 18 deletions crates/proxy/src/runtime.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand All @@ -27,18 +27,59 @@ struct LiveRun {
events: broadcast::Sender<SequencedEvent>,
control: agent_abstraction::RunControl,
cancel: Option<oneshot::Sender<()>>,
completed: watch::Sender<bool>,
}

#[derive(Clone, Debug)]
pub struct RuntimeRegistry {
runs: Arc<RwLock<BTreeMap<RunId, LiveRun>>>,
activity: RunActivity,
executor: tokio::runtime::Handle,
}

#[derive(Clone, Debug)]
struct RunActivity {
active: watch::Sender<usize>,
}

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.
Expand Down Expand Up @@ -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,
Expand All @@ -201,8 +243,10 @@ impl RuntimeRegistry {
events,
control,
cancel: Some(cancel),
completed,
},
);
self.activity.started();
}

let registry = self.clone();
Expand Down Expand Up @@ -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<Attachment, RuntimeError> {
Expand Down Expand Up @@ -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(())
}

Expand Down Expand Up @@ -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);
}
Expand All @@ -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<bool>) {
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<SequencedEvent>,
replay_floor: &mut u64,
Expand Down Expand Up @@ -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;
}
}
4 changes: 1 addition & 3 deletions crates/proxy/src/web.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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");
}
Expand Down
24 changes: 12 additions & 12 deletions crates/proxy/tests/local_socket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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 {
Expand All @@ -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();
}

Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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");
Expand Down
Loading