diff --git a/hive-runtime/src/acp/mod.rs b/hive-runtime/src/acp/mod.rs index fbe29b88..49cce1f8 100644 --- a/hive-runtime/src/acp/mod.rs +++ b/hive-runtime/src/acp/mod.rs @@ -100,7 +100,8 @@ pub struct AcpRuntime { struct Live { conn: Connection, load_session: bool, - session: Option, + /// The recorded session, or a new one whose first prompt is unanswered. + loaded: Option, model: Option, } @@ -117,6 +118,14 @@ impl AcpRuntime { } } + /// Names a new session until its first prompt has been answered, so a + /// failed first prompt is retried on the same session. + fn pending_file(&self) -> PathBuf { + let mut path = self.session_file.clone().into_os_string(); + path.push(".pending"); + PathBuf::from(path) + } + async fn start(&self, config: &Config) -> Result { let (_, servers) = mcp_servers(config)?; let cwd = session_cwd(config); @@ -142,43 +151,47 @@ impl AcpRuntime { Ok(Live { conn, load_session: caps["loadSession"] == Value::Bool(true), - session: None, + loaded: None, model: None, }) } - /// Attach the durable session: the one already loaded, else the one named - /// by the session file, else a new one. Returns its id and whether it is - /// new. A new session is not recorded here: [`Self::turn`] records it once - /// its first prompt, the one carrying the system prompt, has been answered. + /// Attach the durable session: the recorded one (named by the session + /// file), else a new one whose first prompt was never answered (named by + /// the pending file), else a new one. Returns its id and whether its first + /// prompt, the one carrying the system prompt, is still to be answered; + /// [`Self::turn`] records the session once it has been. async fn attach( &self, live: &mut Live, cwd: &Path, servers: &[Value], ) -> Result<(String, bool)> { - let persisted = std::fs::read_to_string(&self.session_file) - .ok() - .map(|s| s.trim().to_owned()) - .filter(|s| !s.is_empty()); - if let Some(id) = persisted { - if live.session.as_deref() == Some(id.as_str()) { - return Ok((id, false)); + let known = [ + (read_id(&self.session_file), false), + (read_id(&self.pending_file()), true), + ] + .into_iter() + .filter_map(|(id, new)| Some((id?, new))); + for (id, new) in known { + if live.loaded.as_deref() == Some(id.as_str()) { + return Ok((id, new)); } - if live.load_session { - let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers }); - match live.conn.request("session/load", params).await { - Ok(response) => { - discard_stale(&mut live.conn)?; - live.model = stream::session_model(&response); - live.session = Some(id.clone()); - return Ok((id, false)); - } - Err(e @ AcpError::Rpc { .. }) => { - tracing::warn!(error = %e, "ACP session/load failed; starting a new session"); - } - Err(e) => return Err(e.into()), + if !live.load_session { + continue; + } + let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers }); + match live.conn.request("session/load", params).await { + Ok(response) => { + discard_stale(&mut live.conn)?; + live.model = stream::session_model(&response); + live.loaded = Some(id.clone()); + return Ok((id, new)); } + Err(e @ AcpError::Rpc { .. }) => { + tracing::warn!(error = %e, session = %id, "ACP session/load failed"); + } + Err(e) => return Err(e.into()), } } let params = json!({ "cwd": cwd, "mcpServers": servers }); @@ -191,15 +204,14 @@ impl AcpRuntime { })? .to_owned(); live.model = stream::session_model(&response); + live.loaded = Some(id.clone()); + write_id(&self.pending_file(), &id)?; Ok((id, true)) } fn record_session(&self, id: &str) -> std::result::Result<(), AcpError> { - if let Some(dir) = self.session_file.parent() { - std::fs::create_dir_all(dir)?; - } - std::fs::write(&self.session_file, id)?; - Ok(()) + write_id(&self.session_file, id)?; + remove_if_present(&self.pending_file()) } async fn turn( @@ -245,7 +257,6 @@ impl AcpRuntime { if created { self.record_session(&session)?; } - live.session = Some(session.clone()); // Updates may still arrive after the reply: an agent can forward its // stream asynchronously. Take them until the agent goes quiet. let settled = tokio::time::Instant::now() + SETTLE_MAX; @@ -308,8 +319,10 @@ impl Runtime for AcpRuntime { } /// Moves the session file aside; the agent keeps its own copy of the - /// session, so nothing is lost. + /// session, so nothing is lost. A new session whose first prompt was never + /// answered is dropped too. fn archive(&self) -> Result> { + remove_if_present(&self.pending_file()).map_err(Error::Acp)?; if !self.session_file.exists() { return Ok(None); } @@ -370,6 +383,28 @@ fn discard_stale(conn: &mut Connection) -> std::result::Result<(), AcpError> { Ok(()) } +fn read_id(file: &Path) -> Option { + std::fs::read_to_string(file) + .ok() + .map(|s| s.trim().to_owned()) + .filter(|s| !s.is_empty()) +} + +fn write_id(file: &Path, id: &str) -> std::result::Result<(), AcpError> { + if let Some(dir) = file.parent() { + std::fs::create_dir_all(dir)?; + } + std::fs::write(file, id)?; + Ok(()) +} + +fn remove_if_present(file: &Path) -> std::result::Result<(), AcpError> { + match std::fs::remove_file(file) { + Err(e) if e.kind() != std::io::ErrorKind::NotFound => Err(e.into()), + _ => Ok(()), + } +} + fn session_cwd(config: &Config) -> PathBuf { config .cwd @@ -402,34 +437,45 @@ mod tests { use super::{AcpError, AcpRuntime}; use crate::{AcpCommand, Error, Runtime}; - /// An ACP agent in plain `sh`: answers `initialize`, numbers its sessions - /// `s1`, `s2`, …, appends every `session/prompt` line to the file named by - /// its argument, and fails the first prompt it is sent. + /// An ACP agent in plain `sh`. It answers `initialize`, numbers its + /// sessions `s1`, `s2`, …, appends every method it is sent to + /// `.methods` (and `start` when it starts), and every + /// `session/prompt` line to ``. It answers every prompt with + /// `end_turn`, except the very first one it is sent across restarts, + /// which it handles per its mode: + /// + /// - `fail`: an error. const AGENT: &str = r#" n=0 +printf 'start\n' >> "$1.methods" while IFS= read -r line; do id=${line#*\"id\":}; id=${id%%[,\}]*} - case $line in - *'"method":"initialize"'*) + m=${line#*\"method\":\"}; m=${m%%\"*} + printf '%s\n' "$m" >> "$1.methods" + case $m in + initialize) printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":1,"agentCapabilities":{"loadSession":true,"mcpCapabilities":{"http":true}}}}\n' "$id" ;; - *'"method":"session/new"'*) + session/new) n=$((n+1)) printf '{"jsonrpc":"2.0","id":%s,"result":{"sessionId":"s%s"}}\n' "$id" "$n" ;; - *'"method":"session/load"'*) + session/load) printf '{"jsonrpc":"2.0","id":%s,"result":{}}\n' "$id" ;; - *'"method":"session/prompt"'*) + session/prompt) printf '%s\n' "$line" >> "$1" - if [ -e "$1.failed" ]; then + if [ -e "$1.once" ]; then printf '{"jsonrpc":"2.0","id":%s,"result":{"stopReason":"end_turn"}}\n' "$id" else - : > "$1.failed" - printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32603,"message":"provider down"}}\n' "$id" + : > "$1.once" + case $2 in + fail) + printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32603,"message":"provider down"}}\n' "$id" ;; + esac fi ;; esac done "#; - fn runtime(dir: &Path) -> AcpRuntime { + fn runtime(dir: &Path, mode: &str) -> AcpRuntime { let script = dir.join("agent.sh"); std::fs::write(&script, AGENT).unwrap(); let command = AcpCommand { @@ -437,6 +483,7 @@ done args: vec![ script.display().to_string(), dir.join("prompts").display().to_string(), + mode.to_owned(), ], env: BTreeMap::new(), }; @@ -461,10 +508,23 @@ done .collect() } + /// How often the agent started, or was sent `method`. + fn count(dir: &Path, method: &str) -> usize { + std::fs::read_to_string(dir.join("prompts.methods")) + .unwrap() + .lines() + .filter(|line| *line == method) + .count() + } + + fn recorded(dir: &Path) -> String { + std::fs::read_to_string(dir.join("session")).unwrap() + } + #[tokio::test] - async fn a_session_whose_first_prompt_failed_is_not_resumed() { + async fn a_session_whose_first_prompt_failed_is_retried_not_resumed() { let dir = tempfile::tempdir().unwrap(); - let (runtime, config) = (runtime(dir.path()), config(dir.path())); + let (runtime, config) = (runtime(dir.path(), "fail"), config(dir.path())); let failed = runtime.run(&config, "one", &NoopSink).await; assert!( @@ -479,25 +539,28 @@ done assert!(!next.created); assert_eq!(carried_system_prompt(dir.path()), [true, true, false]); - assert_eq!( - std::fs::read_to_string(dir.path().join("session")).unwrap(), - "s2" - ); + assert_eq!(recorded(dir.path()), "s1"); + assert_eq!(count(dir.path(), "session/new"), 1); + assert!(!dir.path().join("session.pending").exists()); } #[tokio::test] - async fn after_a_restart_a_session_whose_first_prompt_failed_is_not_loaded() { + async fn after_a_restart_a_session_whose_first_prompt_failed_is_retried() { let dir = tempfile::tempdir().unwrap(); let config = config(dir.path()); - let before = runtime(dir.path()); + let before = runtime(dir.path(), "fail"); assert!(before.run(&config, "one", &NoopSink).await.is_err()); drop(before); - let after = runtime(dir.path()); + let after = runtime(dir.path(), "fail"); let retried = after.run(&config, "two", &NoopSink).await.unwrap(); assert!(retried.created); assert_eq!(carried_system_prompt(dir.path()), [true, true]); + assert_eq!(count(dir.path(), "start"), 2); + assert_eq!(count(dir.path(), "session/new"), 1); + assert_eq!(count(dir.path(), "session/load"), 1); + assert_eq!(recorded(dir.path()), "s1"); } }