1use super::{idle_notification, initialize_response, running_notification};
2use crate::notifications::{SessionPreviewParams, SessionPreviewResponse};
3use agent_client_protocol::schema::v2::{
4 AgentCapabilities, CancelSessionNotification, CloseSessionRequest, CloseSessionResponse, CompactionStatus,
5 CompactionUpdate, ContentChunk, CreateElicitationRequest, ElicitationFormMode, ElicitationSchema,
6 ElicitationSessionScope, Implementation, InitializeRequest, InitializeResponse, ListSessionsRequest,
7 ListSessionsResponse, LoginAuthRequest, LoginAuthResponse, NewSessionRequest, NewSessionResponse, PromptRequest,
8 PromptResponse, ResumeSessionRequest, ResumeSessionResponse, SessionId, SessionInfo, SessionUpdate,
9 SetSessionConfigOptionRequest, SetSessionConfigOptionResponse, StopReason, UpdateSessionNotification,
10};
11use agent_client_protocol::util::MatchDispatchFrom;
12use agent_client_protocol::{
13 self as acp, Agent, Client, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, NullRun, Responder, V2Builder,
14};
15use tokio::sync::mpsc;
16
17pub struct FakeAgent {
18 initialize: InitializeResponse,
19 new_session: Option<NewSessionResponse>,
20 sessions: Option<Vec<SessionInfo>>,
21 previews: Vec<SessionPreviewResponse>,
22 login_method: Option<String>,
23 hold_config: bool,
24 hold_list_sessions: bool,
25 replay: Vec<UpdateSessionNotification>,
26 live: Vec<UpdateSessionNotification>,
27 turn: Option<Turn>,
28 capture: Option<Capture>,
29}
30
31pub struct FakeAgentRequests {
32 pub connection: mpsc::UnboundedReceiver<ConnectionTo<Client>>,
33 pub initialize: mpsc::UnboundedReceiver<InitializeRequest>,
34 pub new_session: mpsc::UnboundedReceiver<NewSessionRequest>,
35 pub login: mpsc::UnboundedReceiver<LoginAuthRequest>,
36 pub config: mpsc::UnboundedReceiver<SetSessionConfigOptionRequest>,
37 pub prompt: mpsc::UnboundedReceiver<(PromptRequest, Responder<PromptResponse>)>,
38 pub resume: mpsc::UnboundedReceiver<(ResumeSessionRequest, Responder<ResumeSessionResponse>)>,
39 pub cancel: mpsc::UnboundedReceiver<CancelSessionNotification>,
40 pub pending_config: mpsc::UnboundedReceiver<Responder<SetSessionConfigOptionResponse>>,
41 pub list_sessions: mpsc::UnboundedReceiver<(ListSessionsRequest, Responder<ListSessionsResponse>)>,
42 pub close_session: mpsc::UnboundedReceiver<CloseSessionRequest>,
43}
44
45impl Default for FakeAgent {
46 fn default() -> Self {
47 Self {
48 initialize: initialize_response(),
49 new_session: None,
50 sessions: None,
51 previews: Vec::new(),
52 login_method: None,
53 hold_config: false,
54 hold_list_sessions: false,
55 replay: Vec::new(),
56 live: Vec::new(),
57 turn: None,
58 capture: None,
59 }
60 }
61}
62
63impl FakeAgent {
64 pub fn remote_server(mut self, info: &crate::notifications::RemoteServerInfo) -> Self {
65 self.initialize.meta = Some(info.to_meta());
66 self
67 }
68
69 pub fn session_preview(mut self, preview: SessionPreviewResponse) -> Self {
70 self.previews.push(preview);
71 self
72 }
73
74 pub fn agent_info(mut self, info: Implementation) -> Self {
75 self.initialize.info = info;
76 self
77 }
78 pub fn capabilities(mut self, capabilities: AgentCapabilities) -> Self {
79 self.initialize.capabilities = capabilities;
80 self
81 }
82 pub fn login_method(mut self, method: &str) -> Self {
83 self.login_method = Some(method.into());
84 self
85 }
86 pub fn hold_config(mut self, hold: bool) -> Self {
87 self.hold_config = hold;
88 self
89 }
90 pub fn hold_list_sessions(mut self, hold: bool) -> Self {
91 self.hold_list_sessions = hold;
92 self
93 }
94 pub fn new_session_response(mut self, response: NewSessionResponse) -> Self {
95 self.new_session = Some(response);
96 self
97 }
98 pub fn sessions(mut self, sessions: Vec<SessionInfo>) -> Self {
99 self.sessions = Some(sessions);
100 self
101 }
102 pub fn replay_message(mut self, session_id: &str, text: &str) -> Self {
103 self.replay.push(message(session_id, text));
104 self
105 }
106 pub fn live_message(mut self, session_id: &str, text: &str) -> Self {
107 self.live.push(message(session_id, text));
108 self
109 }
110 pub fn prompt_reply(mut self, text: &str) -> Self {
111 self.turn = Some(Turn::Reply(text.into()));
112 self
113 }
114
115 pub fn prompt_elicitation(mut self, message: &str) -> Self {
116 self.turn = Some(Turn::Elicit(message.into()));
117 self
118 }
119 pub fn compaction(mut self, session_id: &str, compaction_id: &str, status: CompactionStatus) -> Self {
120 self.replay.push(UpdateSessionNotification::new(
121 session_id,
122 SessionUpdate::CompactionUpdate(CompactionUpdate::new(compaction_id, status)),
123 ));
124 self
125 }
126
127 pub fn capture(mut self) -> (Self, FakeAgentRequests) {
128 let (connection, connection_rx) = mpsc::unbounded_channel();
129 let (initialize, initialize_rx) = mpsc::unbounded_channel();
130 let (new_session, new_session_rx) = mpsc::unbounded_channel();
131 let (login, login_rx) = mpsc::unbounded_channel();
132 let (config, config_rx) = mpsc::unbounded_channel();
133 let (prompt, prompt_rx) = mpsc::unbounded_channel();
134 let (resume, resume_rx) = mpsc::unbounded_channel();
135 let (cancel, cancel_rx) = mpsc::unbounded_channel();
136 let (pending_config, pending_config_rx) = mpsc::unbounded_channel();
137 let (list_sessions, list_sessions_rx) = mpsc::unbounded_channel();
138 let (close_session, close_session_rx) = mpsc::unbounded_channel();
139 self.capture = Some(Capture {
140 connection,
141 initialize,
142 new_session,
143 login,
144 config,
145 prompt,
146 resume,
147 cancel,
148 pending_config,
149 list_sessions,
150 close_session,
151 });
152 (
153 self,
154 FakeAgentRequests {
155 connection: connection_rx,
156 initialize: initialize_rx,
157 new_session: new_session_rx,
158 login: login_rx,
159 config: config_rx,
160 prompt: prompt_rx,
161 resume: resume_rx,
162 cancel: cancel_rx,
163 pending_config: pending_config_rx,
164 list_sessions: list_sessions_rx,
165 close_session: close_session_rx,
166 },
167 )
168 }
169
170 pub fn agent(self) -> V2Builder<Agent, impl HandleDispatchFrom<Client>, NullRun> {
171 Agent.v2().name("fake-agent").with_handler(self)
172 }
173
174 pub async fn build(self) -> Result<crate::client::AcpClient, crate::client::AcpClientError> {
175 let (agent, client) = super::Channel::duplex();
176 tokio::task::spawn_local(self.agent().connect_to(agent));
177 crate::client::connect_acp_client(client, super::initialize_request()).await
178 }
179}
180
181impl HandleDispatchFrom<Client> for FakeAgent {
182 async fn handle_dispatch_from(
183 &mut self,
184 message: Dispatch,
185 cx: ConnectionTo<Client>,
186 ) -> Result<Handled<Dispatch>, acp::Error> {
187 MatchDispatchFrom::new(message, &cx)
188 .if_request(async |request: InitializeRequest, responder| {
189 if let Some(capture) = &self.capture {
190 let _ = capture.connection.send(cx.clone());
191 let _ = capture.initialize.send(request);
192 }
193 responder.respond(self.initialize.clone())
194 })
195 .await
196 .if_request(async |request: NewSessionRequest, responder| {
197 let Some(response) = &self.new_session else {
198 return Ok(Handled::No { message: (request, responder), retry: false });
199 };
200 if let Some(capture) = &self.capture {
201 let _ = capture.new_session.send(request);
202 }
203 responder.respond(response.clone())?;
204 Ok(Handled::Yes)
205 })
206 .await
207 .if_request(async |request: SessionPreviewParams, responder| {
208 let Some(preview) = self.previews.iter().find(|preview| preview.session_id == request.session_id)
209 else {
210 return Ok(Handled::No { message: (request, responder), retry: false });
211 };
212 responder.respond(preview.clone())?;
213 Ok(Handled::Yes)
214 })
215 .await
216 .if_request(async |request: ListSessionsRequest, responder| {
217 if self.hold_list_sessions
218 && let Some(capture) = &self.capture
219 {
220 let _ = capture.list_sessions.send((request, responder));
221 return Ok(Handled::Yes);
222 }
223 let Some(sessions) = &self.sessions else {
224 return Ok(Handled::No { message: (request, responder), retry: false });
225 };
226 responder.respond(ListSessionsResponse::new(sessions.clone()))?;
227 Ok(Handled::Yes)
228 })
229 .await
230 .if_request(async |request: CloseSessionRequest, responder| {
231 if let Some(capture) = &self.capture {
232 let _ = capture.close_session.send(request);
233 }
234 responder.respond(CloseSessionResponse::new())
235 })
236 .await
237 .if_request(async |request: LoginAuthRequest, responder| {
238 let allowed = self.login_method.as_deref() == Some(request.method_id.0.as_ref());
239 if let Some(capture) = &self.capture {
240 let _ = capture.login.send(request);
241 }
242 if allowed {
243 responder.respond(LoginAuthResponse::new())
244 } else {
245 responder.respond_with_error(acp::Error::invalid_params())
246 }
247 })
248 .await
249 .if_request(async |request: SetSessionConfigOptionRequest, responder| {
250 if let Some(capture) = &self.capture {
251 let _ = capture.config.send(request);
252 if self.hold_config {
253 let _ = capture.pending_config.send(responder);
254 return Ok(());
255 }
256 }
257 responder.respond(SetSessionConfigOptionResponse::new(vec![]))
258 })
259 .await
260 .if_request(async |request: PromptRequest, responder| {
261 if let Some(capture) = &self.capture {
262 let _ = capture.prompt.send((request, responder));
263 return Ok(());
264 }
265 responder.respond(PromptResponse::new("user-message"))?;
266 self.run_turn(request.session_id, &cx)
267 })
268 .await
269 .if_request(async |request: ResumeSessionRequest, responder| self.resume(request, responder, &cx))
270 .await
271 .if_notification(async |notification: CancelSessionNotification| {
272 if let Some(capture) = &self.capture {
273 let _ = capture.cancel.send(notification);
274 }
275 Ok(())
276 })
277 .await
278 .done()
279 }
280
281 fn describe_chain(&self) -> impl std::fmt::Debug {
282 "FakeAgent"
283 }
284}
285
286impl FakeAgent {
287 fn resume(
288 &self,
289 request: ResumeSessionRequest,
290 responder: Responder<ResumeSessionResponse>,
291 cx: &ConnectionTo<Client>,
292 ) -> Result<(), acp::Error> {
293 if let Some(capture) = &self.capture {
294 let _ = capture.resume.send((request, responder));
295 return Ok(());
296 }
297 if request.replay_from.is_none() {
298 return responder.respond(ResumeSessionResponse::new());
299 }
300 for notification in &self.replay {
301 cx.send_notification(notification.clone())?;
302 }
303 cx.send_notification(idle_notification(request.session_id, None))?;
304 responder.respond(ResumeSessionResponse::new())?;
305 for notification in &self.live {
306 cx.send_notification(notification.clone())?;
307 }
308 Ok(())
309 }
310
311 fn run_turn(&self, session_id: SessionId, cx: &ConnectionTo<Client>) -> Result<(), acp::Error> {
312 let Some(turn) = &self.turn else {
313 return Ok(());
314 };
315 cx.send_notification(running_notification(session_id.clone()))?;
316 match turn {
317 Turn::Reply(text) => {
318 cx.send_notification(message(&session_id.0, text))?;
319 cx.send_notification(idle_notification(session_id, Some(StopReason::EndTurn)))
320 }
321 Turn::Elicit(prompt) => {
322 let request = CreateElicitationRequest::new(
323 ElicitationFormMode::new(
324 ElicitationSessionScope::new(session_id.clone()),
325 ElicitationSchema::new(),
326 ),
327 prompt.clone(),
328 );
329 let connection = cx.clone();
330 cx.spawn(async move {
331 let echo = match connection.send_request(request).block_task().await {
332 Ok(response) => serde_json::to_string(&response),
333 Err(error) => serde_json::to_string(&error),
334 }
335 .map_err(acp::Error::into_internal_error)?;
336 connection.send_notification(message(&session_id.0, &echo))?;
337 connection.send_notification(idle_notification(session_id, Some(StopReason::EndTurn)))
338 })
339 }
340 }
341 }
342}
343
344enum Turn {
345 Reply(String),
346 Elicit(String),
347}
348
349struct Capture {
350 connection: mpsc::UnboundedSender<ConnectionTo<Client>>,
351 initialize: mpsc::UnboundedSender<InitializeRequest>,
352 new_session: mpsc::UnboundedSender<NewSessionRequest>,
353 login: mpsc::UnboundedSender<LoginAuthRequest>,
354 config: mpsc::UnboundedSender<SetSessionConfigOptionRequest>,
355 prompt: mpsc::UnboundedSender<(PromptRequest, Responder<PromptResponse>)>,
356 resume: mpsc::UnboundedSender<(ResumeSessionRequest, Responder<ResumeSessionResponse>)>,
357 cancel: mpsc::UnboundedSender<CancelSessionNotification>,
358 pending_config: mpsc::UnboundedSender<Responder<SetSessionConfigOptionResponse>>,
359 list_sessions: mpsc::UnboundedSender<(ListSessionsRequest, Responder<ListSessionsResponse>)>,
360 close_session: mpsc::UnboundedSender<CloseSessionRequest>,
361}
362
363fn message(session_id: &str, text: &str) -> UpdateSessionNotification {
364 UpdateSessionNotification::new(
365 SessionId::new(session_id),
366 SessionUpdate::AgentMessageChunk(ContentChunk::new(text.into(), "message")),
367 )
368}