hive-runtime: record a new ACP session only once its first prompt is answered
The session id was written to the session file right after session/new, but the system prompt rides on the first prompt only. If that prompt failed, the next turn (or the next harness start) resumed the recorded session as not-new and the system prompt never reached it. The id is now written after the first session/prompt gets its reply. A failed first prompt leaves nothing recorded, so the next turn starts a new session and sends the system prompt again. Chosen over a separate "system prompt delivered" marker: one file, and "recorded" already means "usable". Tests drive AcpRuntime against a scripted sh agent that fails the first prompt: in-process and across a restart, the retry is a new session carrying the system prompt. Also: PermissionPolicy now sees a PermissionAsk (kind plus the MCP server the tool belongs to, matched by name against the servers handed to the session), so a caller can tell MCP tool calls from other `other` requests. Refs #4391
This commit is contained in:
parent
1b24edf4b4
commit
a2ab40cc69
6 changed files with 240 additions and 61 deletions
1
Cargo.lock
generated
1
Cargo.lock
generated
|
|
@ -1987,6 +1987,7 @@ dependencies = [
|
||||||
"hive-claude",
|
"hive-claude",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
|
"tempfile",
|
||||||
"thiserror 2.0.18",
|
"thiserror 2.0.18",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tracing",
|
"tracing",
|
||||||
|
|
|
||||||
|
|
@ -14,3 +14,6 @@ serde_json.workspace = true
|
||||||
thiserror.workspace = true
|
thiserror.workspace = true
|
||||||
tokio.workspace = true
|
tokio.workspace = true
|
||||||
tracing.workspace = true
|
tracing.workspace = true
|
||||||
|
|
||||||
|
[dev-dependencies]
|
||||||
|
tempfile = "3"
|
||||||
|
|
|
||||||
|
|
@ -28,10 +28,21 @@ const STDERR_TAIL: usize = 20;
|
||||||
const SETTLE_QUIET: std::time::Duration = std::time::Duration::from_millis(500);
|
const SETTLE_QUIET: std::time::Duration = std::time::Duration::from_millis(500);
|
||||||
const SETTLE_MAX: std::time::Duration = std::time::Duration::from_secs(5);
|
const SETTLE_MAX: std::time::Duration = std::time::Duration::from_secs(5);
|
||||||
|
|
||||||
/// Decides the agent's `session/request_permission` requests from the tool
|
/// One `session/request_permission` request, as a [`PermissionPolicy`] sees it.
|
||||||
/// call's ACP `kind` (`read`, `edit`, `execute`, `fetch`, …; `other` when the
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
/// agent gives none). `true` allows the call once, `false` rejects it.
|
pub struct PermissionAsk<'a> {
|
||||||
pub type PermissionPolicy = Arc<dyn Fn(&str) -> bool + Send + Sync>;
|
/// 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 `<server>_<tool>`,
|
||||||
|
/// `<server>__<tool>` or `mcp__<server>__<tool>`.
|
||||||
|
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<dyn Fn(&PermissionAsk<'_>) -> bool + Send + Sync>;
|
||||||
|
|
||||||
/// Why an ACP operation failed.
|
/// Why an ACP operation failed.
|
||||||
#[derive(Debug, thiserror::Error)]
|
#[derive(Debug, thiserror::Error)]
|
||||||
|
|
@ -105,8 +116,10 @@ impl AcpRuntime {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn start(&self, cwd: &Path) -> Result<Live> {
|
async fn start(&self, config: &Config) -> Result<Live> {
|
||||||
let conn = Connection::spawn(&self.command, cwd, self.permit.clone())?;
|
let (_, servers) = mcp_servers(config)?;
|
||||||
|
let cwd = session_cwd(config);
|
||||||
|
let conn = Connection::spawn(&self.command, &cwd, self.permit.clone(), servers)?;
|
||||||
let init = conn
|
let init = conn
|
||||||
.request(
|
.request(
|
||||||
"initialize",
|
"initialize",
|
||||||
|
|
@ -134,15 +147,22 @@ impl AcpRuntime {
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Attach the durable session: the one already loaded, else the one named
|
/// 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.
|
/// by the session file, else a new one. Returns its id and whether it is
|
||||||
async fn attach(&self, live: &mut Live, cwd: &Path, servers: &[Value]) -> Result<bool> {
|
/// 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)
|
let persisted = std::fs::read_to_string(&self.session_file)
|
||||||
.ok()
|
.ok()
|
||||||
.map(|s| s.trim().to_owned())
|
.map(|s| s.trim().to_owned())
|
||||||
.filter(|s| !s.is_empty());
|
.filter(|s| !s.is_empty());
|
||||||
if let Some(id) = persisted {
|
if let Some(id) = persisted {
|
||||||
if live.session.as_deref() == Some(id.as_str()) {
|
if live.session.as_deref() == Some(id.as_str()) {
|
||||||
return Ok(false);
|
return Ok((id, false));
|
||||||
}
|
}
|
||||||
if live.load_session {
|
if live.load_session {
|
||||||
let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers });
|
let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers });
|
||||||
|
|
@ -150,8 +170,8 @@ impl AcpRuntime {
|
||||||
Ok(response) => {
|
Ok(response) => {
|
||||||
discard_stale(&mut live.conn)?;
|
discard_stale(&mut live.conn)?;
|
||||||
live.model = stream::session_model(&response);
|
live.model = stream::session_model(&response);
|
||||||
live.session = Some(id);
|
live.session = Some(id.clone());
|
||||||
return Ok(false);
|
return Ok((id, false));
|
||||||
}
|
}
|
||||||
Err(e @ AcpError::Rpc { .. }) => {
|
Err(e @ AcpError::Rpc { .. }) => {
|
||||||
tracing::warn!(error = %e, "ACP session/load failed; starting a new session");
|
tracing::warn!(error = %e, "ACP session/load failed; starting a new session");
|
||||||
|
|
@ -169,13 +189,16 @@ impl AcpRuntime {
|
||||||
field: "sessionId",
|
field: "sessionId",
|
||||||
})?
|
})?
|
||||||
.to_owned();
|
.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.model = stream::session_model(&response);
|
||||||
live.session = Some(id);
|
Ok((id, true))
|
||||||
Ok(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(
|
async fn turn(
|
||||||
|
|
@ -187,8 +210,7 @@ impl AcpRuntime {
|
||||||
) -> Result<Progress> {
|
) -> Result<Progress> {
|
||||||
let cwd = session_cwd(config);
|
let cwd = session_cwd(config);
|
||||||
let (servers, names) = mcp_servers(config)?;
|
let (servers, names) = mcp_servers(config)?;
|
||||||
let created = self.attach(live, &cwd, &servers).await?;
|
let (session, created) = self.attach(live, &cwd, &servers).await?;
|
||||||
let session = live.session.clone().unwrap_or_default();
|
|
||||||
let text = match (created, &config.system_prompt_file) {
|
let text = match (created, &config.system_prompt_file) {
|
||||||
(true, Some(path)) => match std::fs::read_to_string(path) {
|
(true, Some(path)) => match std::fs::read_to_string(path) {
|
||||||
Ok(system) => format!("{system}\n\n{prompt}"),
|
Ok(system) => format!("{system}\n\n{prompt}"),
|
||||||
|
|
@ -219,6 +241,10 @@ impl AcpRuntime {
|
||||||
reply = &mut reply => break reply.unwrap_or(Err(AcpError::Closed))?,
|
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
|
// Updates may still arrive after the reply: an agent can forward its
|
||||||
// stream asynchronously. Take them until the agent goes quiet.
|
// stream asynchronously. Take them until the agent goes quiet.
|
||||||
let settled = tokio::time::Instant::now() + SETTLE_MAX;
|
let settled = tokio::time::Instant::now() + SETTLE_MAX;
|
||||||
|
|
@ -261,7 +287,7 @@ impl Runtime for AcpRuntime {
|
||||||
}
|
}
|
||||||
let live = match guard.as_mut() {
|
let live = match guard.as_mut() {
|
||||||
Some(live) => live,
|
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;
|
let result = self.turn(live, config, prompt, sink).await;
|
||||||
if let Err(Error::Acp(
|
if let Err(Error::Acp(
|
||||||
|
|
@ -363,3 +389,114 @@ fn mcp_servers(config: &Config) -> Result<(Vec<Value>, Vec<String>)> {
|
||||||
let parsed: Value = serde_json::from_str(&raw).map_err(|e| fail(e.to_string()))?;
|
let parsed: Value = serde_json::from_str(&raw).map_err(|e| fail(e.to_string()))?;
|
||||||
Ok(stream::mcp_servers(&parsed))
|
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<bool> {
|
||||||
|
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]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -12,7 +12,8 @@ use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
|
||||||
use tokio::process::{Child, ChildStdin, Command};
|
use tokio::process::{Child, ChildStdin, Command};
|
||||||
use tokio::sync::{mpsc, oneshot};
|
use tokio::sync::{mpsc, oneshot};
|
||||||
|
|
||||||
use super::{AcpError, PermissionPolicy};
|
use super::stream::mcp_server_of;
|
||||||
|
use super::{AcpError, PermissionAsk, PermissionPolicy};
|
||||||
use crate::spec::AcpCommand;
|
use crate::spec::AcpCommand;
|
||||||
|
|
||||||
/// Something the agent sent that is not a response to one of our requests.
|
/// Something the agent sent that is not a response to one of our requests.
|
||||||
|
|
@ -42,6 +43,7 @@ impl Connection {
|
||||||
command: &AcpCommand,
|
command: &AcpCommand,
|
||||||
cwd: &Path,
|
cwd: &Path,
|
||||||
permit: PermissionPolicy,
|
permit: PermissionPolicy,
|
||||||
|
servers: Vec<String>,
|
||||||
) -> Result<Self, AcpError> {
|
) -> Result<Self, AcpError> {
|
||||||
let mut child = Command::new(&command.command)
|
let mut child = Command::new(&command.command)
|
||||||
.args(&command.args)
|
.args(&command.args)
|
||||||
|
|
@ -80,6 +82,7 @@ impl Connection {
|
||||||
pending: pending.clone(),
|
pending: pending.clone(),
|
||||||
tx,
|
tx,
|
||||||
permit,
|
permit,
|
||||||
|
servers,
|
||||||
};
|
};
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
let mut lines = BufReader::new(stdout).lines();
|
let mut lines = BufReader::new(stdout).lines();
|
||||||
|
|
@ -151,6 +154,7 @@ struct Reader {
|
||||||
pending: Pending,
|
pending: Pending,
|
||||||
tx: mpsc::UnboundedSender<Incoming>,
|
tx: mpsc::UnboundedSender<Incoming>,
|
||||||
permit: PermissionPolicy,
|
permit: PermissionPolicy,
|
||||||
|
servers: Vec<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Reader {
|
impl Reader {
|
||||||
|
|
@ -169,7 +173,7 @@ impl Reader {
|
||||||
let reply = match method {
|
let reply = match method {
|
||||||
"session/request_permission" => json!({
|
"session/request_permission" => json!({
|
||||||
"jsonrpc": "2.0", "id": id,
|
"jsonrpc": "2.0", "id": id,
|
||||||
"result": permission_outcome(&message["params"], &self.permit),
|
"result": permission_outcome(&message["params"], &self.permit, &self.servers),
|
||||||
}),
|
}),
|
||||||
_ => json!({
|
_ => json!({
|
||||||
"jsonrpc": "2.0", "id": id,
|
"jsonrpc": "2.0", "id": id,
|
||||||
|
|
@ -225,15 +229,21 @@ async fn write(stdin: &tokio::sync::Mutex<ChildStdin>, message: &Value) -> Resul
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The `session/request_permission` result: the agent's first option of the
|
/// 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
|
/// matching kind, allow if `permit` accepts the request and reject otherwise.
|
||||||
/// otherwise. With no such option the request is answered `cancelled`.
|
/// With no such option the request is answered `cancelled`. `servers` are the
|
||||||
pub(super) fn permission_outcome(params: &Value, permit: &PermissionPolicy) -> Value {
|
/// MCP servers handed to the session, matched against the tool call's title.
|
||||||
let kind = params
|
pub(super) fn permission_outcome(
|
||||||
.get("toolCall")
|
params: &Value,
|
||||||
.and_then(|t| t.get("kind"))
|
permit: &PermissionPolicy,
|
||||||
.and_then(Value::as_str)
|
servers: &[String],
|
||||||
.unwrap_or("other");
|
) -> Value {
|
||||||
let wanted = if permit(kind) {
|
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"]
|
["allow_once", "allow_always"]
|
||||||
} else {
|
} else {
|
||||||
["reject_once", "reject_always"]
|
["reject_once", "reject_always"]
|
||||||
|
|
@ -257,14 +267,14 @@ pub(super) fn permission_outcome(params: &Value, permit: &PermissionPolicy) -> V
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::permission_outcome;
|
use super::permission_outcome;
|
||||||
use crate::acp::PermissionPolicy;
|
use crate::acp::{PermissionAsk, PermissionPolicy};
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
|
|
||||||
fn request(kind: &str) -> serde_json::Value {
|
fn request(kind: &str, title: &str) -> serde_json::Value {
|
||||||
json!({
|
json!({
|
||||||
"sessionId": "s",
|
"sessionId": "s",
|
||||||
"toolCall": { "toolCallId": "t", "kind": kind },
|
"toolCall": { "toolCallId": "t", "kind": kind, "title": title },
|
||||||
"options": [
|
"options": [
|
||||||
{ "optionId": "once", "kind": "allow_once", "name": "Allow once" },
|
{ "optionId": "once", "kind": "allow_once", "name": "Allow once" },
|
||||||
{ "optionId": "always", "kind": "allow_always", "name": "Always allow" },
|
{ "optionId": "always", "kind": "allow_always", "name": "Always allow" },
|
||||||
|
|
@ -273,34 +283,54 @@ mod tests {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
fn no_execute() -> PermissionPolicy {
|
fn servers() -> Vec<String> {
|
||||||
Arc::new(|kind: &str| kind != "execute")
|
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]
|
#[test]
|
||||||
fn a_permitted_kind_is_allowed_once() {
|
fn a_permitted_request_is_allowed_once() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
permission_outcome(&request("edit"), &no_execute()),
|
permission_outcome(&request("edit", "edit"), &edit_or_mcp(), &servers()),
|
||||||
json!({ "outcome": { "outcome": "selected", "optionId": "once" } })
|
json!({ "outcome": { "outcome": "selected", "optionId": "once" } })
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn a_refused_kind_is_rejected() {
|
fn a_refused_request_is_rejected() {
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
permission_outcome(&request("execute"), &no_execute()),
|
permission_outcome(&request("execute", "bash"), &edit_or_mcp(), &servers()),
|
||||||
json!({ "outcome": { "outcome": "selected", "optionId": "reject" } })
|
json!({ "outcome": { "outcome": "selected", "optionId": "reject" } })
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn a_missing_kind_is_judged_as_other() {
|
fn the_policy_sees_which_mcp_server_a_tool_belongs_to() {
|
||||||
let policy: PermissionPolicy = Arc::new(|kind: &str| kind == "other");
|
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
|
||||||
let mut req = request("x");
|
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");
|
req["toolCall"].as_object_mut().unwrap().remove("kind");
|
||||||
|
permission_outcome(&req, &policy, &servers());
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
permission_outcome(&req, &policy)["outcome"]["optionId"],
|
*seen.lock().unwrap(),
|
||||||
"once"
|
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" },
|
let req = json!({ "toolCall": { "kind": "execute" },
|
||||||
"options": [{ "optionId": "once", "kind": "allow_once" }] });
|
"options": [{ "optionId": "once", "kind": "allow_once" }] });
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
permission_outcome(&req, &no_execute()),
|
permission_outcome(&req, &edit_or_mcp(), &servers()),
|
||||||
json!({ "outcome": { "outcome": "cancelled" } })
|
json!({ "outcome": { "outcome": "cancelled" } })
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -33,10 +33,6 @@ impl StreamMapper {
|
||||||
/// `servers` are the MCP server names handed to the agent, used to give
|
/// `servers` are the MCP server names handed to the agent, used to give
|
||||||
/// their tools claude's `mcp__<server>__<tool>` names.
|
/// their tools claude's `mcp__<server>__<tool>` names.
|
||||||
pub(super) fn new(servers: Vec<String>) -> Self {
|
pub(super) fn new(servers: Vec<String>) -> 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 {
|
Self {
|
||||||
servers,
|
servers,
|
||||||
text: String::new(),
|
text: String::new(),
|
||||||
|
|
@ -226,18 +222,30 @@ pub(super) fn canonical_tool_name(name: &str, servers: &[String]) -> String {
|
||||||
if name.starts_with("mcp__") {
|
if name.starts_with("mcp__") {
|
||||||
return name.to_owned();
|
return name.to_owned();
|
||||||
}
|
}
|
||||||
for server in servers {
|
split_mcp_name(name, servers).map_or_else(
|
||||||
if let Some(rest) = name
|
|| name.to_owned(),
|
||||||
.strip_prefix(server.as_str())
|
|(server, tool)| format!("mcp__{server}__{tool}"),
|
||||||
.and_then(|r| r.strip_prefix('_'))
|
)
|
||||||
{
|
}
|
||||||
|
|
||||||
|
/// Which of `servers` a tool named `<server>_<tool>`, `<server>__<tool>` or
|
||||||
|
/// `mcp__<server>__<tool>` 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);
|
let tool = rest.strip_prefix('_').unwrap_or(rest);
|
||||||
if !tool.is_empty() {
|
(!tool.is_empty()).then_some((server.as_str(), tool))
|
||||||
return format!("mcp__{server}__{tool}");
|
})
|
||||||
}
|
.max_by_key(|(server, _)| server.len())
|
||||||
}
|
|
||||||
}
|
|
||||||
name.to_owned()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn chunk_text(update: &Value) -> &str {
|
fn chunk_text(update: &Value) -> &str {
|
||||||
|
|
|
||||||
|
|
@ -21,7 +21,7 @@ mod acp;
|
||||||
mod claude;
|
mod claude;
|
||||||
mod spec;
|
mod spec;
|
||||||
|
|
||||||
pub use acp::{AcpError, AcpRuntime, PermissionPolicy};
|
pub use acp::{AcpError, AcpRuntime, PermissionAsk, PermissionPolicy};
|
||||||
pub use claude::ClaudeRuntime;
|
pub use claude::ClaudeRuntime;
|
||||||
pub use hive_claude::{
|
pub use hive_claude::{
|
||||||
CompactionPolicy, Config, PercentPolicy, Progress, SessionStore, Sink, Telemetry, TokenUsage,
|
CompactionPolicy, Config, PercentPolicy, Progress, SessionStore, Sink, Telemetry, TokenUsage,
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue