Skip to main content

acp_utils/testing/
fake_agent.rs

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, UserMessage,
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                self.run_turn(request, responder, &cx)
266            })
267            .await
268            .if_request(async |request: ResumeSessionRequest, responder| self.resume(request, responder, &cx))
269            .await
270            .if_notification(async |notification: CancelSessionNotification| {
271                if let Some(capture) = &self.capture {
272                    let _ = capture.cancel.send(notification);
273                }
274                Ok(())
275            })
276            .await
277            .done()
278    }
279
280    fn describe_chain(&self) -> impl std::fmt::Debug {
281        "FakeAgent"
282    }
283}
284
285impl FakeAgent {
286    fn resume(
287        &self,
288        request: ResumeSessionRequest,
289        responder: Responder<ResumeSessionResponse>,
290        cx: &ConnectionTo<Client>,
291    ) -> Result<(), acp::Error> {
292        if let Some(capture) = &self.capture {
293            let _ = capture.resume.send((request, responder));
294            return Ok(());
295        }
296        if request.replay_from.is_none() {
297            return responder.respond(ResumeSessionResponse::new());
298        }
299        for notification in &self.replay {
300            cx.send_notification(notification.clone())?;
301        }
302        cx.send_notification(idle_notification(request.session_id, None))?;
303        responder.respond(ResumeSessionResponse::new())?;
304        for notification in &self.live {
305            cx.send_notification(notification.clone())?;
306        }
307        Ok(())
308    }
309
310    fn run_turn(
311        &self,
312        request: PromptRequest,
313        responder: Responder<PromptResponse>,
314        cx: &ConnectionTo<Client>,
315    ) -> Result<(), acp::Error> {
316        const USER_MESSAGE_ID: &str = "user-message";
317        let session_id = request.session_id;
318        let Some(turn) = &self.turn else {
319            return responder.respond(PromptResponse::new(USER_MESSAGE_ID));
320        };
321        cx.send_notification(running_notification(session_id.clone()))?;
322        let user_message = UserMessage::new(USER_MESSAGE_ID).content(request.prompt);
323        cx.send_notification(UpdateSessionNotification::new(
324            session_id.clone(),
325            SessionUpdate::UserMessage(user_message),
326        ))?;
327        responder.respond(PromptResponse::new(USER_MESSAGE_ID))?;
328        match turn {
329            Turn::Reply(text) => {
330                cx.send_notification(message(&session_id.0, text))?;
331                cx.send_notification(idle_notification(session_id, Some(StopReason::EndTurn)))
332            }
333            Turn::Elicit(prompt) => {
334                let request = CreateElicitationRequest::new(
335                    ElicitationFormMode::new(
336                        ElicitationSessionScope::new(session_id.clone()),
337                        ElicitationSchema::new(),
338                    ),
339                    prompt.clone(),
340                );
341                let connection = cx.clone();
342                cx.spawn(async move {
343                    let echo = match connection.send_request(request).block_task().await {
344                        Ok(response) => serde_json::to_string(&response),
345                        Err(error) => serde_json::to_string(&error),
346                    }
347                    .map_err(acp::Error::into_internal_error)?;
348                    connection.send_notification(message(&session_id.0, &echo))?;
349                    connection.send_notification(idle_notification(session_id, Some(StopReason::EndTurn)))
350                })
351            }
352        }
353    }
354}
355
356enum Turn {
357    Reply(String),
358    Elicit(String),
359}
360
361struct Capture {
362    connection: mpsc::UnboundedSender<ConnectionTo<Client>>,
363    initialize: mpsc::UnboundedSender<InitializeRequest>,
364    new_session: mpsc::UnboundedSender<NewSessionRequest>,
365    login: mpsc::UnboundedSender<LoginAuthRequest>,
366    config: mpsc::UnboundedSender<SetSessionConfigOptionRequest>,
367    prompt: mpsc::UnboundedSender<(PromptRequest, Responder<PromptResponse>)>,
368    resume: mpsc::UnboundedSender<(ResumeSessionRequest, Responder<ResumeSessionResponse>)>,
369    cancel: mpsc::UnboundedSender<CancelSessionNotification>,
370    pending_config: mpsc::UnboundedSender<Responder<SetSessionConfigOptionResponse>>,
371    list_sessions: mpsc::UnboundedSender<(ListSessionsRequest, Responder<ListSessionsResponse>)>,
372    close_session: mpsc::UnboundedSender<CloseSessionRequest>,
373}
374
375fn message(session_id: &str, text: &str) -> UpdateSessionNotification {
376    UpdateSessionNotification::new(
377        SessionId::new(session_id),
378        SessionUpdate::AgentMessageChunk(ContentChunk::new(text.into(), "message")),
379    )
380}