diff --git a/crates/attractor/Cargo.toml b/crates/attractor/Cargo.toml index cb49a0fe4..543b6c067 100644 --- a/crates/attractor/Cargo.toml +++ b/crates/attractor/Cargo.toml @@ -17,7 +17,7 @@ name = "attractor" path = "src/main.rs" [features] -default = [] +default = ["server"] server = ["axum", "tower", "tokio-stream"] [dependencies] diff --git a/crates/attractor/src/cli/backend.rs b/crates/attractor/src/cli/backend.rs index 8ff96dc1a..66812c610 100644 --- a/crates/attractor/src/cli/backend.rs +++ b/crates/attractor/src/cli/backend.rs @@ -4,10 +4,11 @@ use std::sync::Arc; use async_trait::async_trait; use agent::{ - AnthropicProfile, GeminiProfile, LocalExecutionEnvironment, OpenAiProfile, ProviderProfile, - Session, SessionConfig, Turn, + AnthropicProfile, EventData, EventKind, GeminiProfile, LocalExecutionEnvironment, + OpenAiProfile, ProviderProfile, Session, SessionConfig, Turn, }; use llm::client::Client; +use terminal::Styles; use crate::context::Context; use crate::error::AttractorError; @@ -18,12 +19,24 @@ use crate::handler::codergen::{CodergenBackend, CodergenResult}; pub struct AgentBackend { model: String, provider: Option, + verbose: u8, + styles: &'static Styles, } impl AgentBackend { #[must_use] - pub const fn new(model: String, provider: Option) -> Self { - Self { model, provider } + pub const fn new( + model: String, + provider: Option, + verbose: u8, + styles: &'static Styles, + ) -> Self { + Self { + model, + provider, + verbose, + styles, + } } fn build_profile(&self) -> Arc { @@ -59,11 +72,96 @@ impl CodergenBackend for AgentBackend { }; let mut session = Session::new(client, profile, exec_env, config); + + // Subscribe to session events for real-time tool status on stderr. + let verbose = self.verbose; + if verbose >= 1 { + let node_id = node.id.clone(); + let styles = self.styles; + let mut rx = session.subscribe(); + tokio::spawn(async move { + while let Ok(event) = rx.recv().await { + match (&event.kind, &event.data) { + ( + EventKind::ToolCallStart, + EventData::ToolCall { + tool_name, + arguments, + .. + }, + ) => { + eprintln!( + "{dim}[{node_id}]{reset} {dim}\u{25cf}{reset} {bold}{cyan}{tool_name}{reset}{dim}({args}){reset}", + dim = styles.dim, + reset = styles.reset, + bold = styles.bold, + cyan = styles.cyan, + args = format_tool_args(arguments), + ); + } + ( + EventKind::ToolCallEnd, + EventData::ToolCallEnd { + tool_name, + output, + is_error, + .. + }, + ) if verbose >= 2 => { + let label = if *is_error { "error" } else { "result" }; + eprintln!( + "{dim}[{node_id}] [{label}] {tool_name}:{reset}\n{}", + serde_json::to_string_pretty(output) + .unwrap_or_else(|_| output.to_string()), + dim = styles.dim, + reset = styles.reset, + ); + } + (EventKind::Error, EventData::Error { error }) => { + eprintln!( + "{dim}[{node_id}]{reset} {red}\u{2717} {error}{reset}", + dim = styles.dim, + red = styles.red, + reset = styles.reset, + ); + } + _ => {} + } + } + }); + } + session.initialize().await; session.process_input(prompt).await.map_err(|e| { AttractorError::Handler(format!("Agent session failed: {e}")) })?; + // Print session summary to stderr. + if self.verbose >= 1 { + let (mut turn_count, mut tool_call_count, mut total_tokens) = (0usize, 0usize, 0i64); + for turn in session.history().turns() { + if let Turn::Assistant { + tool_calls, usage, .. + } = turn + { + turn_count += 1; + tool_call_count += tool_calls.len(); + total_tokens += usage.total_tokens; + } + } + let token_str = if total_tokens >= 1000 { + format!("{}k tokens", total_tokens / 1000) + } else { + format!("{total_tokens} tokens") + }; + eprintln!( + "{dim}[{node_id}] Done ({turn_count} turns, {tool_call_count} tool calls, {token_str}){reset}", + node_id = node.id, + dim = self.styles.dim, + reset = self.styles.reset, + ); + } + // Extract last assistant response from the session history. let response = session .history() @@ -83,3 +181,23 @@ impl CodergenBackend for AgentBackend { Ok(CodergenResult::Text(response)) } } + +fn format_tool_args(args: &serde_json::Value) -> String { + let Some(obj) = args.as_object() else { + return args.to_string(); + }; + obj.iter() + .map(|(k, v)| match v { + serde_json::Value::String(s) => { + let display = if s.len() > 80 { + format!("{}...", &s[..77]) + } else { + s.clone() + }; + format!("{k}={display:?}") + } + other => format!("{k}={other}"), + }) + .collect::>() + .join(", ") +} diff --git a/crates/attractor/src/cli/mod.rs b/crates/attractor/src/cli/mod.rs index a55b86046..2349071a3 100644 --- a/crates/attractor/src/cli/mod.rs +++ b/crates/attractor/src/cli/mod.rs @@ -1,5 +1,7 @@ pub mod backend; pub mod run; +#[cfg(feature = "server")] +pub mod serve; pub mod validate; use std::path::Path; @@ -24,6 +26,9 @@ pub enum Command { Run(RunArgs), /// Parse and validate a pipeline without executing Validate(ValidateArgs), + /// Start the HTTP API server + #[cfg(feature = "server")] + Serve(ServeArgs), } #[derive(Args)] @@ -66,6 +71,30 @@ pub struct ValidateArgs { pub pipeline: PathBuf, } +#[cfg(feature = "server")] +#[derive(Args)] +pub struct ServeArgs { + /// Port to listen on + #[arg(long, default_value = "3000")] + pub port: u16, + + /// Host address to bind to + #[arg(long, default_value = "127.0.0.1")] + pub host: String, + + /// Override default LLM model + #[arg(long)] + pub model: Option, + + /// Override default LLM provider + #[arg(long)] + pub provider: Option, + + /// Execute with simulated LLM backend + #[arg(long)] + pub dry_run: bool, +} + /// Read a .dot file from disk. /// /// # Errors diff --git a/crates/attractor/src/cli/run.rs b/crates/attractor/src/cli/run.rs index 66664e52c..f37fba9c2 100644 --- a/crates/attractor/src/cli/run.rs +++ b/crates/attractor/src/cli/run.rs @@ -131,6 +131,8 @@ pub async fn run_command(args: RunArgs, styles: &'static Styles) -> anyhow::Resu Some(Box::new(AgentBackend::new( model.clone(), provider.clone(), + args.verbose, + styles, ))) } }); diff --git a/crates/attractor/src/cli/serve.rs b/crates/attractor/src/cli/serve.rs new file mode 100644 index 000000000..e21cdf47f --- /dev/null +++ b/crates/attractor/src/cli/serve.rs @@ -0,0 +1,91 @@ +use std::sync::Arc; + +use terminal::Styles; +use tokio::net::TcpListener; + +use crate::cli::backend::AgentBackend; +use crate::handler::default_registry; +use crate::interviewer::Interviewer; +use crate::server::{build_router, create_app_state}; + +use super::ServeArgs; + +/// Start the HTTP API server. +/// +/// # Errors +/// +/// Returns an error if the server fails to bind or encounters a fatal error. +pub async fn serve_command(args: ServeArgs, styles: &'static Styles) -> anyhow::Result<()> { + // Resolve dry-run mode (same pattern as run.rs) + let dry_run_mode = if args.dry_run { + true + } else { + match llm::client::Client::from_env().await { + Ok(c) if c.provider_names().is_empty() => { + eprintln!( + "{yellow}Warning:{reset} No LLM providers configured. Running in dry-run mode.", + yellow = styles.yellow, reset = styles.reset, + ); + true + } + Ok(_) => false, + Err(e) => { + eprintln!( + "{yellow}Warning:{reset} Failed to initialize LLM client: {e}. Running in dry-run mode.", + yellow = styles.yellow, reset = styles.reset, + ); + true + } + } + }; + + // Resolve model/provider defaults + let provider = args.provider; + let model = args.model.unwrap_or_else(|| match provider.as_deref() { + Some("openai") => "gpt-5.2".to_string(), + Some("gemini") => "gemini-3-pro-preview".to_string(), + _ => "claude-sonnet-4-5".to_string(), + }); + + // Build registry factory + let factory = move |interviewer: Arc| { + let model = model.clone(); + let provider = provider.clone(); + default_registry(interviewer, move || { + if dry_run_mode { + None + } else { + Some(Box::new(AgentBackend::new( + model.clone(), + provider.clone(), + 0, + styles, + ))) + } + }) + }; + + let state = create_app_state(factory); + let router = build_router(state); + + let addr = format!("{}:{}", args.host, args.port); + let listener = TcpListener::bind(&addr).await?; + + eprintln!( + "{bold}Attractor server listening on {green}{addr}{reset}", + bold = styles.bold, + green = styles.green, + reset = styles.reset, + ); + if dry_run_mode { + eprintln!( + "{dim}(dry-run mode){reset}", + dim = styles.dim, + reset = styles.reset, + ); + } + + axum::serve(listener, router).await?; + + Ok(()) +} diff --git a/crates/attractor/src/main.rs b/crates/attractor/src/main.rs index 7f648dce9..41c192071 100644 --- a/crates/attractor/src/main.rs +++ b/crates/attractor/src/main.rs @@ -13,6 +13,10 @@ async fn main() { attractor::cli::Command::Validate(args) => { attractor::cli::validate::validate_command(&args, styles) } + #[cfg(feature = "server")] + attractor::cli::Command::Serve(args) => { + attractor::cli::serve::serve_command(args, styles).await + } }; if let Err(e) = result { diff --git a/crates/attractor/tests/cli.rs b/crates/attractor/tests/cli.rs index a6144d4e9..4a964d686 100644 --- a/crates/attractor/tests/cli.rs +++ b/crates/attractor/tests/cli.rs @@ -61,6 +61,21 @@ fn validate_invalid() { .failure(); } +// -- serve ------------------------------------------------------------------- + +#[test] +fn serve_help() { + attractor() + .args(["serve", "--help"]) + .assert() + .success() + .stdout(predicate::str::contains("--port")) + .stdout(predicate::str::contains("--host")) + .stdout(predicate::str::contains("--dry-run")) + .stdout(predicate::str::contains("--model")) + .stdout(predicate::str::contains("--provider")); +} + // -- run --dry-run ----------------------------------------------------------- #[test] diff --git a/crates/attractor/tests/integration.rs b/crates/attractor/tests/integration.rs index ae5257cf4..5361e8758 100644 --- a/crates/attractor/tests/integration.rs +++ b/crates/attractor/tests/integration.rs @@ -3229,6 +3229,99 @@ mod sse_events { } } +// =========================================================================== +// 18b. Serve command: dry-run registry factory builds a working router +// =========================================================================== + +#[cfg(feature = "server")] +mod serve_dry_run { + use std::sync::Arc; + use std::time::Duration; + + use attractor::handler::default_registry; + use attractor::interviewer::Interviewer; + use attractor::server::{build_router, create_app_state}; + use axum::body::Body; + use axum::http::{Request, StatusCode}; + use tower::ServiceExt; + + const MINIMAL_DOT: &str = r#"digraph Test { + graph [goal="Test"] + start [shape=Mdiamond] + exit [shape=Msquare] + start -> exit + }"#; + + /// Build the router exactly as `serve_command` does in dry-run mode. + fn dry_run_app() -> axum::Router { + let factory = |interviewer: Arc| { + default_registry(interviewer, || None) + }; + let state = create_app_state(factory); + build_router(state) + } + + async fn body_json(body: Body) -> serde_json::Value { + let bytes = axum::body::to_bytes(body, usize::MAX).await.unwrap(); + serde_json::from_slice(&bytes).unwrap() + } + + #[tokio::test] + async fn dry_run_serve_starts_and_runs_pipeline() { + let app = dry_run_app(); + + // POST /pipelines to start a pipeline + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": MINIMAL_DOT})).unwrap(), + )) + .unwrap(); + + let response = app.clone().oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::CREATED); + + let body = body_json(response.into_body()).await; + let pipeline_id = body["id"].as_str().unwrap().to_string(); + assert!(!pipeline_id.is_empty()); + + // Wait for pipeline to complete + tokio::time::sleep(Duration::from_millis(500)).await; + + // GET /pipelines/{id} to verify completion + let req = Request::builder() + .method("GET") + .uri(format!("/pipelines/{pipeline_id}")) + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + + let body = body_json(response.into_body()).await; + assert_eq!(body["status"].as_str().unwrap(), "completed"); + } + + #[tokio::test] + async fn dry_run_serve_rejects_invalid_dot() { + let app = dry_run_app(); + + let req = Request::builder() + .method("POST") + .uri("/pipelines") + .header("content-type", "application/json") + .body(Body::from( + serde_json::to_string(&serde_json::json!({"dot_source": "not valid dot"})).unwrap(), + )) + .unwrap(); + + let response = app.oneshot(req).await.unwrap(); + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + } +} + // =========================================================================== // 19a. Sub-pipeline E2E (TS Scenario 9) // ===========================================================================