mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-09-07 08:27:12 +00:00
serve
This commit is contained in:
parent
dddc6df8c8
commit
3a70e6711c
8 changed files with 357 additions and 5 deletions
|
|
@ -17,7 +17,7 @@ name = "attractor"
|
|||
path = "src/main.rs"
|
||||
|
||||
[features]
|
||||
default = []
|
||||
default = ["server"]
|
||||
server = ["axum", "tower", "tokio-stream"]
|
||||
|
||||
[dependencies]
|
||||
|
|
|
|||
|
|
@ -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(", ")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)))
|
||||
}
|
||||
});
|
||||
|
|
|
|||
91
crates/attractor/src/cli/serve.rs
Normal file
91
crates/attractor/src/cli/serve.rs
Normal 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(())
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
// ===========================================================================
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue