fix: skip stdin drain and set silent render mode when --acp-server is active

This commit is contained in:
2026-07-27 18:59:09 -06:00
parent d462b09f80
commit 577c51b62f
3 changed files with 64 additions and 32 deletions
+53 -25
View File
@@ -1,11 +1,12 @@
use super::types::{METHOD_NOT_FOUND, PARSE_ERROR, Request, Response}; use super::types::{METHOD_NOT_FOUND, PARSE_ERROR, Request, Response};
use crate::client::call_chat_completions_streaming; use crate::client::call_chat_completions_streaming;
use crate::config::{Input, RenderMode, RequestContext}; use crate::config::{Input, RenderMode, RequestContext};
use crate::utils;
use crate::utils::AbortSignal; use crate::utils::AbortSignal;
use anyhow::Result; use anyhow::Result;
use serde_json::{Value, json}; use serde_json::{Value, json};
use std::sync::Arc; use std::sync::Arc;
use tokio::io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader}; use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
pub(crate) struct AcpServerState { pub(crate) struct AcpServerState {
ctx: Option<RequestContext>, ctx: Option<RequestContext>,
@@ -19,22 +20,25 @@ pub async fn run_acp_server(ctx: RequestContext, abort: AbortSignal) -> Result<(
abort, abort,
session_active: false, session_active: false,
}; };
run_acp_server_with_state(tokio::io::stdin(), tokio::io::stdout(), state).await run_acp_server_with_state(tokio::io::stdin(), tokio::io::stdout(), state).await
} }
#[cfg(test)] #[cfg(test)]
pub(crate) async fn run_acp_server_on<R, W>(reader: R, writer: W) -> Result<()> pub(crate) async fn run_acp_server_on<R, W>(reader: R, writer: W) -> Result<()>
where where
R: tokio::io::AsyncRead + Unpin, R: AsyncRead + Unpin,
W: AsyncWrite + Unpin, W: AsyncWrite + Unpin,
{ {
use crate::utils::{create_abort_signal, drain_acp_permissions}; use crate::utils::{create_abort_signal, drain_acp_permissions};
drain_acp_permissions(); drain_acp_permissions();
let state = AcpServerState { let state = AcpServerState {
ctx: None, ctx: None,
abort: create_abort_signal(), abort: create_abort_signal(),
session_active: false, session_active: false,
}; };
run_acp_server_with_state(reader, writer, state).await run_acp_server_with_state(reader, writer, state).await
} }
@@ -44,7 +48,7 @@ async fn run_acp_server_with_state<R, W>(
mut state: AcpServerState, mut state: AcpServerState,
) -> Result<()> ) -> Result<()>
where where
R: tokio::io::AsyncRead + Unpin, R: AsyncRead + Unpin,
W: AsyncWrite + Unpin, W: AsyncWrite + Unpin,
{ {
let reader = BufReader::new(reader); let reader = BufReader::new(reader);
@@ -55,10 +59,12 @@ where
if line.is_empty() { if line.is_empty() {
continue; continue;
} }
if let Some(response) = dispatch(&line, &mut state).await { if let Some(response) = dispatch(&line, &mut state).await {
for params in crate::utils::drain_acp_permissions() { for params in utils::drain_acp_permissions() {
emit_notification(&mut writer, "session/request_permission", params).await?; emit_notification(&mut writer, "session/request_permission", params).await?;
} }
emit(&mut writer, &response).await?; emit(&mut writer, &response).await?;
} }
} }
@@ -72,7 +78,7 @@ async fn dispatch(raw: &str, state: &mut AcpServerState) -> Option<Response> {
Err(_) => return Some(Response::err(None, PARSE_ERROR, "Parse error")), Err(_) => return Some(Response::err(None, PARSE_ERROR, "Parse error")),
}; };
// session/cancel is a notification — handle it regardless of whether an id is present. // session/cancel is a notification. Handle it regardless of whether an id is present.
if req.method == "session/cancel" { if req.method == "session/cancel" {
handle_session_cancel(state); handle_session_cancel(state);
return if req.id.is_some() { return if req.id.is_some() {
@@ -117,6 +123,7 @@ fn handle_session_new(req: Request, state: &mut AcpServerState) -> Response {
); );
} }
state.session_active = true; state.session_active = true;
Response::ok(req.id, json!({ "sessionId": "default" })) Response::ok(req.id, json!({ "sessionId": "default" }))
} }
@@ -162,7 +169,9 @@ async fn run_prompt_turn(
let (output, tool_results) = let (output, tool_results) =
call_chat_completions_streaming(&input, client.as_ref(), ctx, abort).await?; call_chat_completions_streaming(&input, client.as_ref(), ctx, abort).await?;
let app = Arc::clone(&ctx.app.config); let app = Arc::clone(&ctx.app.config);
ctx.after_chat_completion(app.as_ref(), &input, &output, &tool_results)?; ctx.after_chat_completion(app.as_ref(), &input, &output, &tool_results)?;
Ok(output) Ok(output)
} }
@@ -210,15 +219,16 @@ async fn emit<W: AsyncWrite + Unpin>(writer: &mut W, response: &Response) -> Res
line.push('\n'); line.push('\n');
writer.write_all(line.as_bytes()).await?; writer.write_all(line.as_bytes()).await?;
writer.flush().await?; writer.flush().await?;
Ok(()) Ok(())
} }
async fn emit_notification<W: AsyncWrite + Unpin>( async fn emit_notification<W: AsyncWrite + Unpin>(
writer: &mut W, writer: &mut W,
method: &str, method: &str,
params: serde_json::Value, params: Value,
) -> Result<()> { ) -> Result<()> {
let frame = serde_json::json!({ let frame = json!({
"jsonrpc": "2.0", "jsonrpc": "2.0",
"method": method, "method": method,
"params": params, "params": params,
@@ -227,12 +237,14 @@ async fn emit_notification<W: AsyncWrite + Unpin>(
line.push('\n'); line.push('\n');
writer.write_all(line.as_bytes()).await?; writer.write_all(line.as_bytes()).await?;
writer.flush().await?; writer.flush().await?;
Ok(()) Ok(())
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use std::str;
#[tokio::test] #[tokio::test]
async fn all_stdout_is_valid_json_rpc() { async fn all_stdout_is_valid_json_rpc() {
@@ -246,8 +258,8 @@ mod tests {
.unwrap(); .unwrap();
for line in output.split(|&b| b == b'\n').filter(|l| !l.is_empty()) { for line in output.split(|&b| b == b'\n').filter(|l| !l.is_empty()) {
let s = std::str::from_utf8(line).expect("non-UTF8 in ACP stdout"); let s = str::from_utf8(line).expect("non-UTF8 in ACP stdout");
let _: serde_json::Value = serde_json::from_str(s) let _: Value = serde_json::from_str(s)
.unwrap_or_else(|_| panic!("ACP stdout not valid JSON: {s}")); .unwrap_or_else(|_| panic!("ACP stdout not valid JSON: {s}"));
} }
} }
@@ -264,7 +276,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["error"]["code"], METHOD_NOT_FOUND); assert_eq!(v["error"]["code"], METHOD_NOT_FOUND);
assert_eq!(v["id"], 2); assert_eq!(v["id"], 2);
} }
@@ -278,7 +291,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["error"]["code"], PARSE_ERROR); assert_eq!(v["error"]["code"], PARSE_ERROR);
} }
@@ -289,9 +303,11 @@ mod tests {
"\n", "\n",
); );
let mut output = Vec::new(); let mut output = Vec::new();
run_acp_server_on(input.as_bytes(), &mut output) run_acp_server_on(input.as_bytes(), &mut output)
.await .await
.unwrap(); .unwrap();
assert!(output.is_empty()); assert!(output.is_empty());
} }
@@ -307,7 +323,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["id"], 1); assert_eq!(v["id"], 1);
assert_eq!(v["result"]["name"], "coyote"); assert_eq!(v["result"]["name"], "coyote");
assert!(v["result"]["version"].is_string()); assert!(v["result"]["version"].is_string());
@@ -325,7 +342,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["id"], 10); assert_eq!(v["id"], 10);
assert_eq!(v["result"]["sessionId"], "default"); assert_eq!(v["result"]["sessionId"], "default");
} }
@@ -345,8 +363,9 @@ mod tests {
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let mut lines = s.lines().filter(|l| !l.is_empty()); let mut lines = s.lines().filter(|l| !l.is_empty());
let first: serde_json::Value = serde_json::from_str(lines.next().unwrap()).unwrap();
let second: serde_json::Value = serde_json::from_str(lines.next().unwrap()).unwrap(); let first: Value = serde_json::from_str(lines.next().unwrap()).unwrap();
let second: Value = serde_json::from_str(lines.next().unwrap()).unwrap();
assert!(first["result"]["sessionId"].is_string()); assert!(first["result"]["sessionId"].is_string());
assert_eq!(second["error"]["code"], -32000); assert_eq!(second["error"]["code"], -32000);
} }
@@ -363,7 +382,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["id"], 3); assert_eq!(v["id"], 3);
assert_eq!(v["error"]["code"], -32000); assert_eq!(v["error"]["code"], -32000);
} }
@@ -383,8 +403,9 @@ mod tests {
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let mut lines = s.lines().filter(|l| !l.is_empty()); let mut lines = s.lines().filter(|l| !l.is_empty());
let first: serde_json::Value = serde_json::from_str(lines.next().unwrap()).unwrap();
let second: serde_json::Value = serde_json::from_str(lines.next().unwrap()).unwrap(); let first: Value = serde_json::from_str(lines.next().unwrap()).unwrap();
let second: Value = serde_json::from_str(lines.next().unwrap()).unwrap();
assert_eq!(first["result"]["sessionId"], "default"); assert_eq!(first["result"]["sessionId"], "default");
assert_eq!(second["id"], 2); assert_eq!(second["id"], 2);
assert!( assert!(
@@ -397,9 +418,11 @@ mod tests {
async fn session_cancel_notification_produces_no_output() { async fn session_cancel_notification_produces_no_output() {
let input = concat!(r#"{"jsonrpc":"2.0","method":"session/cancel"}"#, "\n",); let input = concat!(r#"{"jsonrpc":"2.0","method":"session/cancel"}"#, "\n",);
let mut output = Vec::new(); let mut output = Vec::new();
run_acp_server_on(input.as_bytes(), &mut output) run_acp_server_on(input.as_bytes(), &mut output)
.await .await
.unwrap(); .unwrap();
assert!(output.is_empty()); assert!(output.is_empty());
} }
@@ -415,7 +438,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["id"], 99); assert_eq!(v["id"], 99);
assert!(v["result"].is_object()); assert!(v["result"].is_object());
} }
@@ -432,7 +456,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["id"], 5); assert_eq!(v["id"], 5);
assert_eq!(v["error"]["code"], -32602); assert_eq!(v["error"]["code"], -32602);
} }
@@ -452,8 +477,9 @@ mod tests {
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let mut lines = s.lines().filter(|l| !l.is_empty()); let mut lines = s.lines().filter(|l| !l.is_empty());
let _first: serde_json::Value = serde_json::from_str(lines.next().unwrap()).unwrap();
let second: serde_json::Value = serde_json::from_str(lines.next().unwrap()).unwrap(); let _first: Value = serde_json::from_str(lines.next().unwrap()).unwrap();
let second: Value = serde_json::from_str(lines.next().unwrap()).unwrap();
assert_eq!(second["id"], 2); assert_eq!(second["id"], 2);
assert_eq!(second["error"]["code"], -32000); assert_eq!(second["error"]["code"], -32000);
} }
@@ -470,7 +496,8 @@ mod tests {
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["id"], 6); assert_eq!(v["id"], 6);
assert!(v["error"].is_object()); assert!(v["error"].is_object());
} }
@@ -481,13 +508,14 @@ mod tests {
emit_notification( emit_notification(
&mut output, &mut output,
"session/request_permission", "session/request_permission",
serde_json::json!({"action": "confirm", "question": "Proceed?"}), json!({"action": "confirm", "question": "Proceed?"}),
) )
.await .await
.unwrap(); .unwrap();
let s = String::from_utf8(output).unwrap(); let s = String::from_utf8(output).unwrap();
let v: serde_json::Value = serde_json::from_str(s.trim()).unwrap(); let v: Value = serde_json::from_str(s.trim()).unwrap();
assert_eq!(v["jsonrpc"], "2.0"); assert_eq!(v["jsonrpc"], "2.0");
assert_eq!(v["method"], "session/request_permission"); assert_eq!(v["method"], "session/request_permission");
assert!(v["params"]["action"].is_string()); assert!(v["params"]["action"].is_string());
+7 -4
View File
@@ -41,7 +41,7 @@ use clap_complete::CompleteEnv;
use client::ClientConfig; use client::ClientConfig;
use inquire::{Select, Text, set_global_render_config}; use inquire::{Select, Text, set_global_render_config};
use log::{LevelFilter, warn}; use log::{LevelFilter, warn};
use log4rs::append::console::ConsoleAppender; use log4rs::append::console::{ConsoleAppender, Target};
use log4rs::append::rolling_file::RollingFileAppender; use log4rs::append::rolling_file::RollingFileAppender;
use log4rs::append::rolling_file::policy::compound::CompoundPolicy; use log4rs::append::rolling_file::policy::compound::CompoundPolicy;
use log4rs::append::rolling_file::policy::compound::roll::fixed_window::FixedWindowRoller; use log4rs::append::rolling_file::policy::compound::roll::fixed_window::FixedWindowRoller;
@@ -76,7 +76,7 @@ async fn main() -> Result<()> {
return Ok(()); return Ok(());
} }
let text = cli.text()?; let text = if cli.acp_server { None } else { cli.text()? };
let working_mode = if !cli.acp_server && text.is_none() && cli.file.is_empty() { let working_mode = if !cli.acp_server && text.is_none() && cli.file.is_empty() {
WorkingMode::Repl WorkingMode::Repl
} else { } else {
@@ -87,10 +87,13 @@ async fn main() -> Result<()> {
if cli.headless && !cli.acp_server && text.is_none() && cli.file.is_empty() { if cli.headless && !cli.acp_server && text.is_none() && cli.file.is_empty() {
bail!("--headless requires a prompt argument; REPL mode is not supported"); bail!("--headless requires a prompt argument; REPL mode is not supported");
} }
unsafe { unsafe {
env::set_var("AUTO_CONFIRM", "true"); env::set_var("AUTO_CONFIRM", "true");
} }
HEADLESS.store(true, Ordering::SeqCst); HEADLESS.store(true, Ordering::SeqCst);
if cli.acp_server { if cli.acp_server {
ACP_SERVER.store(true, Ordering::SeqCst); ACP_SERVER.store(true, Ordering::SeqCst);
} }
@@ -233,7 +236,7 @@ async fn main() -> Result<()> {
} }
} }
if cli.headless { if cli.headless || cli.acp_server {
ctx.render_mode = RenderMode::Silent; ctx.render_mode = RenderMode::Silent;
} }
@@ -715,7 +718,7 @@ fn setup_logger(acp_mode: bool) -> Result<Option<PathBuf>> {
None => { None => {
let mut builder = ConsoleAppender::builder().encoder(encoder); let mut builder = ConsoleAppender::builder().encoder(encoder);
if acp_mode { if acp_mode {
builder = builder.target(log4rs::append::console::Target::Stderr); builder = builder.target(Target::Stderr);
} }
let console_appender = builder.build(); let console_appender = builder.build();
log4rs::init_config(init_console_logger(log_level, log_filter, console_appender))?; log4rs::init_config(init_console_logger(log_level, log_filter, console_appender))?;
+4 -3
View File
@@ -32,6 +32,7 @@ use fancy_regex::Regex;
use fuzzy_matcher::{FuzzyMatcher, skim::SkimMatcherV2}; use fuzzy_matcher::{FuzzyMatcher, skim::SkimMatcherV2};
use is_terminal::IsTerminal; use is_terminal::IsTerminal;
use nu_ansi_term::Color; use nu_ansi_term::Color;
use serde_json::Value;
use std::borrow::Cow; use std::borrow::Cow;
use std::collections::VecDeque; use std::collections::VecDeque;
use std::sync::atomic::AtomicBool; use std::sync::atomic::AtomicBool;
@@ -48,15 +49,15 @@ pub static IS_STDOUT_TERMINAL: LazyLock<bool> = LazyLock::new(|| std::io::stdout
pub static HEADLESS: AtomicBool = AtomicBool::new(false); pub static HEADLESS: AtomicBool = AtomicBool::new(false);
pub static ACP_SERVER: AtomicBool = AtomicBool::new(false); pub static ACP_SERVER: AtomicBool = AtomicBool::new(false);
static ACP_PERMISSION_QUEUE: Mutex<VecDeque<serde_json::Value>> = Mutex::new(VecDeque::new()); static ACP_PERMISSION_QUEUE: Mutex<VecDeque<Value>> = Mutex::new(VecDeque::new());
pub fn queue_acp_permission(notification: serde_json::Value) { pub fn queue_acp_permission(notification: Value) {
if let Ok(mut q) = ACP_PERMISSION_QUEUE.lock() { if let Ok(mut q) = ACP_PERMISSION_QUEUE.lock() {
q.push_back(notification); q.push_back(notification);
} }
} }
pub fn drain_acp_permissions() -> Vec<serde_json::Value> { pub fn drain_acp_permissions() -> Vec<Value> {
ACP_PERMISSION_QUEUE ACP_PERMISSION_QUEUE
.lock() .lock()
.map(|mut q| q.drain(..).collect()) .map(|mut q| q.drain(..).collect())