use std::sync::Arc;
use axum::{
extract::{
Path, State, WebSocketUpgrade,
ws::{Message, WebSocket},
},
response::IntoResponse,
};
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize};
use tokio::sync::Mutex;
use tracing::{debug, info};
use crate::AppState;
use crate::acp::chat_persistence;
use crate::acp::permission::PermissionRequestEvent;
use crate::acp::terminal::TerminalActivity;
use crate::acp::{AcpClient, ImageInput, ResourceInput};
use crate::api::agents::load_agent;
const MAX_PROMPT_IMAGES: usize = 3;
const MAX_AT_REFERENCES: usize = 8;
const MAX_AT_FILE_BYTES: usize = 64 * 1024;
fn extract_at_paths(text: &str) -> Vec<String> {
static RE: std::sync::OnceLock<regex::Regex> = std::sync::OnceLock::new();
let re = RE.get_or_init(|| regex::Regex::new(r"(?:^|\s)@([^\s@]+)").unwrap());
let mut out: Vec<String> = Vec::new();
for cap in re.captures_iter(text) {
let p = &cap[1];
if !out.iter().any(|e| e == p) {
out.push(p.to_string());
if out.len() >= MAX_AT_REFERENCES {
break;
}
}
}
out
}
async fn resolve_at_references(
db: &sqlx::SqlitePool,
session_id: &str,
text: &str,
) -> Vec<ResourceInput> {
let paths = extract_at_paths(text);
if paths.is_empty() {
return Vec::new();
}
let row: Option<(String,)> = sqlx::query_as("SELECT workspace_path FROM sessions WHERE id = ?")
.bind(session_id)
.fetch_optional(db)
.await
.ok()
.flatten();
let Some((ws_path,)) = row else {
return Vec::new();
};
let base = std::path::PathBuf::from(ws_path);
let mut out = Vec::new();
for rel in paths {
let abs = match crate::fs::sanitize_path(&base, &rel) {
Ok(p) => p,
Err(e) => {
debug!("@ 引用跳过(路径无效): {}: {}", rel, e);
continue;
}
};
if abs.is_dir() {
debug!("@ 引用跳过(是目录): {}", rel);
continue;
}
let content = match tokio::fs::read_to_string(&abs).await {
Ok(c) => c,
Err(e) => {
debug!("@ 引用跳过(读取失败): {}: {}", rel, e);
continue;
}
};
let text = if content.len() > MAX_AT_FILE_BYTES {
let mut end = MAX_AT_FILE_BYTES;
while !content.is_char_boundary(end) {
end -= 1;
}
format!("{}\n… [content truncated at 64KB]", &content[..end])
} else {
content
};
out.push(ResourceInput { uri: format!("file://{}", abs.display()), label: rel, text });
}
out
}
pub async fn ws_acp_handler(
ws: WebSocketUpgrade,
Path(session_id): Path<String>,
State(state): State<AppState>,
) -> impl IntoResponse {
info!("ACP WS upgrade request: session_id={}", session_id);
ws.on_upgrade(move |socket| handle_acp_ws(socket, session_id, state))
}
#[derive(Debug, Deserialize)]
#[serde(tag = "type")]
enum AcpClientMessage {
#[serde(rename = "prompt")]
Prompt {
text: String,
#[serde(default)]
images: Vec<ImageInput>,
},
#[serde(rename = "cancel")]
Cancel,
#[serde(rename = "load_session")]
LoadSession,
#[serde(rename = "permission_response")]
PermissionResponse { id: String, option_id: String },
#[serde(rename = "set_config_option")]
SetConfigOption { config_id: String, value: String },
}
#[derive(Debug, Serialize)]
#[serde(tag = "type")]
enum AcpServerMessage<'a> {
#[serde(rename = "error")]
Error {
#[serde(skip_serializing_if = "Option::is_none")]
code: Option<&'a str>,
message: &'a str,
},
#[serde(rename = "session_update")]
SessionUpdate { data: serde_json::Value },
#[serde(rename = "prompt_done")]
PromptDone { stop_reason: &'a str },
#[serde(rename = "prompt_error")]
PromptError { message: &'a str },
#[serde(rename = "terminal_activity")]
TerminalActivity {
id: String,
command: String,
args: Vec<String>,
status: String,
exit_code: Option<u32>,
},
#[serde(rename = "replay_start")]
ReplayStart,
#[serde(rename = "replay_end")]
ReplayEnd,
#[serde(rename = "process_alive")]
ProcessAlive { alive: bool },
#[serde(rename = "permission_request")]
PermissionRequest { id: &'a str, request: &'a serde_json::Value },
#[serde(rename = "capabilities")]
Capabilities { image: bool },
}
fn extract_text_from_notification(data: &serde_json::Value) -> Option<String> {
let update = data.get("update")?;
let obj = update.as_object()?;
let chunk = if let Some(c) = obj.get("AgentMessageChunk") {
c
} else if obj.get("sessionUpdate").and_then(|v| v.as_str()) == Some("agent_message_chunk") {
update
} else {
return None;
};
let content = chunk.get("content")?;
if let Some(text_obj) = content.get("Text").or_else(|| content.get("text"))
&& let Some(t) = text_obj.get("text").and_then(|v| v.as_str())
{
return Some(t.to_string());
}
if let Some(t) = content.get("text").and_then(|v| v.as_str()) {
return Some(t.to_string());
}
if let Some(t) = chunk.get("text").and_then(|v| v.as_str()) {
return Some(t.to_string());
}
None
}
async fn spawn_notify_task(
mut rx: tokio::sync::broadcast::Receiver<
agent_client_protocol::schema::v1::SessionNotification,
>,
notify_tx: tokio::sync::mpsc::Sender<Message>,
buf: Arc<Mutex<String>>,
) {
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(notification) => {
let data = serde_json::to_value(¬ification).unwrap_or_default();
if let Some(text) = extract_text_from_notification(&data) {
buf.lock().await.push_str(&text);
}
let msg = serde_json::to_string(&AcpServerMessage::SessionUpdate { data })
.unwrap_or_default();
if notify_tx.send(Message::Text(msg.into())).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
tracing::warn!(
"ACP WS subscriber lagged by {} messages; dropped stale updates",
n
);
}
}
}
});
}
async fn spawn_crash_task(
mut rx: tokio::sync::broadcast::Receiver<String>,
notify_tx: tokio::sync::mpsc::Sender<Message>,
) {
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(reason) => {
let msg =
serde_json::to_string(&AcpServerMessage::PromptError { message: &reason })
.unwrap_or_default();
if notify_tx.send(Message::Text(msg.into())).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
tracing::warn!(
"ACP crash-event subscriber lagged by {} messages; dropped stale events",
n
);
}
}
}
});
}
async fn spawn_terminal_task(
mut rx: tokio::sync::broadcast::Receiver<TerminalActivity>,
notify_tx: tokio::sync::mpsc::Sender<Message>,
) {
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(ev) => {
let (id, command, args, status, exit_code) = match ev {
TerminalActivity::Created { id, command, args } => {
(id, command, args, "created".to_string(), None)
}
TerminalActivity::Exited { id, exit_code } => {
(id, String::new(), Vec::new(), "exited".to_string(), exit_code)
}
};
let msg = serde_json::to_string(&AcpServerMessage::TerminalActivity {
id,
command,
args,
status,
exit_code,
})
.unwrap_or_default();
if notify_tx.send(Message::Text(msg.into())).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
tracing::warn!(
"ACP terminal-event subscriber lagged by {} messages; dropped stale events",
n
);
}
}
}
});
}
async fn spawn_permission_task(
mut rx: tokio::sync::broadcast::Receiver<PermissionRequestEvent>,
notify_tx: tokio::sync::mpsc::Sender<Message>,
) {
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(event) => {
let msg = serde_json::to_string(&AcpServerMessage::PermissionRequest {
id: &event.id,
request: &event.request,
})
.unwrap_or_default();
if notify_tx.send(Message::Text(msg.into())).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
tracing::warn!(
"ACP permission-event subscriber lagged by {} messages; dropped stale events",
n
);
}
}
}
});
}
async fn handle_acp_ws(socket: WebSocket, session_id: String, state: AppState) {
let (mut ws_tx, mut ws_rx) = socket.split();
let (notify_tx, mut notify_rx) = tokio::sync::mpsc::channel::<Message>(64);
let assistant_buf: Arc<Mutex<String>> = Arc::new(Mutex::new(String::new()));
let mut client: Option<Arc<AcpClient>> = match state.acp_supervisor.get(&session_id).await {
Some(c) => {
info!("ACP WS connected: session_id={} (supervisor hit)", session_id);
let rx = c.session_update_subscribe();
spawn_notify_task(rx, notify_tx.clone(), assistant_buf.clone()).await;
let perm_rx = c.permission_subscribe();
spawn_permission_task(perm_rx, notify_tx.clone()).await;
let crash_rx = c.crash_subscribe();
spawn_crash_task(crash_rx, notify_tx.clone()).await;
let term_rx = c.terminal_event_subscribe();
spawn_terminal_task(term_rx, notify_tx.clone()).await;
if let Some(notif) = c.initial_config_notification() {
let data = serde_json::to_value(¬if).unwrap_or_default();
let msg = serde_json::to_string(&AcpServerMessage::SessionUpdate { data })
.unwrap_or_default();
let _ = notify_tx.send(Message::Text(msg.into())).await;
}
if let Some(notif) = c.initial_commands_notification() {
let data = serde_json::to_value(¬if).unwrap_or_default();
let msg = serde_json::to_string(&AcpServerMessage::SessionUpdate { data })
.unwrap_or_default();
let _ = notify_tx.send(Message::Text(msg.into())).await;
}
let msg = serde_json::to_string(&AcpServerMessage::Capabilities {
image: c.supports_image(),
})
.unwrap_or_default();
let _ = notify_tx.send(Message::Text(msg.into())).await;
Some(c)
}
None => {
info!("ACP WS: session_id={} not in supervisor, keeping alive for restore", session_id);
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: Some("session_not_found"),
message: "ACP session not found",
})
.unwrap();
let _ = ws_tx.send(Message::Text(msg.into())).await;
None
}
};
let mut proc_rx = state.acp_supervisor.process_event_subscribe();
let _ = notify_tx
.send(Message::Text(
serde_json::to_string(&AcpServerMessage::ProcessAlive { alive: client.is_some() })
.unwrap_or_default()
.into(),
))
.await;
let db = state.db.clone();
let sid = session_id.clone();
loop {
tokio::select! {
msg = ws_rx.next() => {
match msg {
Some(Ok(Message::Text(text))) => {
match serde_json::from_str::<AcpClientMessage>(&text) {
Ok(AcpClientMessage::Prompt { text: prompt_text, images }) => {
let Some(ref c) = client else {
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: Some("session_not_found"),
message: "no active ACP session",
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
continue;
};
if images.len() > MAX_PROMPT_IMAGES {
let msg = serde_json::to_string(&AcpServerMessage::PromptError {
message: &format!("too many images (max {})", MAX_PROMPT_IMAGES),
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
continue;
}
if !images.is_empty() && !c.supports_image() {
let msg = serde_json::to_string(&AcpServerMessage::PromptError {
message: "agent does not support image input",
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
continue;
}
let blocks_json = if images.is_empty() {
None
} else {
let mut arr = Vec::new();
if !prompt_text.is_empty() {
arr.push(serde_json::json!({
"type": "text", "text": prompt_text,
}));
}
for img in &images {
arr.push(serde_json::json!({
"type": "image",
"mimeType": img.mime_type,
"data": img.data,
}));
}
serde_json::to_string(&arr).ok()
};
let _ = chat_persistence::insert_message(
&db, &sid, "user", &prompt_text, blocks_json.as_deref(),
).await;
let c = c.clone();
let resources = resolve_at_references(&db, &sid, &prompt_text).await;
c.mark_prompt_active();
let tx = notify_tx.clone();
let db2 = db.clone();
let sid2 = sid.clone();
let buf2 = assistant_buf.clone();
tokio::spawn(async move {
match c.send_prompt(&prompt_text, images, resources).await {
Ok(resp) => {
c.mark_prompt_idle();
tokio::task::yield_now().await;
let assistant_text = buf2.lock().await.drain(..).collect::<String>();
if !assistant_text.is_empty() {
let _ = chat_persistence::insert_message(
&db2, &sid2, "assistant", &assistant_text, None,
).await;
}
let reason = format!("{:?}", resp.stop_reason);
let msg = serde_json::to_string(
&AcpServerMessage::PromptDone { stop_reason: &reason },
).unwrap_or_default();
let _ = tx.send(Message::Text(msg.into())).await;
}
Err(e) => {
c.mark_prompt_idle();
buf2.lock().await.clear();
let err_msg = format!("{}", e);
let msg = serde_json::to_string(
&AcpServerMessage::PromptError { message: &err_msg },
).unwrap_or_default();
let _ = tx.send(Message::Text(msg.into())).await;
}
}
});
}
Ok(AcpClientMessage::Cancel) => {
if let Some(ref c) = client {
c.mark_prompt_idle();
if let Err(e) = c.cancel() {
let err_msg = format!("取消 agent 失败: {}", e);
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: Some("cancel_failed"),
message: &err_msg,
})
.unwrap_or_default();
let _ = notify_tx.send(Message::Text(msg.into())).await;
}
}
}
Ok(AcpClientMessage::LoadSession) => {
let row: Option<(String, String, String)> = sqlx::query_as(
"SELECT agent_id, acp_session_id, workspace_path FROM sessions WHERE id = ? AND runtime_kind = 'acp'",
)
.bind(&sid)
.fetch_optional(&db)
.await
.ok()
.flatten();
let Some((agent_id, acp_sid, ws_path)) = row else {
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: None,
message: "session row not found or not ACP",
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
continue;
};
let Some(agent) = load_agent(&db, &agent_id).await else {
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: None,
message: "agent config not found",
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
continue;
};
let cwd = std::path::PathBuf::from(&ws_path);
match AcpClient::spawn_and_load(agent, cwd.clone(), acp_sid.clone()).await {
Ok(new_client) => {
let new_client = Arc::new(new_client);
if !new_client.supports_load_session() {
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: Some("load_not_supported"),
message: "agent does not support session/load",
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
let c = Arc::try_unwrap(new_client).ok();
if let Some(c) = c { c.disconnect().await; }
continue;
}
if let Some(old) = state.acp_supervisor.dispose(&sid).await
&& let Ok(c) = Arc::try_unwrap(old) {
c.disconnect().await;
}
state.acp_supervisor.insert(sid.clone(), new_client.clone()).await;
let perm_rx = new_client.permission_subscribe();
spawn_permission_task(perm_rx, notify_tx.clone()).await;
let crash_rx = new_client.crash_subscribe();
spawn_crash_task(crash_rx, notify_tx.clone()).await;
let term_rx = new_client.terminal_event_subscribe();
spawn_terminal_task(term_rx, notify_tx.clone()).await;
client = Some(new_client.clone());
let cap_msg = serde_json::to_string(&AcpServerMessage::Capabilities {
image: new_client.supports_image(),
}).unwrap_or_default();
let _ = notify_tx.send(Message::Text(cap_msg.into())).await;
let replay_msg = serde_json::to_string(&AcpServerMessage::ReplayStart).unwrap_or_default();
let _ = ws_tx.send(Message::Text(replay_msg.into())).await;
let tx = notify_tx.clone();
let buf2 = assistant_buf.clone();
tokio::spawn(async move {
let mut replay_rx = new_client.session_update_subscribe();
let result = new_client.load_session(&acp_sid, cwd).await;
loop {
match replay_rx.try_recv() {
Ok(notif) => {
let data = serde_json::to_value(¬if).unwrap_or_default();
let frame = serde_json::to_string(
&AcpServerMessage::SessionUpdate { data },
)
.unwrap_or_default();
if tx.send(Message::Text(frame.into())).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::TryRecvError::Empty) => break,
Err(tokio::sync::broadcast::error::TryRecvError::Lagged(n)) => {
tracing::warn!("ACP replay subscriber lagged by {} messages", n);
break;
}
Err(tokio::sync::broadcast::error::TryRecvError::Closed) => break,
}
}
let msg = match result {
Ok(()) => serde_json::to_string(&AcpServerMessage::ReplayEnd).unwrap_or_default(),
Err(e) => serde_json::to_string(&AcpServerMessage::Error {
code: Some("load_failed"),
message: &format!("session/load failed: {}", e),
}).unwrap_or_default(),
};
let _ = tx.send(Message::Text(msg.into())).await;
spawn_notify_task(new_client.session_update_subscribe(), tx.clone(), buf2).await;
});
}
Err(e) => {
let err_msg = format!("failed to spawn agent: {}", e);
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: Some("spawn_failed"),
message: &err_msg,
}).unwrap_or_default();
let _ = ws_tx.send(Message::Text(msg.into())).await;
}
}
}
Ok(AcpClientMessage::PermissionResponse { id, option_id }) => {
if let Some(ref c) = client {
c.resolve_permission(&id, &option_id).await;
}
}
Ok(AcpClientMessage::SetConfigOption { config_id, value }) => {
if let Some(ref c) = client
&& let Err(e) = c.set_config_option(&config_id, &value).await {
let err_msg = format!("配置�? {} 设置失败: {}", config_id, e);
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: Some("config_option_failed"),
message: &err_msg,
})
.unwrap_or_default();
let _ = notify_tx.send(Message::Text(msg.into())).await;
}
}
Err(e) => {
let err_msg = format!("invalid message: {}", e);
let msg = serde_json::to_string(&AcpServerMessage::Error {
code: None,
message: &err_msg,
})
.unwrap_or_default();
if ws_tx.send(Message::Text(msg.into())).await.is_err() {
break;
}
}
}
}
Some(Ok(Message::Close(_))) | None => break,
other => {
tracing::warn!(
session_id = %session_id,
?other,
"received unsupported websocket frame (non-text/non-close); ignoring"
);
}
}
}
msg = notify_rx.recv() => {
match msg {
Some(ws_msg) => {
if ws_tx.send(ws_msg).await.is_err() {
break;
}
}
None => break,
}
}
msg = proc_rx.recv() => {
match msg {
Ok(evt) if evt.session_id == session_id => {
let frame = serde_json::to_string(&AcpServerMessage::ProcessAlive {
alive: evt.alive,
})
.unwrap_or_default();
if notify_tx.send(Message::Text(frame.into())).await.is_err() {
break;
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
tracing::warn!(
session_id = %session_id,
skipped = n,
"process-alive channel lagged; process events may be stale"
);
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
tracing::debug!(session_id = %session_id, "process-alive channel closed");
}
Ok(_) => {}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::extract_at_paths;
#[test]
fn extracts_basic_paths() {
assert_eq!(
extract_at_paths("看看 @src/main.rs 和 @README.md 的内容"),
vec!["src/main.rs", "README.md"]
);
}
#[test]
fn extracts_at_start_and_after_newline() {
assert_eq!(extract_at_paths("@a.txt first"), vec!["a.txt"]);
assert_eq!(extract_at_paths("line1\n@b.txt"), vec!["b.txt"]);
}
#[test]
fn dedupes_preserving_order() {
assert_eq!(extract_at_paths("@a.rs @b.rs @a.rs"), vec!["a.rs", "b.rs"]);
}
#[test]
fn caps_at_max_references() {
let text = (1..=10).map(|i| format!("@f{}.rs", i)).collect::<Vec<_>>().join(" ");
assert_eq!(extract_at_paths(&text).len(), super::MAX_AT_REFERENCES);
}
#[test]
fn ignores_email_like_tokens() {
assert_eq!(extract_at_paths("联系 user@example.com 谢谢"), Vec::<String>::new());
}
#[test]
fn ignores_bare_at() {
assert_eq!(extract_at_paths("@ 后面是空格"), Vec::<String>::new());
assert_eq!(extract_at_paths("no refs here"), Vec::<String>::new());
}
}