mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-07 03:00:29 +00:00
fix(acp): tolerate clean stdio exit after final response
This commit is contained in:
parent
cbf81d79fb
commit
d6d7f18c23
3 changed files with 227 additions and 15 deletions
|
|
@ -10,7 +10,9 @@ use agent_client_protocol::{
|
|||
};
|
||||
use fabro_sandbox::{
|
||||
Error as SandboxError, Result as SandboxResult, Sandbox, StderrCollector, StdioProcessHandle,
|
||||
StdioProcessTermination,
|
||||
};
|
||||
use fabro_types::CommandTermination;
|
||||
use futures::io::BufReader;
|
||||
use futures::sink::unfold;
|
||||
use futures::{AsyncBufReadExt, AsyncWriteExt, Stream};
|
||||
|
|
@ -20,6 +22,8 @@ use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt};
|
|||
|
||||
use crate::command::AcpCommand;
|
||||
|
||||
const CLEAN_EXIT_PROTOCOL_GRACE: Duration = Duration::from_millis(500);
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct TransportState {
|
||||
handle: Arc<TokioMutex<Option<StdioProcessHandle>>>,
|
||||
|
|
@ -128,12 +132,12 @@ impl ConnectTo<Client> for SandboxAcpTransport {
|
|||
},
|
||||
));
|
||||
|
||||
let protocol = agent_client_protocol::ConnectTo::<Client>::connect_to(
|
||||
let mut protocol = Box::pin(agent_client_protocol::ConnectTo::<Client>::connect_to(
|
||||
Lines::new(outgoing_sink, incoming_lines),
|
||||
client,
|
||||
);
|
||||
));
|
||||
tokio::select! {
|
||||
result = protocol => {
|
||||
result = &mut protocol => {
|
||||
if let Err(err) = handle.terminate().await {
|
||||
tracing::warn!(error = %err, "Failed to terminate ACP process after protocol completion");
|
||||
}
|
||||
|
|
@ -143,14 +147,30 @@ impl ConnectTo<Client> for SandboxAcpTransport {
|
|||
termination = handle.wait() => {
|
||||
let termination = termination.map_err(ProtocolError::into_internal_error)?;
|
||||
let stderr = stderr.tail_string().await;
|
||||
let exit_code = termination
|
||||
.exit_code
|
||||
.map_or_else(|| "unknown".to_string(), |code| code.to_string());
|
||||
Err(internal_error(format!(
|
||||
"ACP process exited before protocol completed: termination={}, exit_code={exit_code}, stderr={stderr}",
|
||||
termination.termination,
|
||||
)))
|
||||
if termination.termination == CommandTermination::Exited
|
||||
&& termination.exit_code == Some(0)
|
||||
{
|
||||
// Stdio agents commonly exit immediately after writing their final response.
|
||||
// Process wait can observe that exit before the line reader drains stdout.
|
||||
if let Ok(result) = timeout(CLEAN_EXIT_PROTOCOL_GRACE, &mut protocol).await {
|
||||
return result;
|
||||
}
|
||||
}
|
||||
Err(process_exited_before_protocol_completed(termination, &stderr))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn process_exited_before_protocol_completed(
|
||||
termination: StdioProcessTermination,
|
||||
stderr: &str,
|
||||
) -> ProtocolError {
|
||||
let exit_code = termination
|
||||
.exit_code
|
||||
.map_or_else(|| "unknown".to_string(), |code| code.to_string());
|
||||
internal_error(format!(
|
||||
"ACP process exited before protocol completed: termination={}, exit_code={exit_code}, stderr={stderr}",
|
||||
termination.termination,
|
||||
))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,10 +5,11 @@ use std::time::Duration;
|
|||
|
||||
use agent_client_protocol::schema::StopReason;
|
||||
use fabro_acp::{AcpError, AcpRunRequest, AcpRunResult, resolve_acp_command, run_acp_turn};
|
||||
use fabro_sandbox::test_support::MockSandbox;
|
||||
use fabro_sandbox::test_support::{MockSandbox, MockStdioProcess};
|
||||
use fabro_sandbox::{LocalSandbox, Sandbox, shell_quote};
|
||||
use fabro_util::error::collect_chain;
|
||||
use tokio::fs::{read_to_string, write};
|
||||
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader, DuplexStream};
|
||||
use tokio::process::Command;
|
||||
use tokio::sync::Notify;
|
||||
use tokio::time::{sleep, timeout};
|
||||
|
|
@ -61,6 +62,30 @@ async fn stdio_spawn_failure_returns_sandbox_error() {
|
|||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn clean_stdio_exit_after_final_response_completes_turn() {
|
||||
let sandbox = MockSandbox::linux();
|
||||
sandbox.set_stdio_process(mock_acp_stdio_process("end_turn"));
|
||||
let sandbox: Arc<dyn Sandbox> = Arc::new(sandbox);
|
||||
let command = resolve_acp_command(Some("mock-acp-agent")).expect("resolve ACP command");
|
||||
|
||||
let result = run_acp_turn(AcpRunRequest {
|
||||
command,
|
||||
prompt: "hello".to_string(),
|
||||
cwd: "/workspace".to_string(),
|
||||
timeout_ms: Some(ACP_TEST_TIMEOUT_MS),
|
||||
env: HashMap::new(),
|
||||
sandbox,
|
||||
cancel_token: CancellationToken::new(),
|
||||
on_activity: None,
|
||||
})
|
||||
.await
|
||||
.expect("clean ACP process exit should not preempt final protocol response");
|
||||
|
||||
assert_eq!(result.text, "hello from acp");
|
||||
assert_eq!(result.stop_reason, StopReason::EndTurn);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn session_lifecycle_initializes_sends_prompt_and_aggregates_text() {
|
||||
let tempdir = tempfile::tempdir().expect("create tempdir");
|
||||
|
|
@ -447,3 +472,96 @@ async fn process_is_running(pid: &str) -> bool {
|
|||
.find(|ch| !ch.is_whitespace())
|
||||
.is_none_or(|state| !matches!(state, 'Z' | 'z'))
|
||||
}
|
||||
|
||||
fn mock_acp_stdio_process(stop_reason: &'static str) -> MockStdioProcess {
|
||||
MockStdioProcess::new(move |stdin, mut stdout, _stderr| {
|
||||
tokio::spawn(async move {
|
||||
let mut lines = BufReader::new(stdin).lines();
|
||||
while let Some(line) = lines.next_line().await.expect("read mock ACP stdin") {
|
||||
let message: serde_json::Value =
|
||||
serde_json::from_str(&line).expect("parse mock ACP request");
|
||||
let method = message
|
||||
.get("method")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.expect("mock ACP request method");
|
||||
let id = message
|
||||
.get("id")
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
|
||||
match method {
|
||||
"initialize" => {
|
||||
write_acp_response(
|
||||
&mut stdout,
|
||||
id,
|
||||
serde_json::json!({
|
||||
"protocolVersion": 1,
|
||||
"agentCapabilities": {},
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
"session/new" => {
|
||||
write_acp_response(
|
||||
&mut stdout,
|
||||
id,
|
||||
serde_json::json!({ "sessionId": "sess-1" }),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
"session/prompt" => {
|
||||
write_acp_message(
|
||||
&mut stdout,
|
||||
serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"method": "session/update",
|
||||
"params": {
|
||||
"sessionId": "sess-1",
|
||||
"update": {
|
||||
"sessionUpdate": "agent_message_chunk",
|
||||
"content": { "type": "text", "text": "hello from acp" }
|
||||
}
|
||||
}
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
write_acp_response(
|
||||
&mut stdout,
|
||||
id,
|
||||
serde_json::json!({ "stopReason": stop_reason }),
|
||||
)
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
});
|
||||
})
|
||||
}
|
||||
|
||||
async fn write_acp_response(
|
||||
stdout: &mut DuplexStream,
|
||||
id: serde_json::Value,
|
||||
result: serde_json::Value,
|
||||
) {
|
||||
write_acp_message(
|
||||
stdout,
|
||||
serde_json::json!({
|
||||
"jsonrpc": "2.0",
|
||||
"id": id,
|
||||
"result": result,
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
async fn write_acp_message(stdout: &mut DuplexStream, message: serde_json::Value) {
|
||||
let mut line = serde_json::to_vec(&message).expect("serialize mock ACP message");
|
||||
line.push(b'\n');
|
||||
stdout
|
||||
.write_all(&line)
|
||||
.await
|
||||
.expect("write mock ACP stdout");
|
||||
stdout.flush().await.expect("flush mock ACP stdout");
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
use std::collections::HashMap;
|
||||
use std::sync::Mutex;
|
||||
use std::time::Duration;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use fabro_types::CommandTermination;
|
||||
use tokio::fs;
|
||||
use tokio::io::duplex;
|
||||
use tokio::io::{DuplexStream, duplex};
|
||||
use tokio::time::sleep;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::sandbox::StdioProcessControl;
|
||||
|
|
@ -43,6 +45,7 @@ pub struct MockSandbox {
|
|||
pub delete_calls: Mutex<u32>,
|
||||
pub event_callback: Option<SandboxEventCallback>,
|
||||
pub stdio_process_error: Option<String>,
|
||||
pub stdio_process: Mutex<Option<MockStdioProcess>>,
|
||||
}
|
||||
|
||||
impl MockSandbox {
|
||||
|
|
@ -69,6 +72,13 @@ impl MockSandbox {
|
|||
.lock()
|
||||
.expect("delete_calls lock poisoned")
|
||||
}
|
||||
|
||||
pub fn set_stdio_process(&self, process: MockStdioProcess) {
|
||||
*self
|
||||
.stdio_process
|
||||
.lock()
|
||||
.expect("stdio_process lock poisoned") = Some(process);
|
||||
}
|
||||
}
|
||||
|
||||
impl MockSandbox {
|
||||
|
|
@ -108,11 +118,48 @@ impl Default for MockSandbox {
|
|||
delete_calls: Mutex::new(0),
|
||||
event_callback: None,
|
||||
stdio_process_error: None,
|
||||
stdio_process: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct MockStdioProcessControl;
|
||||
type StdioProcessDriver =
|
||||
Box<dyn FnOnce(DuplexStream, DuplexStream, StderrCollector) + Send + 'static>;
|
||||
|
||||
pub struct MockStdioProcess {
|
||||
exit_code: Option<i32>,
|
||||
wait_delay: Duration,
|
||||
driver: StdioProcessDriver,
|
||||
}
|
||||
|
||||
impl MockStdioProcess {
|
||||
pub fn new(
|
||||
driver: impl FnOnce(DuplexStream, DuplexStream, StderrCollector) + Send + 'static,
|
||||
) -> Self {
|
||||
Self {
|
||||
exit_code: Some(0),
|
||||
wait_delay: Duration::ZERO,
|
||||
driver: Box::new(driver),
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_exit_code(mut self, exit_code: Option<i32>) -> Self {
|
||||
self.exit_code = exit_code;
|
||||
self
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn with_wait_delay(mut self, wait_delay: Duration) -> Self {
|
||||
self.wait_delay = wait_delay;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
struct MockStdioProcessControl {
|
||||
exit_code: Option<i32>,
|
||||
wait_delay: Duration,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl StdioProcessControl for MockStdioProcessControl {
|
||||
|
|
@ -121,7 +168,10 @@ impl StdioProcessControl for MockStdioProcessControl {
|
|||
}
|
||||
|
||||
async fn wait(&self) -> crate::Result<StdioProcessTermination> {
|
||||
Ok(StdioProcessTermination::exited(Some(0)))
|
||||
if !self.wait_delay.is_zero() {
|
||||
sleep(self.wait_delay).await;
|
||||
}
|
||||
Ok(StdioProcessTermination::exited(self.exit_code))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -233,13 +283,37 @@ impl Sandbox for MockSandbox {
|
|||
return Err(crate::Error::message(error.clone()));
|
||||
}
|
||||
|
||||
if let Some(process) = self
|
||||
.stdio_process
|
||||
.lock()
|
||||
.expect("stdio_process lock poisoned")
|
||||
.take()
|
||||
{
|
||||
let (stdin, stdin_reader) = duplex(4096);
|
||||
let (stdout_writer, stdout) = duplex(4096);
|
||||
let stderr = StderrCollector::new(DEFAULT_EXEC_OUTPUT_TAIL_BYTES);
|
||||
(process.driver)(stdin_reader, stdout_writer, stderr.clone());
|
||||
return Ok(StdioProcess {
|
||||
stdin: Box::pin(stdin),
|
||||
stdout: Box::pin(stdout),
|
||||
stderr,
|
||||
handle: StdioProcessHandle::new(MockStdioProcessControl {
|
||||
exit_code: process.exit_code,
|
||||
wait_delay: process.wait_delay,
|
||||
}),
|
||||
});
|
||||
}
|
||||
|
||||
let (stdin, _stdin_read) = duplex(1024);
|
||||
let (_stdout_write, stdout) = duplex(1024);
|
||||
Ok(StdioProcess {
|
||||
stdin: Box::pin(stdin),
|
||||
stdout: Box::pin(stdout),
|
||||
stderr: StderrCollector::new(DEFAULT_EXEC_OUTPUT_TAIL_BYTES),
|
||||
handle: StdioProcessHandle::new(MockStdioProcessControl),
|
||||
handle: StdioProcessHandle::new(MockStdioProcessControl {
|
||||
exit_code: Some(0),
|
||||
wait_delay: Duration::ZERO,
|
||||
}),
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue