1use super::error::AcpClientError;
2use super::event::AcpEvent;
3use crate::notifications::{
4 AuthMethodsUpdatedParams, ContextClearedParams, GitDiffEventPayload, McpNotification, SubAgentProgressParams,
5};
6use agent_client_protocol::schema::ProtocolVersion;
7use agent_client_protocol::schema::v2::{
8 AuthMethod, CancelSessionNotification, ClientCapabilities, CreateElicitationRequest, ElicitationCapabilities,
9 ElicitationFormCapabilities, ElicitationUrlCapabilities, Implementation, InitializeRequest, InitializeResponse,
10 NewSessionRequest, NewSessionResponse, PermissionOptionId, PermissionOptionKind, PromptCapabilities, PromptRequest,
11 PromptResponse, RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse,
12 ResumeSessionRequest, ResumeSessionResponse, SelectedPermissionOutcome, SessionCapabilities,
13 UpdateSessionNotification,
14};
15use agent_client_protocol::util::MatchDispatchFrom;
16use agent_client_protocol::{
17 self as acp, Client, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, V2ConnectionTo,
18};
19use futures::future::{AbortHandle, abortable};
20use std::future::Future;
21use std::sync::Arc;
22use tokio::sync::{mpsc, oneshot};
23use tokio_util::sync::CancellationToken;
24use tracing::info;
25
26#[derive(Clone)]
27pub struct AcpClientHandle {
28 cx: V2ConnectionTo<acp::Agent>,
29 connection: Arc<ClientConnection>,
30}
31
32pub struct AcpClient {
33 pub initialize_response: InitializeResponse,
34 pub event_rx: mpsc::UnboundedReceiver<AcpEvent>,
35 pub handle: AcpClientHandle,
36}
37
38pub async fn connect_acp_client(
40 agent: impl ConnectTo<Client> + 'static,
41 init_request: InitializeRequest,
42) -> Result<AcpClient, AcpClientError> {
43 let (event_tx, event_rx) = mpsc::unbounded_channel();
44 let (init_tx, init_rx) = oneshot::channel();
45 let closed = CancellationToken::new();
46 let events = ConnectionEvents { event_tx, closed: closed.clone() };
47 let (driver, abort) = abortable(run_client_connection(agent, init_request, init_tx, events));
48 spawn(async move {
49 let _ = driver.await;
50 });
51 let connection = Arc::new(ClientConnection { abort, closed });
52 let (initialize_response, cx) = await_response(init_rx).await?;
53 Ok(AcpClient { initialize_response, event_rx, handle: AcpClientHandle { cx, connection } })
54}
55
56pub fn initialize_request(info: Implementation) -> InitializeRequest {
57 InitializeRequest::new(ProtocolVersion::V2, info).capabilities(ClientCapabilities::new().elicitation(
58 ElicitationCapabilities::new().form(ElicitationFormCapabilities::new()).url(ElicitationUrlCapabilities::new()),
59 ))
60}
61
62impl AcpClient {
63 pub fn agent_name(&self) -> String {
65 let info = &self.initialize_response.info;
66 info.title.as_deref().unwrap_or(&info.name).to_string()
67 }
68
69 pub fn prompt_capabilities(&self) -> Option<&PromptCapabilities> {
70 self.session_capabilities().and_then(|session| session.prompt.as_ref())
71 }
72
73 pub fn session_capabilities(&self) -> Option<&SessionCapabilities> {
74 self.initialize_response.capabilities.session.as_ref()
75 }
76
77 pub fn auth_methods(&self) -> &[AuthMethod] {
78 &self.initialize_response.auth_methods
79 }
80}
81
82impl AcpClientHandle {
83 pub async fn disconnect(&self) {
85 self.connection.abort.abort();
86 self.connection.closed.cancelled().await;
87 }
88
89 pub fn prompt(
90 &self,
91 request: PromptRequest,
92 ) -> impl Future<Output = Result<PromptResponse, AcpClientError>> + Send + use<> {
93 self.request(request)
94 }
95
96 pub fn new_session(
97 &self,
98 request: NewSessionRequest,
99 ) -> impl Future<Output = Result<NewSessionResponse, AcpClientError>> + Send + use<> {
100 self.request(request)
101 }
102
103 pub fn request<R: acp::JsonRpcRequest>(
105 &self,
106 request: R,
107 ) -> impl Future<Output = Result<R::Response, AcpClientError>> + Send + use<R> {
108 let sent = self.cx.send_request(request);
109 async move { sent.block_task().await.map_err(AcpClientError::Protocol) }
110 }
111
112 pub fn resume_session(
113 &self,
114 request: ResumeSessionRequest,
115 ) -> impl Future<Output = Result<ResumeSessionResponse, AcpClientError>> + Send + use<> {
116 self.request(request)
117 }
118
119 pub fn cancel(&self, request: CancelSessionNotification) -> Result<(), AcpClientError> {
120 self.notify(request)
121 }
122
123 pub fn notify(&self, request: impl acp::JsonRpcNotification) -> Result<(), AcpClientError> {
124 self.cx.send_notification(request).map_err(AcpClientError::Protocol)
125 }
126}
127
128struct ClientConnection {
129 abort: AbortHandle,
130 closed: CancellationToken,
131}
132
133impl Drop for ClientConnection {
134 fn drop(&mut self) {
135 self.abort.abort();
136 }
137}
138
139struct ConnectionEvents {
140 event_tx: mpsc::UnboundedSender<AcpEvent>,
141 closed: CancellationToken,
142}
143
144impl Drop for ConnectionEvents {
145 fn drop(&mut self) {
146 let _ = self.event_tx.send(AcpEvent::ConnectionClosed);
147 self.closed.cancel();
148 }
149}
150
151async fn run_client_connection(
152 agent: impl ConnectTo<Client> + 'static,
153 init_request: InitializeRequest,
154 init_tx: oneshot::Sender<Result<(InitializeResponse, V2ConnectionTo<acp::Agent>), AcpClientError>>,
155 events: ConnectionEvents,
156) {
157 let connection_result = Client
158 .v2()
159 .name(&init_request.info.name)
160 .with_handler(ClientHandlers(events.event_tx.clone()))
161 .connect_with(agent, async move |cx: V2ConnectionTo<acp::Agent>| {
162 let result = cx.send_request(init_request).block_task().await.map_err(AcpClientError::Protocol);
163 let _ = init_tx.send(result.map(|response| {
164 info!("ACP initialized: protocol={:?}, agent_info={:?}", response.protocol_version, response.info);
165 (response, cx.clone())
166 }));
167 cx.incoming_closed().await;
168 Ok(())
169 })
170 .await;
171 if let Err(error) = connection_result {
172 tracing::warn!("ACP connection exited with error: {error:?}");
173 }
174}
175
176struct ClientHandlers(mpsc::UnboundedSender<AcpEvent>);
177
178impl HandleDispatchFrom<acp::Agent> for ClientHandlers {
179 async fn handle_dispatch_from(
180 &mut self,
181 message: Dispatch,
182 cx: ConnectionTo<acp::Agent>,
183 ) -> Result<Handled<Dispatch>, acp::Error> {
184 let emit = |event| {
185 if let Err(error) = self.0.send(event)
186 && let AcpEvent::ElicitationRequest { responder, .. } = error.0
187 {
188 let _ = responder.respond_with_error(acp::Error::internal_error());
189 }
190 Ok::<_, acp::Error>(())
191 };
192 MatchDispatchFrom::new(message, &cx)
193 .if_request(async |request: RequestPermissionRequest, responder| {
194 let outcome = auto_approve_option(&request).map_or(RequestPermissionOutcome::Cancelled, |option| {
195 RequestPermissionOutcome::Selected(SelectedPermissionOutcome::new(option))
196 });
197 let _ = responder.respond(RequestPermissionResponse::new(outcome));
198 Ok(())
199 })
200 .await
201 .if_request(async |params: CreateElicitationRequest, responder| {
202 emit(AcpEvent::ElicitationRequest { params: Box::new(params), responder })
203 })
204 .await
205 .if_notification(async |params: UpdateSessionNotification| emit(params.into()))
206 .await
207 .if_notification(async |params: ContextClearedParams| emit(AcpEvent::ContextCleared(params)))
208 .await
209 .if_notification(async |params: SubAgentProgressParams| emit(AcpEvent::SubAgentProgress(params)))
210 .await
211 .if_notification(async |params: GitDiffEventPayload| emit(AcpEvent::GitDiffEvent(params)))
212 .await
213 .if_notification(async |params: McpNotification| emit(AcpEvent::McpNotification(params)))
214 .await
215 .if_notification(async |params: AuthMethodsUpdatedParams| emit(AcpEvent::AuthMethodsUpdated(params)))
216 .await
217 .done()
218 }
219
220 fn describe_chain(&self) -> impl std::fmt::Debug {
221 "ClientHandlers"
222 }
223}
224
225#[cfg(not(target_family = "wasm"))]
226fn spawn(future: impl Future<Output = ()> + Send + 'static) {
227 tokio::spawn(future);
228}
229
230#[cfg(target_family = "wasm")]
231fn spawn(future: impl Future<Output = ()> + 'static) {
232 wasm_bindgen_futures::spawn_local(future);
233}
234
235async fn await_response<T>(receiver: oneshot::Receiver<Result<T, AcpClientError>>) -> Result<T, AcpClientError> {
236 receiver.await.map_err(|_| AcpClientError::AgentCrashed("ACP task ended before responding".to_string()))?
237}
238
239fn auto_approve_option(req: &RequestPermissionRequest) -> Option<PermissionOptionId> {
240 req.options
241 .iter()
242 .find(|option| matches!(option.kind, PermissionOptionKind::AllowOnce | PermissionOptionKind::AllowAlways))
243 .or_else(|| req.options.first())
244 .map(|option| option.option_id.clone())
245}