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
|
|
@ -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<dyn Fn(&str) -> 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 `<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.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
|
|
@ -105,8 +116,10 @@ impl AcpRuntime {
|
|||
}
|
||||
}
|
||||
|
||||
async fn start(&self, cwd: &Path) -> Result<Live> {
|
||||
let conn = Connection::spawn(&self.command, cwd, self.permit.clone())?;
|
||||
async fn start(&self, config: &Config) -> Result<Live> {
|
||||
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<bool> {
|
||||
/// 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<Progress> {
|
||||
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<Value>, Vec<String>)> {
|
|||
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<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]);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue