agent_client_protocol_conductor/
lib.rs1use std::path::PathBuf;
67use std::str::FromStr;
68
69mod conductor;
71mod debug_logger;
73mod snoop;
74pub mod trace;
76
77pub use self::conductor::*;
78
79use clap::{Parser, Subcommand};
80
81#[cfg(feature = "unstable_protocol_v2")]
82use agent_client_protocol::schema::v2;
83use agent_client_protocol::{AcpAgent, Stdio};
84use agent_client_protocol::{Client, Conductor, DynConnectTo, schema::v1::InitializeRequest};
85use tracing::Instrument;
86use tracing_subscriber::{EnvFilter, layer::SubscriberExt, util::SubscriberInitExt};
87
88#[derive(Debug)]
95pub struct CommandLineComponents(pub Vec<AcpAgent>);
96
97impl InstantiateProxies for CommandLineComponents {
98 fn instantiate_proxies(
99 self: Box<Self>,
100 req: InitializeRequest,
101 ) -> futures::future::BoxFuture<
102 'static,
103 Result<(InitializeRequest, Vec<DynConnectTo<Conductor>>), agent_client_protocol::Error>,
104 > {
105 Box::pin(async move {
106 let proxies = self.0.into_iter().map(DynConnectTo::new).collect();
107 Ok((req, proxies))
108 })
109 }
110
111 #[cfg(feature = "unstable_protocol_v2")]
112 fn instantiate_v2_proxies(
113 self: Box<Self>,
114 req: v2::InitializeRequest,
115 ) -> futures::future::BoxFuture<
116 'static,
117 Result<(v2::InitializeRequest, Vec<DynConnectTo<Conductor>>), agent_client_protocol::Error>,
118 > {
119 Box::pin(async move {
120 let proxies = self.0.into_iter().map(DynConnectTo::new).collect();
121 Ok((req, proxies))
122 })
123 }
124}
125
126impl InstantiateProxiesAndAgent for CommandLineComponents {
127 fn instantiate_proxies_and_agent(
128 self: Box<Self>,
129 req: InitializeRequest,
130 ) -> futures::future::BoxFuture<
131 'static,
132 Result<
133 (
134 InitializeRequest,
135 Vec<DynConnectTo<Conductor>>,
136 DynConnectTo<Client>,
137 ),
138 agent_client_protocol::Error,
139 >,
140 > {
141 Box::pin(async move {
142 let mut iter = self.0.into_iter().peekable();
143 let mut proxies: Vec<DynConnectTo<Conductor>> = Vec::new();
144
145 while let Some(component) = iter.next() {
147 if iter.peek().is_some() {
148 proxies.push(DynConnectTo::new(component));
149 } else {
150 let agent = DynConnectTo::new(component);
152 return Ok((req, proxies, agent));
153 }
154 }
155
156 Err(agent_client_protocol::util::internal_error(
157 "no agent component in list",
158 ))
159 })
160 }
161
162 #[cfg(feature = "unstable_protocol_v2")]
163 fn instantiate_v2_proxies_and_agent(
164 self: Box<Self>,
165 req: v2::InitializeRequest,
166 ) -> futures::future::BoxFuture<
167 'static,
168 Result<
169 (
170 v2::InitializeRequest,
171 Vec<DynConnectTo<Conductor>>,
172 DynConnectTo<Client>,
173 ),
174 agent_client_protocol::Error,
175 >,
176 > {
177 Box::pin(async move {
178 let mut iter = self.0.into_iter().peekable();
179 let mut proxies = Vec::new();
180
181 while let Some(component) = iter.next() {
182 if iter.peek().is_some() {
183 proxies.push(DynConnectTo::new(component));
184 } else {
185 return Ok((req, proxies, DynConnectTo::new(component)));
186 }
187 }
188
189 Err(agent_client_protocol::util::internal_error(
190 "no agent component in list",
191 ))
192 })
193 }
194}
195
196struct TraceHandleWriter(agent_client_protocol_trace_viewer::TraceHandle);
198
199impl trace::WriteEvent for TraceHandleWriter {
200 fn write_event(&mut self, event: &trace::TraceEvent) -> std::io::Result<()> {
201 let value = serde_json::to_value(event).map_err(std::io::Error::other)?;
202 self.0.push(value);
203 Ok(())
204 }
205}
206
207#[derive(Parser, Debug)]
208#[command(author, version, about, long_about = None)]
209pub struct ConductorArgs {
210 #[arg(long)]
212 pub debug: bool,
213
214 #[arg(long)]
216 pub debug_dir: Option<PathBuf>,
217
218 #[arg(long)]
221 pub log: Option<String>,
222
223 #[arg(long)]
226 pub trace: Option<PathBuf>,
227
228 #[arg(long)]
231 pub serve: bool,
232
233 #[command(subcommand)]
234 pub command: ConductorCommand,
235}
236
237#[derive(Subcommand, Debug)]
238pub enum ConductorCommand {
239 Agent {
241 #[arg(short, long, default_value = "conductor")]
243 name: String,
244
245 components: Vec<String>,
247 },
248
249 Proxy {
251 #[arg(short, long, default_value = "conductor")]
253 name: String,
254
255 proxies: Vec<String>,
257 },
258}
259
260impl ConductorArgs {
261 pub async fn main(self) -> anyhow::Result<()> {
263 let pid = std::process::id();
264 let cwd = std::env::current_dir()
265 .map_or_else(|_| "<unknown>".to_string(), |p| p.display().to_string());
266
267 let debug_logger = if self.debug {
269 let components = match &self.command {
271 ConductorCommand::Agent { components, .. } => components.clone(),
272 ConductorCommand::Proxy { proxies, .. } => proxies.clone(),
273 };
274
275 Some(
277 debug_logger::DebugLogger::new(self.debug_dir.clone(), &components)
278 .await
279 .map_err(|e| anyhow::anyhow!("Failed to create debug logger: {e}"))?,
280 )
281 } else {
282 None
283 };
284
285 if let Some(debug_logger) = &debug_logger {
286 let log_level = self.log.as_deref().unwrap_or("info");
288
289 let tracing_writer = debug_logger.create_tracing_writer();
291 tracing_subscriber::registry()
292 .with(EnvFilter::new(log_level))
293 .with(
294 tracing_subscriber::fmt::layer()
295 .with_target(true)
296 .with_writer(move || tracing_writer.clone()),
297 )
298 .init();
299
300 tracing::info!(pid = %pid, cwd = %cwd, level = %log_level, "Conductor starting with debug logging");
301 }
302
303 let (trace_writer, _viewer_server) = match (&self.trace, self.serve) {
305 (Some(trace_path), false) => {
307 let writer = trace::TraceWriter::from_path(trace_path)
308 .map_err(|e| anyhow::anyhow!("Failed to create trace writer: {e}"))?;
309 (Some(writer), None)
310 }
311 (None, true) => {
313 let (handle, server) = agent_client_protocol_trace_viewer::serve_memory(
314 agent_client_protocol_trace_viewer::TraceViewerConfig::default(),
315 )?;
316 let writer = trace::TraceWriter::new(TraceHandleWriter(handle));
317 (Some(writer), Some(tokio::spawn(server)))
318 }
319 (Some(trace_path), true) => {
321 let writer = trace::TraceWriter::from_path(trace_path)
322 .map_err(|e| anyhow::anyhow!("Failed to create trace writer: {e}"))?;
323 let server = agent_client_protocol_trace_viewer::serve_file(
324 trace_path.clone(),
325 agent_client_protocol_trace_viewer::TraceViewerConfig::default(),
326 );
327 (Some(writer), Some(tokio::spawn(server)))
328 }
329 (None, false) => (None, None),
331 };
332
333 self.run(debug_logger.as_ref(), trace_writer)
334 .instrument(tracing::info_span!("conductor", pid = %pid, cwd = %cwd))
335 .await
336 .map_err(|err| anyhow::anyhow!("{err}"))
337 }
338
339 async fn run(
340 self,
341 debug_logger: Option<&debug_logger::DebugLogger>,
342 trace_writer: Option<trace::TraceWriter>,
343 ) -> Result<(), agent_client_protocol::Error> {
344 match self.command {
345 ConductorCommand::Agent { name, components } => {
346 initialize_conductor(
347 debug_logger,
348 trace_writer,
349 name,
350 components,
351 ConductorImpl::new_agent,
352 )
353 .await
354 }
355 ConductorCommand::Proxy { name, proxies } => {
356 initialize_conductor(
357 debug_logger,
358 trace_writer,
359 name,
360 proxies,
361 ConductorImpl::new_proxy,
362 )
363 .await
364 }
365 }
366 }
367}
368
369async fn initialize_conductor<Host: ConductorHostRole>(
370 debug_logger: Option<&debug_logger::DebugLogger>,
371 trace_writer: Option<trace::TraceWriter>,
372 name: String,
373 components: Vec<String>,
374 new_conductor: impl FnOnce(String, CommandLineComponents) -> ConductorImpl<Host>,
375) -> Result<(), agent_client_protocol::Error> {
376 let providers: Vec<AcpAgent> = components
378 .into_iter()
379 .enumerate()
380 .map(|(i, s)| {
381 let mut agent = AcpAgent::from_str(&s)?;
382 if let Some(logger) = debug_logger {
383 agent = agent.with_debug(logger.create_callback(i.to_string()));
384 }
385 Ok(agent)
386 })
387 .collect::<Result<Vec<_>, agent_client_protocol::Error>>()?;
388
389 let stdio = if let Some(logger) = debug_logger {
391 Stdio::new().with_debug(logger.create_callback("C".to_string()))
392 } else {
393 Stdio::new()
394 };
395
396 let mut conductor = new_conductor(name, CommandLineComponents(providers));
398 if let Some(writer) = trace_writer {
399 conductor = conductor.with_trace_writer(writer);
400 }
401
402 conductor.run(stdio).await
403}