hive-runtime: retry an ACP session whose first prompt failed on that session
A new session is only recorded once its first prompt is answered, so a failed first prompt made the next turn run `session/new` again and left the first session behind in the agent: in the same process, and after a restart. The new session's id is now kept in `<session file>.pending` until that prompt is answered. The next turn attaches it (in-process, or through `session/load` after a restart) and still prepends the system prompt, since the session has not had an answered prompt yet. Recording the session removes the pending file, and `archive` drops it. Raised by argus in the #4812 review. Refs #4391
This commit is contained in:
parent
bed72ce280
commit
8ec51ca774
1 changed files with 117 additions and 54 deletions
|
|
@ -100,7 +100,8 @@ pub struct AcpRuntime {
|
||||||
struct Live {
|
struct Live {
|
||||||
conn: Connection,
|
conn: Connection,
|
||||||
load_session: bool,
|
load_session: bool,
|
||||||
session: Option<String>,
|
/// The recorded session, or a new one whose first prompt is unanswered.
|
||||||
|
loaded: Option<String>,
|
||||||
model: Option<String>,
|
model: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -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<Live> {
|
async fn start(&self, config: &Config) -> Result<Live> {
|
||||||
let (_, servers) = mcp_servers(config)?;
|
let (_, servers) = mcp_servers(config)?;
|
||||||
let cwd = session_cwd(config);
|
let cwd = session_cwd(config);
|
||||||
|
|
@ -142,45 +151,49 @@ impl AcpRuntime {
|
||||||
Ok(Live {
|
Ok(Live {
|
||||||
conn,
|
conn,
|
||||||
load_session: caps["loadSession"] == Value::Bool(true),
|
load_session: caps["loadSession"] == Value::Bool(true),
|
||||||
session: None,
|
loaded: None,
|
||||||
model: None,
|
model: None,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Attach the durable session: the one already loaded, else the one named
|
/// Attach the durable session: the recorded one (named by the session
|
||||||
/// by the session file, else a new one. Returns its id and whether it is
|
/// file), else a new one whose first prompt was never answered (named by
|
||||||
/// new. A new session is not recorded here: [`Self::turn`] records it once
|
/// the pending file), else a new one. Returns its id and whether its first
|
||||||
/// its first prompt, the one carrying the system prompt, has been answered.
|
/// prompt, the one carrying the system prompt, is still to be answered;
|
||||||
|
/// [`Self::turn`] records the session once it has been.
|
||||||
async fn attach(
|
async fn attach(
|
||||||
&self,
|
&self,
|
||||||
live: &mut Live,
|
live: &mut Live,
|
||||||
cwd: &Path,
|
cwd: &Path,
|
||||||
servers: &[Value],
|
servers: &[Value],
|
||||||
) -> Result<(String, bool)> {
|
) -> Result<(String, bool)> {
|
||||||
let persisted = std::fs::read_to_string(&self.session_file)
|
let known = [
|
||||||
.ok()
|
(read_id(&self.session_file), false),
|
||||||
.map(|s| s.trim().to_owned())
|
(read_id(&self.pending_file()), true),
|
||||||
.filter(|s| !s.is_empty());
|
]
|
||||||
if let Some(id) = persisted {
|
.into_iter()
|
||||||
if live.session.as_deref() == Some(id.as_str()) {
|
.filter_map(|(id, new)| Some((id?, new)));
|
||||||
return Ok((id, false));
|
for (id, new) in known {
|
||||||
|
if live.loaded.as_deref() == Some(id.as_str()) {
|
||||||
|
return Ok((id, new));
|
||||||
|
}
|
||||||
|
if !live.load_session {
|
||||||
|
continue;
|
||||||
}
|
}
|
||||||
if live.load_session {
|
|
||||||
let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers });
|
let params = json!({ "sessionId": id, "cwd": cwd, "mcpServers": servers });
|
||||||
match live.conn.request("session/load", params).await {
|
match live.conn.request("session/load", params).await {
|
||||||
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.clone());
|
live.loaded = Some(id.clone());
|
||||||
return Ok((id, false));
|
return Ok((id, new));
|
||||||
}
|
}
|
||||||
Err(e @ AcpError::Rpc { .. }) => {
|
Err(e @ AcpError::Rpc { .. }) => {
|
||||||
tracing::warn!(error = %e, "ACP session/load failed; starting a new session");
|
tracing::warn!(error = %e, session = %id, "ACP session/load failed");
|
||||||
}
|
}
|
||||||
Err(e) => return Err(e.into()),
|
Err(e) => return Err(e.into()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
let params = json!({ "cwd": cwd, "mcpServers": servers });
|
let params = json!({ "cwd": cwd, "mcpServers": servers });
|
||||||
let response = live.conn.request("session/new", params).await?;
|
let response = live.conn.request("session/new", params).await?;
|
||||||
let id = response["sessionId"]
|
let id = response["sessionId"]
|
||||||
|
|
@ -191,15 +204,14 @@ impl AcpRuntime {
|
||||||
})?
|
})?
|
||||||
.to_owned();
|
.to_owned();
|
||||||
live.model = stream::session_model(&response);
|
live.model = stream::session_model(&response);
|
||||||
|
live.loaded = Some(id.clone());
|
||||||
|
write_id(&self.pending_file(), &id)?;
|
||||||
Ok((id, true))
|
Ok((id, true))
|
||||||
}
|
}
|
||||||
|
|
||||||
fn record_session(&self, id: &str) -> std::result::Result<(), AcpError> {
|
fn record_session(&self, id: &str) -> std::result::Result<(), AcpError> {
|
||||||
if let Some(dir) = self.session_file.parent() {
|
write_id(&self.session_file, id)?;
|
||||||
std::fs::create_dir_all(dir)?;
|
remove_if_present(&self.pending_file())
|
||||||
}
|
|
||||||
std::fs::write(&self.session_file, id)?;
|
|
||||||
Ok(())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn turn(
|
async fn turn(
|
||||||
|
|
@ -245,7 +257,6 @@ impl AcpRuntime {
|
||||||
if created {
|
if created {
|
||||||
self.record_session(&session)?;
|
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;
|
||||||
|
|
@ -308,8 +319,10 @@ impl Runtime for AcpRuntime {
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Moves the session file aside; the agent keeps its own copy of the
|
/// 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<Option<PathBuf>> {
|
fn archive(&self) -> Result<Option<PathBuf>> {
|
||||||
|
remove_if_present(&self.pending_file()).map_err(Error::Acp)?;
|
||||||
if !self.session_file.exists() {
|
if !self.session_file.exists() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
@ -370,6 +383,28 @@ fn discard_stale(conn: &mut Connection) -> std::result::Result<(), AcpError> {
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn read_id(file: &Path) -> Option<String> {
|
||||||
|
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 {
|
fn session_cwd(config: &Config) -> PathBuf {
|
||||||
config
|
config
|
||||||
.cwd
|
.cwd
|
||||||
|
|
@ -402,34 +437,45 @@ mod tests {
|
||||||
use super::{AcpError, AcpRuntime};
|
use super::{AcpError, AcpRuntime};
|
||||||
use crate::{AcpCommand, Error, Runtime};
|
use crate::{AcpCommand, Error, Runtime};
|
||||||
|
|
||||||
/// An ACP agent in plain `sh`: answers `initialize`, numbers its sessions
|
/// An ACP agent in plain `sh`. It answers `initialize`, numbers its
|
||||||
/// `s1`, `s2`, …, appends every `session/prompt` line to the file named by
|
/// sessions `s1`, `s2`, …, appends every method it is sent to
|
||||||
/// its argument, and fails the first prompt it is sent.
|
/// `<log>.methods` (and `start` when it starts), and every
|
||||||
|
/// `session/prompt` line to `<log>`. 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#"
|
const AGENT: &str = r#"
|
||||||
n=0
|
n=0
|
||||||
|
printf 'start\n' >> "$1.methods"
|
||||||
while IFS= read -r line; do
|
while IFS= read -r line; do
|
||||||
id=${line#*\"id\":}; id=${id%%[,\}]*}
|
id=${line#*\"id\":}; id=${id%%[,\}]*}
|
||||||
case $line in
|
m=${line#*\"method\":\"}; m=${m%%\"*}
|
||||||
*'"method":"initialize"'*)
|
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" ;;
|
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))
|
n=$((n+1))
|
||||||
printf '{"jsonrpc":"2.0","id":%s,"result":{"sessionId":"s%s"}}\n' "$id" "$n" ;;
|
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" ;;
|
printf '{"jsonrpc":"2.0","id":%s,"result":{}}\n' "$id" ;;
|
||||||
*'"method":"session/prompt"'*)
|
session/prompt)
|
||||||
printf '%s\n' "$line" >> "$1"
|
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"
|
printf '{"jsonrpc":"2.0","id":%s,"result":{"stopReason":"end_turn"}}\n' "$id"
|
||||||
else
|
else
|
||||||
: > "$1.failed"
|
: > "$1.once"
|
||||||
printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32603,"message":"provider down"}}\n' "$id"
|
case $2 in
|
||||||
|
fail)
|
||||||
|
printf '{"jsonrpc":"2.0","id":%s,"error":{"code":-32603,"message":"provider down"}}\n' "$id" ;;
|
||||||
|
esac
|
||||||
fi ;;
|
fi ;;
|
||||||
esac
|
esac
|
||||||
done
|
done
|
||||||
"#;
|
"#;
|
||||||
|
|
||||||
fn runtime(dir: &Path) -> AcpRuntime {
|
fn runtime(dir: &Path, mode: &str) -> AcpRuntime {
|
||||||
let script = dir.join("agent.sh");
|
let script = dir.join("agent.sh");
|
||||||
std::fs::write(&script, AGENT).unwrap();
|
std::fs::write(&script, AGENT).unwrap();
|
||||||
let command = AcpCommand {
|
let command = AcpCommand {
|
||||||
|
|
@ -437,6 +483,7 @@ done
|
||||||
args: vec![
|
args: vec![
|
||||||
script.display().to_string(),
|
script.display().to_string(),
|
||||||
dir.join("prompts").display().to_string(),
|
dir.join("prompts").display().to_string(),
|
||||||
|
mode.to_owned(),
|
||||||
],
|
],
|
||||||
env: BTreeMap::new(),
|
env: BTreeMap::new(),
|
||||||
};
|
};
|
||||||
|
|
@ -461,10 +508,23 @@ done
|
||||||
.collect()
|
.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]
|
#[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 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;
|
let failed = runtime.run(&config, "one", &NoopSink).await;
|
||||||
assert!(
|
assert!(
|
||||||
|
|
@ -479,25 +539,28 @@ done
|
||||||
assert!(!next.created);
|
assert!(!next.created);
|
||||||
|
|
||||||
assert_eq!(carried_system_prompt(dir.path()), [true, true, false]);
|
assert_eq!(carried_system_prompt(dir.path()), [true, true, false]);
|
||||||
assert_eq!(
|
assert_eq!(recorded(dir.path()), "s1");
|
||||||
std::fs::read_to_string(dir.path().join("session")).unwrap(),
|
assert_eq!(count(dir.path(), "session/new"), 1);
|
||||||
"s2"
|
assert!(!dir.path().join("session.pending").exists());
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[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 dir = tempfile::tempdir().unwrap();
|
||||||
let config = config(dir.path());
|
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());
|
assert!(before.run(&config, "one", &NoopSink).await.is_err());
|
||||||
drop(before);
|
drop(before);
|
||||||
|
|
||||||
let after = runtime(dir.path());
|
let after = runtime(dir.path(), "fail");
|
||||||
let retried = after.run(&config, "two", &NoopSink).await.unwrap();
|
let retried = after.run(&config, "two", &NoopSink).await.unwrap();
|
||||||
assert!(retried.created);
|
assert!(retried.created);
|
||||||
|
|
||||||
assert_eq!(carried_system_prompt(dir.path()), [true, true]);
|
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");
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue