diff --git a/Cargo.lock b/Cargo.lock index 14729ee6..daf6faa8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1987,6 +1987,7 @@ dependencies = [ "hive-claude", "serde", "serde_json", + "tempfile", "thiserror 2.0.18", "tokio", "tracing", diff --git a/hive-runtime/Cargo.toml b/hive-runtime/Cargo.toml index 7039ef75..56b1701b 100644 --- a/hive-runtime/Cargo.toml +++ b/hive-runtime/Cargo.toml @@ -14,3 +14,6 @@ serde_json.workspace = true thiserror.workspace = true tokio.workspace = true tracing.workspace = true + +[dev-dependencies] +tempfile = "3" diff --git a/hive-runtime/src/acp/mod.rs b/hive-runtime/src/acp/mod.rs index 223f34a1..73b60828 100644 --- a/hive-runtime/src/acp/mod.rs +++ b/hive-runtime/src/acp/mod.rs @@ -28,10 +28,21 @@ const STDERR_TAIL: usize = 20; const SETTLE_QUIET: std::time::Duration = std::time::Duration::from_millis(500); const SETTLE_MAX: std::time::Duration = std::time::Duration::from_secs(5); -/// Decides the agent's `session/request_permission` requests from the tool -/// call's ACP `kind` (`read`, `edit`, `execute`, `fetch`, …; `other` when the -/// agent gives none). `true` allows the call once, `false` rejects it. -pub type PermissionPolicy = Arc bool + Send + Sync>; +/// One `session/request_permission` request, as a [`PermissionPolicy`] sees it. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PermissionAsk<'a> { + /// The tool call's ACP `kind` (`read`, `edit`, `execute`, `fetch`, …; + /// `other` when the agent gives none). + pub kind: &'a str, + /// The MCP server, of those handed to the session, whose tool this is: + /// the tool call's title names it as `_`, + /// `__` or `mcp____`. + pub mcp_server: Option<&'a str>, +} + +/// Decides the agent's `session/request_permission` requests. `true` allows +/// the call once, `false` rejects it. +pub type PermissionPolicy = Arc) -> bool + Send + Sync>; /// Why an ACP operation failed. #[derive(Debug, thiserror::Error)] @@ -105,8 +116,10 @@ impl AcpRuntime { } } - async fn start(&self, cwd: &Path) -> Result { - let conn = Connection::spawn(&self.command, cwd, self.permit.clone())?; + async fn start(&self, config: &Config) -> Result { + let (_, servers) = mcp_servers(config)?; + let cwd = session_cwd(config); + let conn = Connection::spawn(&self.command, &cwd, self.permit.clone(), servers)?; let init = conn .request( "initialize", @@ -134,15 +147,22 @@ impl AcpRuntime { } /// Attach the durable session: the one already loaded, else the one named - /// by the session file, else a new one. Returns whether it is new. - async fn attach(&self, live: &mut Live, cwd: &Path, servers: &[Value]) -> Result { + /// 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. + 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(false); + return Ok((id, false)); } if live.load_session { let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers }); @@ -150,8 +170,8 @@ impl AcpRuntime { Ok(response) => { discard_stale(&mut live.conn)?; live.model = stream::session_model(&response); - live.session = Some(id); - return Ok(false); + 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"); @@ -169,13 +189,16 @@ impl AcpRuntime { field: "sessionId", })? .to_owned(); - if let Some(dir) = self.session_file.parent() { - std::fs::create_dir_all(dir).map_err(AcpError::Io)?; - } - std::fs::write(&self.session_file, &id).map_err(AcpError::Io)?; live.model = stream::session_model(&response); - live.session = Some(id); - Ok(true) + 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(()) } async fn turn( @@ -187,8 +210,7 @@ impl AcpRuntime { ) -> Result { let cwd = session_cwd(config); let (servers, names) = mcp_servers(config)?; - let created = self.attach(live, &cwd, &servers).await?; - let session = live.session.clone().unwrap_or_default(); + let (session, created) = self.attach(live, &cwd, &servers).await?; let text = match (created, &config.system_prompt_file) { (true, Some(path)) => match std::fs::read_to_string(path) { Ok(system) => format!("{system}\n\n{prompt}"), @@ -219,6 +241,10 @@ impl AcpRuntime { reply = &mut reply => break reply.unwrap_or(Err(AcpError::Closed))?, } }; + 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; @@ -261,7 +287,7 @@ impl Runtime for AcpRuntime { } let live = match guard.as_mut() { Some(live) => live, - None => guard.insert(self.start(&session_cwd(config)).await?), + None => guard.insert(self.start(config).await?), }; let result = self.turn(live, config, prompt, sink).await; if let Err(Error::Acp( @@ -363,3 +389,114 @@ fn mcp_servers(config: &Config) -> Result<(Vec, Vec)> { let parsed: Value = serde_json::from_str(&raw).map_err(|e| fail(e.to_string()))?; Ok(stream::mcp_servers(&parsed)) } + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + use std::path::Path; + use std::sync::Arc; + + use hive_claude::{Config, NoopSink}; + + 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. + const AGENT: &str = r#" +n=0 +while IFS= read -r line; do + id=${line#*\"id\":}; id=${id%%[,\}]*} + case $line in + *'"method":"initialize"'*) + printf '{"jsonrpc":"2.0","id":%s,"result":{"protocolVersion":1,"agentCapabilities":{"loadSession":true,"mcpCapabilities":{"http":true}}}}\n' "$id" ;; + *'"method":"session/new"'*) + n=$((n+1)) + printf '{"jsonrpc":"2.0","id":%s,"result":{"sessionId":"s%s"}}\n' "$id" "$n" ;; + *'"method":"session/load"'*) + printf '{"jsonrpc":"2.0","id":%s,"result":{}}\n' "$id" ;; + *'"method":"session/prompt"'*) + printf '%s\n' "$line" >> "$1" + if [ -e "$1.failed" ]; 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" + fi ;; + esac +done +"#; + + fn runtime(dir: &Path) -> AcpRuntime { + let script = dir.join("agent.sh"); + std::fs::write(&script, AGENT).unwrap(); + let command = AcpCommand { + command: "/bin/sh".into(), + args: vec![ + script.display().to_string(), + dir.join("prompts").display().to_string(), + ], + env: BTreeMap::new(), + }; + AcpRuntime::new(command, dir.join("session"), Arc::new(|_: &_| false)) + } + + fn config(dir: &Path) -> Config { + std::fs::write(dir.join("system.md"), "SYSTEM PROMPT").unwrap(); + Config { + cwd: Some(dir.to_path_buf()), + system_prompt_file: Some(dir.join("system.md")), + ..Config::default() + } + } + + /// Which of the prompts the agent was sent carried the system prompt. + fn carried_system_prompt(dir: &Path) -> Vec { + std::fs::read_to_string(dir.join("prompts")) + .unwrap() + .lines() + .map(|line| line.contains("SYSTEM PROMPT")) + .collect() + } + + #[tokio::test] + async fn a_session_whose_first_prompt_failed_is_not_resumed() { + let dir = tempfile::tempdir().unwrap(); + let (runtime, config) = (runtime(dir.path()), config(dir.path())); + + let failed = runtime.run(&config, "one", &NoopSink).await; + assert!( + matches!(failed, Err(Error::Acp(AcpError::Rpc { .. }))), + "{failed:?}" + ); + assert!(!dir.path().join("session").exists()); + + let retried = runtime.run(&config, "two", &NoopSink).await.unwrap(); + assert!(retried.created); + let next = runtime.run(&config, "three", &NoopSink).await.unwrap(); + 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" + ); + } + + #[tokio::test] + async fn after_a_restart_a_session_whose_first_prompt_failed_is_not_loaded() { + let dir = tempfile::tempdir().unwrap(); + let config = config(dir.path()); + + let before = runtime(dir.path()); + assert!(before.run(&config, "one", &NoopSink).await.is_err()); + drop(before); + + let after = runtime(dir.path()); + let retried = after.run(&config, "two", &NoopSink).await.unwrap(); + assert!(retried.created); + + assert_eq!(carried_system_prompt(dir.path()), [true, true]); + } +} diff --git a/hive-runtime/src/acp/rpc.rs b/hive-runtime/src/acp/rpc.rs index be265c9a..076024b6 100644 --- a/hive-runtime/src/acp/rpc.rs +++ b/hive-runtime/src/acp/rpc.rs @@ -12,7 +12,8 @@ use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; use tokio::process::{Child, ChildStdin, Command}; use tokio::sync::{mpsc, oneshot}; -use super::{AcpError, PermissionPolicy}; +use super::stream::mcp_server_of; +use super::{AcpError, PermissionAsk, PermissionPolicy}; use crate::spec::AcpCommand; /// Something the agent sent that is not a response to one of our requests. @@ -42,6 +43,7 @@ impl Connection { command: &AcpCommand, cwd: &Path, permit: PermissionPolicy, + servers: Vec, ) -> Result { let mut child = Command::new(&command.command) .args(&command.args) @@ -80,6 +82,7 @@ impl Connection { pending: pending.clone(), tx, permit, + servers, }; tokio::spawn(async move { let mut lines = BufReader::new(stdout).lines(); @@ -151,6 +154,7 @@ struct Reader { pending: Pending, tx: mpsc::UnboundedSender, permit: PermissionPolicy, + servers: Vec, } impl Reader { @@ -169,7 +173,7 @@ impl Reader { let reply = match method { "session/request_permission" => json!({ "jsonrpc": "2.0", "id": id, - "result": permission_outcome(&message["params"], &self.permit), + "result": permission_outcome(&message["params"], &self.permit, &self.servers), }), _ => json!({ "jsonrpc": "2.0", "id": id, @@ -225,15 +229,21 @@ async fn write(stdin: &tokio::sync::Mutex, message: &Value) -> Resul } /// The `session/request_permission` result: the agent's first option of the -/// matching kind, allow if `permit` accepts the tool call's kind and reject -/// otherwise. With no such option the request is answered `cancelled`. -pub(super) fn permission_outcome(params: &Value, permit: &PermissionPolicy) -> Value { - let kind = params - .get("toolCall") - .and_then(|t| t.get("kind")) - .and_then(Value::as_str) - .unwrap_or("other"); - let wanted = if permit(kind) { +/// matching kind, allow if `permit` accepts the request and reject otherwise. +/// With no such option the request is answered `cancelled`. `servers` are the +/// MCP servers handed to the session, matched against the tool call's title. +pub(super) fn permission_outcome( + params: &Value, + permit: &PermissionPolicy, + servers: &[String], +) -> Value { + let tool_call = params.get("toolCall"); + let field = |k: &str| tool_call.and_then(|t| t.get(k)).and_then(Value::as_str); + let ask = PermissionAsk { + kind: field("kind").unwrap_or("other"), + mcp_server: field("title").and_then(|title| mcp_server_of(title, servers)), + }; + let wanted = if permit(&ask) { ["allow_once", "allow_always"] } else { ["reject_once", "reject_always"] @@ -257,14 +267,14 @@ pub(super) fn permission_outcome(params: &Value, permit: &PermissionPolicy) -> V #[cfg(test)] mod tests { use super::permission_outcome; - use crate::acp::PermissionPolicy; + use crate::acp::{PermissionAsk, PermissionPolicy}; use serde_json::json; use std::sync::Arc; - fn request(kind: &str) -> serde_json::Value { + fn request(kind: &str, title: &str) -> serde_json::Value { json!({ "sessionId": "s", - "toolCall": { "toolCallId": "t", "kind": kind }, + "toolCall": { "toolCallId": "t", "kind": kind, "title": title }, "options": [ { "optionId": "once", "kind": "allow_once", "name": "Allow once" }, { "optionId": "always", "kind": "allow_always", "name": "Always allow" }, @@ -273,34 +283,54 @@ mod tests { }) } - fn no_execute() -> PermissionPolicy { - Arc::new(|kind: &str| kind != "execute") + fn servers() -> Vec { + vec!["hyperhive".to_owned()] + } + + /// Allows `edit` and MCP tools, nothing else. + fn edit_or_mcp() -> PermissionPolicy { + Arc::new(|ask: &PermissionAsk<'_>| ask.kind == "edit" || ask.mcp_server.is_some()) } #[test] - fn a_permitted_kind_is_allowed_once() { + fn a_permitted_request_is_allowed_once() { assert_eq!( - permission_outcome(&request("edit"), &no_execute()), + permission_outcome(&request("edit", "edit"), &edit_or_mcp(), &servers()), json!({ "outcome": { "outcome": "selected", "optionId": "once" } }) ); } #[test] - fn a_refused_kind_is_rejected() { + fn a_refused_request_is_rejected() { assert_eq!( - permission_outcome(&request("execute"), &no_execute()), + permission_outcome(&request("execute", "bash"), &edit_or_mcp(), &servers()), json!({ "outcome": { "outcome": "selected", "optionId": "reject" } }) ); } #[test] - fn a_missing_kind_is_judged_as_other() { - let policy: PermissionPolicy = Arc::new(|kind: &str| kind == "other"); - let mut req = request("x"); + fn the_policy_sees_which_mcp_server_a_tool_belongs_to() { + let seen = Arc::new(std::sync::Mutex::new(Vec::new())); + let record = seen.clone(); + let policy: PermissionPolicy = Arc::new(move |ask: &PermissionAsk<'_>| { + record + .lock() + .unwrap() + .push((ask.kind.to_owned(), ask.mcp_server.map(str::to_owned))); + true + }); + permission_outcome(&request("other", "hyperhive_send"), &policy, &servers()); + permission_outcome(&request("other", "skill"), &policy, &servers()); + let mut req = request("x", "read"); req["toolCall"].as_object_mut().unwrap().remove("kind"); + permission_outcome(&req, &policy, &servers()); assert_eq!( - permission_outcome(&req, &policy)["outcome"]["optionId"], - "once" + *seen.lock().unwrap(), + vec![ + ("other".to_owned(), Some("hyperhive".to_owned())), + ("other".to_owned(), None), + ("other".to_owned(), None), + ] ); } @@ -309,7 +339,7 @@ mod tests { let req = json!({ "toolCall": { "kind": "execute" }, "options": [{ "optionId": "once", "kind": "allow_once" }] }); assert_eq!( - permission_outcome(&req, &no_execute()), + permission_outcome(&req, &edit_or_mcp(), &servers()), json!({ "outcome": { "outcome": "cancelled" } }) ); } diff --git a/hive-runtime/src/acp/stream.rs b/hive-runtime/src/acp/stream.rs index aadf9067..9460eed4 100644 --- a/hive-runtime/src/acp/stream.rs +++ b/hive-runtime/src/acp/stream.rs @@ -33,10 +33,6 @@ impl StreamMapper { /// `servers` are the MCP server names handed to the agent, used to give /// their tools claude's `mcp____` names. pub(super) fn new(servers: Vec) -> Self { - let mut servers = servers; - // Longest first, so a server whose name prefixes another's never - // claims the other's tools. - servers.sort_by_key(|s| std::cmp::Reverse(s.len())); Self { servers, text: String::new(), @@ -226,18 +222,30 @@ pub(super) fn canonical_tool_name(name: &str, servers: &[String]) -> String { if name.starts_with("mcp__") { return name.to_owned(); } - for server in servers { - if let Some(rest) = name - .strip_prefix(server.as_str()) - .and_then(|r| r.strip_prefix('_')) - { + split_mcp_name(name, servers).map_or_else( + || name.to_owned(), + |(server, tool)| format!("mcp__{server}__{tool}"), + ) +} + +/// Which of `servers` a tool named `_`, `__` or +/// `mcp____` belongs to, if any. +pub(super) fn mcp_server_of<'a>(name: &str, servers: &'a [String]) -> Option<&'a str> { + split_mcp_name(name, servers).map(|(server, _)| server) +} + +/// `(server, tool)` for an MCP tool name, matching the longest server name so +/// one that prefixes another's never claims the other's tools. +fn split_mcp_name<'a, 'n>(name: &'n str, servers: &'a [String]) -> Option<(&'a str, &'n str)> { + let bare = name.strip_prefix("mcp__").unwrap_or(name); + servers + .iter() + .filter_map(|server| { + let rest = bare.strip_prefix(server.as_str())?.strip_prefix('_')?; let tool = rest.strip_prefix('_').unwrap_or(rest); - if !tool.is_empty() { - return format!("mcp__{server}__{tool}"); - } - } - } - name.to_owned() + (!tool.is_empty()).then_some((server.as_str(), tool)) + }) + .max_by_key(|(server, _)| server.len()) } fn chunk_text(update: &Value) -> &str { diff --git a/hive-runtime/src/lib.rs b/hive-runtime/src/lib.rs index f9142a17..50064625 100644 --- a/hive-runtime/src/lib.rs +++ b/hive-runtime/src/lib.rs @@ -21,7 +21,7 @@ mod acp; mod claude; mod spec; -pub use acp::{AcpError, AcpRuntime, PermissionPolicy}; +pub use acp::{AcpError, AcpRuntime, PermissionAsk, PermissionPolicy}; pub use claude::ClaudeRuntime; pub use hive_claude::{ CompactionPolicy, Config, PercentPolicy, Progress, SessionStore, Sink, Telemetry, TokenUsage,