Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

56 changes: 36 additions & 20 deletions crates/infinity-agent-cli/src/acp_server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ struct SessionConnection {
/// ACP session ID (what the client sees).
acp_id: String,
/// Daemon session ID (what the daemon sees). None until Connected arrives.
daemon_id: Option<String>,
daemon_id: Option<infinity_protocol::ThreadRef>,
/// ACP connection for sending notifications.
cx: ConnectionTo<Client>,
/// Pending prompt responder.
Expand All @@ -60,7 +60,7 @@ impl SessionConnection {
/// Global state shared across all ACP request handlers.
struct GlobalState {
/// Sessions known from the daemon Welcome/SessionsUpdated.
sessions: HashMap<String, SessionInfo>,
sessions: HashMap<infinity_protocol::ThreadRef, SessionInfo>,
/// Maps daemon session ID → ACP session ID.
daemon_to_acp: HashMap<String, String>,
/// Maps ACP session ID → daemon session ID.
Expand Down Expand Up @@ -165,7 +165,9 @@ pub async fn run() -> Result<(), BoxError> {
info.status != infinity_protocol::SessionStatus::Archived
})
.map(|(id, info)| {
let acp_id = s.daemon_to_acp.get(id).unwrap_or(id).clone();
let rendered = id.to_string();
let acp_id =
s.daemon_to_acp.get(&rendered).unwrap_or(&rendered).clone();
agent_client_protocol::schema::SessionInfo::new(
acp_id,
std::env::current_dir().unwrap_or_default(),
Expand Down Expand Up @@ -202,19 +204,23 @@ pub async fn run() -> Result<(), BoxError> {
cx: ConnectionTo<Client>| {
let mut s = state.lock().expect("bug: lock poisoned");
let acp_id = req.session_id.0.to_string();
let daemon_id = s
let daemon_id: infinity_protocol::ThreadRef = s
.acp_to_daemon
.get(&acp_id)
.cloned()
.unwrap_or_else(|| acp_id.clone());
.unwrap_or_else(|| acp_id.clone())
.as_str()
.into();
if !s.sessions.contains_key(&daemon_id) {
return responder.respond_with_error(
agent_client_protocol::schema::Error::invalid_params()
.data(format!("session '{}' not found", acp_id)),
);
}
s.daemon_to_acp.insert(daemon_id.clone(), acp_id.clone());
s.acp_to_daemon.insert(acp_id.clone(), daemon_id.clone());
s.daemon_to_acp
.insert(daemon_id.to_string(), acp_id.clone());
s.acp_to_daemon
.insert(acp_id.clone(), daemon_id.to_string());
save_session_mappings(&s);

let (cmd_tx, cmd_rx) = mpsc::unbounded_channel();
Expand Down Expand Up @@ -309,11 +315,13 @@ pub async fn run() -> Result<(), BoxError> {
}

// Unknown session — try connecting by daemon ID.
let daemon_id = s
let daemon_id: infinity_protocol::ThreadRef = s
.acp_to_daemon
.get(&acp_id)
.cloned()
.unwrap_or_else(|| acp_id.clone());
.unwrap_or_else(|| acp_id.clone())
.as_str()
.into();
let (cmd_tx, cmd_rx) = mpsc::unbounded_channel();
s.session_connections.insert(acp_id.clone(), cmd_tx);
drop(s);
Expand Down Expand Up @@ -384,7 +392,8 @@ pub async fn run() -> Result<(), BoxError> {
s.sessions = sessions;
// Notify all active sessions of info updates.
for (daemon_id, info) in &s.sessions {
let acp_id = s.daemon_to_acp.get(daemon_id).unwrap_or(daemon_id);
let rendered = daemon_id.to_string();
let acp_id = s.daemon_to_acp.get(&rendered).unwrap_or(&rendered);
if let Some(ref cx) = s.cx {
let notif = SessionNotification::new(
acp_id.clone(),
Expand Down Expand Up @@ -416,7 +425,7 @@ pub async fn run() -> Result<(), BoxError> {
/// How a session connection starts.
enum SessionStart {
Load {
daemon_id: String,
daemon_id: infinity_protocol::ThreadRef,
responder: Responder<LoadSessionResponse>,
},
Create {
Expand All @@ -425,7 +434,7 @@ enum SessionStart {
responder: Responder<PromptResponse>,
},
ConnectAndPrompt {
daemon_id: String,
daemon_id: infinity_protocol::ThreadRef,
text: String,
responder: Responder<PromptResponse>,
},
Expand Down Expand Up @@ -462,7 +471,7 @@ async fn run_session_connection(
responder,
} => {
let msg = ClientMessage::Connect {
session_id: daemon_id.clone(),
root_thread_id: daemon_id.clone(),
thread_id: None,
keeps_session_alive: true,
};
Expand Down Expand Up @@ -518,18 +527,25 @@ async fn run_session_connection(
let msg = recv(&mut framed).await?;
match msg {
DaemonMessage::Connected {
session_id, title, ..
root_thread_id: session_id,
title,
..
} => {
conn.daemon_id = Some(session_id.clone());
// Register mapping.
{
let mut g = global.lock().expect("bug: lock poisoned");
g.daemon_to_acp.insert(session_id.clone(), acp_id.clone());
g.acp_to_daemon.insert(acp_id.clone(), session_id.clone());
g.daemon_to_acp
.insert(session_id.to_string(), acp_id.clone());
g.acp_to_daemon
.insert(acp_id.clone(), session_id.to_string());
save_session_mappings(&g);
}

let input_msg = ClientMessage::UserInput { session_id, text };
let input_msg = ClientMessage::UserInput {
thread_id: session_id,
text,
};
send(&mut framed, &input_msg).await?;
conn.pending_prompt = Some(responder);

Expand Down Expand Up @@ -559,7 +575,7 @@ async fn run_session_connection(
responder,
} => {
let msg = ClientMessage::Connect {
session_id: daemon_id.clone(),
root_thread_id: daemon_id.clone(),
thread_id: None,
keeps_session_alive: true,
};
Expand All @@ -571,7 +587,7 @@ async fn run_session_connection(
match msg {
DaemonMessage::Connected { .. } => {
let input_msg = ClientMessage::UserInput {
session_id: daemon_id,
thread_id: daemon_id,
text,
};
send(&mut framed, &input_msg).await?;
Expand Down Expand Up @@ -615,7 +631,7 @@ async fn run_session_connection(
}
if let Some(ref daemon_id) = conn.daemon_id {
let msg = ClientMessage::UserInput {
session_id: daemon_id.clone(),
thread_id: daemon_id.clone(),
text,
};
if let Err(e) = send(&mut framed, &msg).await {
Expand Down
54 changes: 29 additions & 25 deletions crates/infinity-agent-cli/src/daemon_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,8 @@ use std::collections::HashMap;
use bytes::Bytes;
use futures_util::{SinkExt, StreamExt};
use infinity_protocol::{
ClientMessage, DaemonMessage, ModelRef, SessionInfo, TokenUsage, length_delimited_codec,
ClientMessage, DaemonMessage, ModelRef, SessionInfo, ThreadRef, TokenUsage,
length_delimited_codec,
};
use std::path::PathBuf;
use tokio::io::AsyncBufReadExt;
Expand Down Expand Up @@ -35,7 +36,7 @@ fn convert_token_usage(u: &TokenUsage) -> infinity_provider_protocol::Usage {

/// Convert a DaemonMessage into a (thread_id, DisplayEvent) tuple.
/// Returns None for messages that are handled separately (Connected, Welcome, etc.).
fn daemon_msg_to_display(msg: DaemonMessage) -> Option<(Option<String>, DisplayEvent)> {
fn daemon_msg_to_display(msg: DaemonMessage) -> Option<(Option<ThreadRef>, DisplayEvent)> {
Some(match msg {
DaemonMessage::StartOutput { thread_id } => (thread_id, DisplayEvent::StartOutput),
DaemonMessage::TextChunk { thread_id, chunk } => {
Expand Down Expand Up @@ -372,15 +373,15 @@ pub async fn run_headless(message: String) -> Result<(), BoxError> {
// Wait for Connected to get the session ID.
let session_id = loop {
match recv(&mut framed).await? {
DaemonMessage::Connected { session_id, .. } => break session_id,
DaemonMessage::Connected { root_thread_id, .. } => break root_thread_id,
DaemonMessage::Error { text, .. } => return Err(text.into()),
_ => continue, // skip SessionsUpdated, etc.
}
};

// Send the user's message.
let msg = ClientMessage::UserInput {
session_id: session_id.clone(),
thread_id: session_id.clone(),
text: message,
};
framed.send(Bytes::from(serde_json::to_vec(&msg)?)).await?;
Expand Down Expand Up @@ -540,15 +541,15 @@ where
_ => return Err("expected Welcome from daemon".into()),
};

let (display_tx, display_rx) = mpsc::unbounded_channel::<(Option<String>, DisplayEvent)>();
let (display_tx, display_rx) = mpsc::unbounded_channel::<(Option<ThreadRef>, DisplayEvent)>();
let (input_tx, input_rx) = mpsc::unbounded_channel::<String>();
let (load_session_tx, load_session_rx) = mpsc::unbounded_channel::<(Option<String>, bool)>();
let (load_session_tx, load_session_rx) = mpsc::unbounded_channel::<(Option<ThreadRef>, bool)>();
let (model_switch_tx, model_switch_rx) = mpsc::unbounded_channel::<usize>();
let (session_tx, session_rx) = mpsc::unbounded_channel::<SessionChanged>();
let (model_switched_tx, model_switched_rx) =
mpsc::unbounded_channel::<crate::terminal::ModelSwitched>();
let (sessions_updated_tx, sessions_updated_rx) =
mpsc::unbounded_channel::<HashMap<String, SessionInfo>>();
mpsc::unbounded_channel::<HashMap<ThreadRef, SessionInfo>>();
let (soft_detach_tx, soft_detach_rx) = mpsc::unbounded_channel::<()>();
let (detach_result_tx, detach_result_rx) = mpsc::unbounded_channel::<DetachResult>();
let (choice_answered_tx, choice_answered_rx) = mpsc::unbounded_channel::<(String, usize)>();
Expand All @@ -559,27 +560,28 @@ where

// If --session was provided, connect to it immediately (supports prefix matching).
if let Some(ref session_id) = session {
let matches: Vec<&String> = sessions
let matches: Vec<(&ThreadRef, String)> = sessions
.keys()
.filter(|k| k.starts_with(session_id.as_str()))
.map(|k| (k, k.to_string()))
.filter(|(_, rendered)| rendered.starts_with(session_id.as_str()))
.collect();
let resolved = match matches.len() {
0 => return Err(format!("no session found matching prefix '{session_id}'").into()),
1 => matches[0].clone(),
1 => matches[0].0.clone(),
_ => {
return Err(format!(
"ambiguous session prefix '{session_id}' — matches: {}",
matches
.iter()
.map(|s| s.as_str())
.map(|(_, rendered)| rendered.as_str())
.collect::<Vec<_>>()
.join(", ")
)
.into());
}
};
to_daemon.send(ClientMessage::Connect {
session_id: resolved,
root_thread_id: resolved,
thread_id: None,
keeps_session_alive: true,
})?;
Expand Down Expand Up @@ -608,7 +610,7 @@ where
choice_answered_tx,
));

let mut active_session: Option<String> = None;
let mut active_session: Option<ThreadRef> = None;
let mut pending_input: Vec<String> = Vec::new();
// The model most recently selected via the model picker. Stored locally so
// it can be passed when creating new sessions, even if no session is active.
Expand Down Expand Up @@ -637,12 +639,12 @@ where
break;
};
match msg {
DaemonMessage::Connected { session_id, title, total_tokens_used, model_name, context_window, provider_id, .. } => {
active_session = Some(session_id.clone());
let _ = session_tx.send(SessionChanged { session_id, title, total_tokens_used, model_name, context_window, provider_id });
DaemonMessage::Connected { root_thread_id, title, total_tokens_used, model_name, context_window, provider_id, .. } => {
active_session = Some(root_thread_id.clone());
let _ = session_tx.send(SessionChanged { session_id: root_thread_id, title, total_tokens_used, model_name, context_window, provider_id });
for text in pending_input.drain(..) {
let sid = active_session.as_ref().expect("bug: active_session should be set after Connected").clone();
let _ = to_daemon.send(ClientMessage::UserInput { session_id: sid, text });
let _ = to_daemon.send(ClientMessage::UserInput { thread_id: sid, text });
}
}
DaemonMessage::ModelSwitched { thread_id, model_name, context_window, provider_id } => {
Expand Down Expand Up @@ -699,15 +701,15 @@ where
let Some(()) = msg else { break };
if let Some(ref sid) = active_session {
pending_soft_detach = true;
let _ = to_daemon.send(ClientMessage::SoftDetach { session_id: sid.clone() });
let _ = to_daemon.send(ClientMessage::SoftDetach { root_thread_id: sid.clone() });
}
}

maybe_target = load_session_rx.recv() => {
let Some((maybe_target, shut_down_old)) = maybe_target else { break };
if let Some(ref sid) = active_session {
if shut_down_old {
let _ = to_daemon.send(ClientMessage::ShutdownSession { session_id: sid.clone() });
let _ = to_daemon.send(ClientMessage::ShutdownSession { root_thread_id: sid.clone() });
} else {
let _ = to_daemon.send(ClientMessage::Disconnect);
}
Expand All @@ -721,7 +723,7 @@ where
}

if let Some(target) = maybe_target {
let _ = to_daemon.send(ClientMessage::Connect { session_id: target, thread_id: None, keeps_session_alive: true });
let _ = to_daemon.send(ClientMessage::Connect { root_thread_id: target, thread_id: None, keeps_session_alive: true });
} // if none, will be created on next user input
}

Expand All @@ -736,7 +738,7 @@ where
selected_model = Some(model.clone());
if let Some(sid) = &active_session {
let _ = to_daemon.send(ClientMessage::SwitchModel {
session_id: sid.clone(), model,
thread_id: sid.clone(), model,
});
}
}
Expand All @@ -756,12 +758,12 @@ where
let Some(text) = text else { break };
if let Some(ref sid) = active_session {
if text == "__compact__" {
let _ = to_daemon.send(ClientMessage::TriggerCompaction { session_id: sid.clone() });
let _ = to_daemon.send(ClientMessage::TriggerCompaction { root_thread_id: sid.clone() });
} else if text == "__archive__" {
let _ = to_daemon.send(ClientMessage::ArchiveSession { session_id: sid.clone() });
let _ = to_daemon.send(ClientMessage::ArchiveSession { root_thread_id: sid.clone() });
active_session = None;
} else {
let _ = to_daemon.send(ClientMessage::UserInput { session_id: sid.clone(), text });
let _ = to_daemon.send(ClientMessage::UserInput { thread_id: sid.clone(), text });
}
} else {
pending_input.push(text);
Expand Down Expand Up @@ -789,7 +791,9 @@ where
if keep_running {
let _ = to_daemon.send(ClientMessage::Disconnect);
} else {
let _ = to_daemon.send(ClientMessage::ShutdownSession { session_id: sid });
let _ = to_daemon.send(ClientMessage::ShutdownSession {
root_thread_id: sid,
});
}
}

Expand Down
4 changes: 2 additions & 2 deletions crates/infinity-agent-cli/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -250,11 +250,11 @@ async fn async_main(cli: Cli) -> Result<(), BoxError> {
);
}
};
if remotes.iter().any(|r| r.name == name) {
if remotes.iter().any(|r| r.name.as_str() == name) {
return Err(format!("remote '{name}' already exists").into());
}
remotes.push(infinity_daemon::remote::RemoteConfig {
name: name.clone(),
name: name.clone().into(),
ssh_args: ssh_args.to_vec(),
});
if let Some(parent) = path.parent() {
Expand Down
Loading
Loading