This commit is contained in:
Bryan Helmkamp 2026-02-23 12:20:48 -05:00
parent dddc6df8c8
commit 3a70e6711c
8 changed files with 357 additions and 5 deletions

View file

@ -17,7 +17,7 @@ name = "attractor"
path = "src/main.rs"
[features]
default = []
default = ["server"]
server = ["axum", "tower", "tokio-stream"]
[dependencies]

View file

@ -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<String>,
verbose: u8,
styles: &'static Styles,
}
impl AgentBackend {
#[must_use]
pub const fn new(model: String, provider: Option<String>) -> Self {
Self { model, provider }
pub const fn new(
model: String,
provider: Option<String>,
verbose: u8,
styles: &'static Styles,
) -> Self {
Self {
model,
provider,
verbose,
styles,
}
}
fn build_profile(&self) -> Arc<dyn ProviderProfile> {
@ -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::<Vec<_>>()
.join(", ")
}

View file

@ -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<String>,
/// Override default LLM provider
#[arg(long)]
pub provider: Option<String>,
/// Execute with simulated LLM backend
#[arg(long)]
pub dry_run: bool,
}
/// Read a .dot file from disk.
///
/// # Errors

View file

@ -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,
)))
}
});

View file

@ -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<dyn Interviewer>| {
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(())
}

View file

@ -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 {

View file

@ -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]

View file

@ -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<dyn Interviewer>| {
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)
// ===========================================================================