diff --git a/Cargo.lock b/Cargo.lock index 7619ec293..a1d0981de 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -167,6 +167,25 @@ dependencies = [ "uuid", ] +[[package]] +name = "arc-cli" +version = "0.1.0" +dependencies = [ + "anyhow", + "arc-agent", + "arc-attractor", + "arc-llm", + "arc-util", + "assert_cmd", + "clap", + "dotenvy", + "httpmock", + "predicates", + "serde_json", + "tempfile", + "tokio", +] + [[package]] name = "arc-git-storage" version = "0.1.0" @@ -182,7 +201,6 @@ name = "arc-llm" version = "0.1.0" dependencies = [ "anyhow", - "assert_cmd", "async-trait", "base64", "bytes", @@ -190,13 +208,10 @@ dependencies = [ "dotenvy", "futures", "http", - "httpmock", - "predicates", "rand 0.8.5", "reqwest 0.12.28", "serde", "serde_json", - "tempfile", "thiserror 2.0.18", "tokio", "tokio-stream", diff --git a/crates/arc-agent/Cargo.toml b/crates/arc-agent/Cargo.toml index cd9277f08..6eba64abc 100644 --- a/crates/arc-agent/Cargo.toml +++ b/crates/arc-agent/Cargo.toml @@ -16,11 +16,6 @@ docker = ["bollard", "tar"] [lib] doctest = false -[[bin]] -name = "arc-agent" -path = "src/main.rs" -test = false - [dependencies] clap.workspace = true anyhow.workspace = true diff --git a/crates/arc-agent/src/cli.rs b/crates/arc-agent/src/cli.rs index 434bdf07c..0ed1f3172 100644 --- a/crates/arc-agent/src/cli.rs +++ b/crates/arc-agent/src/cli.rs @@ -3,7 +3,7 @@ use crate::{ ProviderProfile, Session, SessionConfig, ToolApprovalFn, Turn, subagent::{SessionFactory, SubAgentManager}, }; -use clap::{Parser, ValueEnum}; +use clap::{Args, Parser, ValueEnum}; use arc_llm::client::Client; use arc_llm::provider::{ModelId, Provider}; use std::io::{IsTerminal, Write}; @@ -11,54 +11,60 @@ use std::path::PathBuf; use std::sync::{Arc, Mutex}; use arc_util::terminal::Styles; -/// Minimal CLI for the agent agentic loop. -#[derive(Parser)] -#[command(name = "arc-agent")] -struct Cli { +/// Public arguments for the agent command, usable from an external CLI. +#[derive(Args)] +pub struct AgentArgs { /// Task prompt - prompt: String, + pub prompt: String, /// LLM provider (anthropic, openai, gemini, kimi, zai, minimax) #[arg(long, default_value = "anthropic")] - provider: String, + pub provider: String, /// Model name (defaults per provider) #[arg(long)] - model: Option, + pub model: Option, /// Permission level for tool execution #[arg(long, default_value = "read-write", value_enum)] - permissions: PermissionLevel, + pub permissions: PermissionLevel, /// Skip interactive prompts; deny tools outside permission level #[arg(long)] - auto_approve: bool, + pub auto_approve: bool, /// Print LLM request/response debug info to stderr #[arg(long)] - debug: bool, + pub debug: bool, /// Print full LLM request/response JSON to stderr #[arg(long)] - verbose: bool, + pub verbose: bool, /// Directory containing skill files (overrides default discovery) #[arg(long)] - skills_dir: Option, + pub skills_dir: Option, /// Output format (text for human-readable, json for NDJSON event stream) #[arg(long, default_value = "text", value_enum)] - output_format: OutputFormat, + pub output_format: OutputFormat, +} + +#[derive(Parser)] +#[command(name = "arc-agent")] +struct Cli { + #[command(flatten)] + args: AgentArgs, } #[derive(Clone, Copy, Debug, ValueEnum)] -enum OutputFormat { +pub enum OutputFormat { Text, Json, } #[derive(Clone, Copy, Debug, ValueEnum)] -enum PermissionLevel { +pub enum PermissionLevel { ReadOnly, ReadWrite, Full, @@ -330,15 +336,12 @@ impl arc_llm::middleware::Middleware for VerboseMiddleware { } } -pub async fn run() -> anyhow::Result<()> { - let _ = dotenvy::dotenv(); - let cli = Cli::parse(); - +pub async fn run_with_args(args: AgentArgs) -> anyhow::Result<()> { // Resolve color support once, leak to get 'static lifetime for use across threads let styles: &'static Styles = Box::leak(Box::new(Styles::detect_stderr())); // Parse provider string to enum early for compile-time safety - let provider: Provider = cli.provider.parse().map_err(|e: String| anyhow::anyhow!("{e}"))?; + let provider: Provider = args.provider.parse().map_err(|e: String| anyhow::anyhow!("{e}"))?; // Validate provider API key if !validate_api_key(provider) { @@ -350,14 +353,14 @@ pub async fn run() -> anyhow::Result<()> { .await .map_err(|e| anyhow::anyhow!("Failed to create LLM client: {e}"))?; - if cli.verbose { + if args.verbose { client.add_middleware(Arc::new(VerboseMiddleware { styles })); - } else if cli.debug { + } else if args.debug { client.add_middleware(Arc::new(DebugMiddleware { styles })); } // Resolve model and build profile - let model = cli + let model = args .model .as_deref() .unwrap_or_else(|| default_model(provider)); @@ -373,12 +376,12 @@ pub async fn run() -> anyhow::Result<()> { let env: Arc = Arc::new(LocalExecutionEnvironment::new(cwd)); // Build tool approval callback - let is_interactive = std::io::stdin().is_terminal() && !cli.auto_approve; - let tool_approval = build_tool_approval(cli.permissions, is_interactive, styles); + let is_interactive = std::io::stdin().is_terminal() && !args.auto_approve; + let tool_approval = build_tool_approval(args.permissions, is_interactive, styles); let config = SessionConfig { tool_approval: Some(tool_approval), - skill_dirs: cli.skills_dir.map(|d| vec![d]), + skill_dirs: args.skills_dir.map(|d| vec![d]), ..SessionConfig::default() }; @@ -432,8 +435,8 @@ pub async fn run() -> anyhow::Result<()> { }); // Subscribe to events - let verbose = cli.verbose; - let output_format = cli.output_format; + let verbose = args.verbose; + let output_format = args.output_format; let mut rx = session.subscribe(); tokio::spawn(async move { match output_format { @@ -524,7 +527,7 @@ pub async fn run() -> anyhow::Result<()> { // Initialize and run session.initialize().await; - let result = session.process_input(&cli.prompt).await; + let result = session.process_input(&args.prompt).await; if matches!(output_format, OutputFormat::Text) { // Print assistant text to stdout @@ -539,6 +542,12 @@ pub async fn run() -> anyhow::Result<()> { Ok(()) } +pub async fn run() -> anyhow::Result<()> { + let _ = dotenvy::dotenv(); + let cli = Cli::parse(); + run_with_args(cli.args).await +} + #[cfg(test)] mod tests { use super::*; diff --git a/crates/arc-agent/src/main.rs b/crates/arc-agent/src/main.rs deleted file mode 100644 index 8b1405a35..000000000 --- a/crates/arc-agent/src/main.rs +++ /dev/null @@ -1,12 +0,0 @@ -use std::process::ExitCode; - -#[tokio::main] -async fn main() -> ExitCode { - match arc_agent::cli::run().await { - Ok(()) => ExitCode::SUCCESS, - Err(e) => { - eprintln!("[error] {e}"); - ExitCode::FAILURE - } - } -} diff --git a/crates/arc-agent/tests/cli_integration.rs b/crates/arc-agent/tests/cli_integration.rs deleted file mode 100644 index 0daa54ca1..000000000 --- a/crates/arc-agent/tests/cli_integration.rs +++ /dev/null @@ -1,66 +0,0 @@ -use std::process::Command; - -#[test] -fn no_args_prints_usage() { - let output = Command::new(env!("CARGO_BIN_EXE_arc-agent")) - .env_clear() - .output() - .expect("failed to execute"); - - assert!(!output.status.success()); - let stderr = String::from_utf8_lossy(&output.stderr); - assert!( - stderr.contains("Usage:"), - "expected stderr to contain 'Usage:', got: {stderr}" - ); -} - -#[test] -fn help_flag_prints_help() { - let output = Command::new(env!("CARGO_BIN_EXE_arc-agent")) - .env_clear() - .arg("--help") - .output() - .expect("failed to execute"); - - assert!(output.status.success()); - let stdout = String::from_utf8_lossy(&output.stdout); - assert!( - stdout.contains("Task prompt"), - "expected stdout to contain 'Task prompt', got: {stdout}" - ); -} - -#[test] -fn missing_api_key_exits_with_error() { - let tmp = std::env::temp_dir(); - let output = Command::new(env!("CARGO_BIN_EXE_arc-agent")) - .env_clear() - .current_dir(&tmp) - .arg("test prompt") - .output() - .expect("failed to execute"); - - assert!(!output.status.success()); - let stderr = String::from_utf8_lossy(&output.stderr); - assert!( - stderr.contains("API key not set"), - "expected stderr to contain 'API key not set', got: {stderr}" - ); -} - -#[test] -fn invalid_permissions_value() { - let output = Command::new(env!("CARGO_BIN_EXE_arc-agent")) - .env_clear() - .args(["--permissions", "bogus", "test prompt"]) - .output() - .expect("failed to execute"); - - assert!(!output.status.success()); - let stderr = String::from_utf8_lossy(&output.stderr); - assert!( - stderr.contains("invalid value"), - "expected stderr to contain 'invalid value', got: {stderr}" - ); -} diff --git a/crates/arc-attractor/Cargo.toml b/crates/arc-attractor/Cargo.toml index 2988d63e1..11f985bd9 100644 --- a/crates/arc-attractor/Cargo.toml +++ b/crates/arc-attractor/Cargo.toml @@ -12,11 +12,6 @@ readme = "README.md" [lib] doctest = false -[[bin]] -name = "arc-attractor" -path = "src/main.rs" -test = false - [features] default = ["server"] server = ["axum", "tower", "tokio-stream"] diff --git a/crates/arc-attractor/src/main.rs b/crates/arc-attractor/src/main.rs deleted file mode 100644 index 62996662a..000000000 --- a/crates/arc-attractor/src/main.rs +++ /dev/null @@ -1,29 +0,0 @@ -use clap::Parser; -use arc_util::terminal::Styles; - -#[tokio::main] -async fn main() { - dotenvy::dotenv().ok(); - - let styles: &'static Styles = Box::leak(Box::new(Styles::detect_stderr())); - let cli = arc_attractor::cli::Cli::parse(); - - let result = match cli.command { - arc_attractor::cli::Command::Run(args) => arc_attractor::cli::run::run_command(args, styles).await, - arc_attractor::cli::Command::Validate(args) => { - arc_attractor::cli::validate::validate_command(&args, styles) - } - #[cfg(feature = "server")] - arc_attractor::cli::Command::Serve(args) => { - arc_attractor::cli::serve::serve_command(args, styles).await - } - }; - - if let Err(e) = result { - eprintln!( - "{red}Error:{reset} {e:#}", - red = styles.red, reset = styles.reset, - ); - std::process::exit(1); - } -} diff --git a/crates/arc-attractor/tests/cli.rs b/crates/arc-attractor/tests/cli.rs deleted file mode 100644 index 81162023e..000000000 --- a/crates/arc-attractor/tests/cli.rs +++ /dev/null @@ -1,190 +0,0 @@ -use assert_cmd::Command; -use predicates::prelude::*; - -#[allow(deprecated)] -fn attractor() -> Command { - Command::cargo_bin("arc-attractor").unwrap() -} - -// -- validate ---------------------------------------------------------------- - -#[test] -fn validate_simple() { - attractor() - .args(["validate", "../../test/simple.dot"]) - .assert() - .success() - .stderr(predicate::str::contains("Validation: OK")); -} - -#[test] -fn validate_branching() { - attractor() - .args(["validate", "../../test/branching.dot"]) - .assert() - .success() - .stderr(predicate::str::contains("Validation: OK")); -} - -#[test] -fn validate_conditions() { - attractor() - .args(["validate", "../../test/conditions.dot"]) - .assert() - .success() - .stderr(predicate::str::contains("Validation: OK")); -} - -#[test] -fn validate_parallel() { - attractor() - .args(["validate", "../../test/parallel.dot"]) - .assert() - .success() - .stderr(predicate::str::contains("Validation: OK")); -} - -#[test] -fn validate_styled() { - attractor() - .args(["validate", "../../test/styled.dot"]) - .assert() - .success() - .stderr(predicate::str::contains("Validation: OK")); -} - -#[test] -fn validate_legacy_tool() { - attractor() - .args(["validate", "../../test/legacy_tool.dot"]) - .assert() - .success() - .stderr(predicate::str::contains("Validation: OK")); -} - -#[test] -fn validate_invalid() { - attractor() - .args(["validate", "../../test/invalid.dot"]) - .assert() - .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] -fn dry_run_simple() { - attractor() - .args(["run", "--dry-run", "--auto-approve", "../../test/simple.dot"]) - .assert() - .success(); -} - -#[test] -fn dry_run_branching() { - attractor() - .args(["run", "--dry-run", "--auto-approve", "../../test/branching.dot"]) - .assert() - .success(); -} - -#[test] -fn dry_run_conditions() { - attractor() - .args(["run", "--dry-run", "--auto-approve", "../../test/conditions.dot"]) - .assert() - .success(); -} - -#[test] -fn dry_run_parallel() { - attractor() - .args(["run", "--dry-run", "--auto-approve", "../../test/parallel.dot"]) - .assert() - .success(); -} - -#[test] -fn dry_run_styled() { - attractor() - .args(["run", "--dry-run", "--auto-approve", "../../test/styled.dot"]) - .assert() - .success(); -} - -#[test] -fn dry_run_legacy_tool() { - attractor() - .args(["run", "--dry-run", "--auto-approve", "../../test/legacy_tool.dot"]) - .assert() - .success(); -} - -// -- NDJSON logging ---------------------------------------------------------- - -#[test] -fn dry_run_writes_ndjson_and_live_json() { - let tmp = tempfile::tempdir().unwrap(); - let logs_dir = tmp.path().join("logs"); - - attractor() - .args([ - "run", - "--dry-run", - "--auto-approve", - "--logs-dir", - logs_dir.to_str().unwrap(), - "../../test/simple.dot", - ]) - .assert() - .success(); - - // progress.ndjson must exist and contain valid JSON lines - let ndjson_path = logs_dir.join("progress.ndjson"); - assert!(ndjson_path.exists(), "progress.ndjson should exist"); - let ndjson_content = std::fs::read_to_string(&ndjson_path).unwrap(); - let lines: Vec<&str> = ndjson_content.lines().collect(); - assert!(!lines.is_empty(), "progress.ndjson should have at least one line"); - - // Every line must be valid JSON with timestamp, run_id, and event keys - let first_line: serde_json::Value = serde_json::from_str(lines[0]).unwrap(); - assert!(first_line.get("timestamp").is_some(), "line should have timestamp"); - assert!(first_line.get("run_id").is_some(), "line should have run_id"); - assert!(first_line.get("event").is_some(), "line should have event"); - - // Events should contain PipelineStarted (may not be first due to exec env events) - let has_pipeline_started = lines.iter().any(|line| { - let parsed: serde_json::Value = serde_json::from_str(line).unwrap(); - parsed["event"].get("PipelineStarted").is_some() - }); - assert!(has_pipeline_started, "events should contain PipelineStarted"); - - // run_id should be non-empty after PipelineStarted - let last_line: serde_json::Value = serde_json::from_str(lines[lines.len() - 1]).unwrap(); - let run_id = last_line["run_id"].as_str().unwrap(); - assert!(!run_id.is_empty(), "run_id should be non-empty"); - - // live.json must exist and contain valid JSON matching the last NDJSON line - let live_path = logs_dir.join("live.json"); - assert!(live_path.exists(), "live.json should exist"); - let live_content: serde_json::Value = - serde_json::from_str(&std::fs::read_to_string(&live_path).unwrap()).unwrap(); - assert!(live_content.get("timestamp").is_some()); - assert!(live_content.get("run_id").is_some()); - assert!(live_content.get("event").is_some()); -} diff --git a/crates/arc-cli/Cargo.toml b/crates/arc-cli/Cargo.toml new file mode 100644 index 000000000..ca2c5aad2 --- /dev/null +++ b/crates/arc-cli/Cargo.toml @@ -0,0 +1,27 @@ +[package] +name = "arc-cli" +edition.workspace = true +version.workspace = true +license.workspace = true +description = "Unified CLI for the Arc AI framework" + +[[bin]] +name = "arc" +path = "src/main.rs" + +[dependencies] +arc-llm = { path = "../arc-llm" } +arc-agent = { path = "../arc-agent" } +arc-attractor = { path = "../arc-attractor", features = ["server"] } +arc-util = { path = "../arc-util" } +clap.workspace = true +anyhow.workspace = true +dotenvy.workspace = true +tokio.workspace = true + +[dev-dependencies] +assert_cmd = "2" +predicates = "3" +tempfile = "3" +serde_json.workspace = true +httpmock = "0.8" diff --git a/crates/arc-cli/src/main.rs b/crates/arc-cli/src/main.rs new file mode 100644 index 000000000..d927c1986 --- /dev/null +++ b/crates/arc-cli/src/main.rs @@ -0,0 +1,73 @@ +use anyhow::Result; +use clap::{Parser, Subcommand}; + +#[derive(Parser)] +#[command(name = "arc", version)] +struct Cli { + /// Skip loading .env file + #[arg(long, global = true)] + no_dotenv: bool, + + #[command(subcommand)] + command: Command, +} + +#[derive(Subcommand)] +enum Command { + /// LLM prompt and model operations + Llm { + #[command(subcommand)] + command: LlmCommand, + }, + /// Run an agentic coding session + Agent(arc_agent::cli::AgentArgs), + /// Launch a pipeline + Run(arc_attractor::cli::RunArgs), + /// Validate a pipeline + Validate(arc_attractor::cli::ValidateArgs), + /// Start the HTTP API server + Serve(arc_attractor::cli::ServeArgs), +} + +#[derive(Subcommand)] +enum LlmCommand { + /// Execute a prompt + Prompt(arc_llm::cli::PromptArgs), + /// Manage models + Models { + #[command(subcommand)] + command: Option, + }, +} + +#[tokio::main] +async fn main() -> Result<()> { + let cli = Cli::parse(); + if !cli.no_dotenv { + dotenvy::dotenv().ok(); + } + + match cli.command { + Command::Llm { command } => match command { + LlmCommand::Prompt(args) => arc_llm::cli::run_prompt(args).await?, + LlmCommand::Models { command } => arc_llm::cli::run_models(command).await?, + }, + Command::Agent(args) => arc_agent::cli::run_with_args(args).await?, + Command::Run(args) => { + let styles: &'static arc_util::terminal::Styles = + Box::leak(Box::new(arc_util::terminal::Styles::detect_stderr())); + arc_attractor::cli::run::run_command(args, styles).await?; + } + Command::Validate(args) => { + let styles = arc_util::terminal::Styles::detect_stderr(); + arc_attractor::cli::validate::validate_command(&args, &styles)?; + } + Command::Serve(args) => { + let styles: &'static arc_util::terminal::Styles = + Box::leak(Box::new(arc_util::terminal::Styles::detect_stderr())); + arc_attractor::cli::serve::serve_command(args, styles).await?; + } + } + + Ok(()) +} diff --git a/crates/arc-cli/tests/cli.rs b/crates/arc-cli/tests/cli.rs new file mode 100644 index 000000000..e4a9fe7e5 --- /dev/null +++ b/crates/arc-cli/tests/cli.rs @@ -0,0 +1,512 @@ +use assert_cmd::Command; +use predicates::prelude::*; + +#[allow(deprecated)] +fn arc() -> Command { + Command::cargo_bin("arc").unwrap() +} + +// == LLM: models ============================================================== + +#[test] +fn models_list_prints_all_models() { + arc() + .args(["llm", "models", "list"]) + .assert() + .success() + .stdout(predicate::str::contains("claude-opus-4-6")) + .stdout(predicate::str::contains("claude-sonnet-4-5")) + .stdout(predicate::str::contains("gpt-5.2")) + .stdout(predicate::str::contains("gemini-3.1-pro-preview")) + .stdout(predicate::str::contains("anthropic")) + .stdout(predicate::str::contains("openai")) + .stdout(predicate::str::contains("gemini")); +} + +#[test] +fn models_list_filters_by_provider() { + let assert = arc() + .args(["llm", "models", "list", "--provider", "anthropic"]) + .assert() + .success() + .stdout(predicate::str::contains("claude-opus-4-6")) + .stdout(predicate::str::contains("claude-sonnet-4-5")); + + // Should NOT contain other providers + assert + .stdout(predicate::str::contains("gpt-5.2").not()) + .stdout(predicate::str::contains("gemini-3.1-pro-preview").not()); +} + +#[test] +fn models_list_filters_by_query() { + arc() + .args(["llm", "models", "list", "--query", "opus"]) + .assert() + .success() + .stdout(predicate::str::contains("claude-opus-4-6")) + .stdout(predicate::str::contains("claude-sonnet-4-5").not()); +} + +#[test] +fn models_list_query_is_case_insensitive() { + arc() + .args(["llm", "models", "list", "--query", "OPUS"]) + .assert() + .success() + .stdout(predicate::str::contains("claude-opus-4-6")); +} + +#[test] +fn models_list_query_matches_aliases() { + arc() + .args(["llm", "models", "list", "--query", "codex"]) + .assert() + .success() + .stdout(predicate::str::contains("gpt-5.2-codex")); +} + +#[test] +fn models_bare_defaults_to_list() { + arc() + .args(["llm", "models"]) + .assert() + .success() + .stdout(predicate::str::contains("claude-opus-4-6")) + .stdout(predicate::str::contains("gpt-5.2")) + .stdout(predicate::str::contains("gemini-3.1-pro-preview")); +} + +#[test] +fn models_sync_downloads_and_saves() { + let server = httpmock::MockServer::start(); + let mock_response = serde_json::json!({ + "data": [{"id": "test-model", "name": "Test Model"}] + }); + server.mock(|when, then| { + when.method("GET").path("/api/v1/models"); + then.status(200) + .header("content-type", "application/json") + .body(serde_json::to_string(&mock_response).unwrap()); + }); + + let dir = tempfile::tempdir().unwrap(); + let output_path = dir.path().join("models.json"); + + arc() + .args([ + "llm", + "models", + "sync", + "--url", + &server.url("/api/v1/models"), + "--output", + output_path.to_str().unwrap(), + ]) + .assert() + .success() + .stderr(predicate::str::contains("Saved models to")); + + let contents = std::fs::read_to_string(&output_path).unwrap(); + let expected = serde_json::to_string_pretty(&mock_response).unwrap(); + assert_eq!(contents, expected); +} + +#[test] +fn models_sync_reports_http_errors() { + let server = httpmock::MockServer::start(); + server.mock(|when, then| { + when.method("GET").path("/api/v1/models"); + then.status(500); + }); + + let dir = tempfile::tempdir().unwrap(); + let output_path = dir.path().join("models.json"); + + arc() + .args([ + "llm", + "models", + "sync", + "--url", + &server.url("/api/v1/models"), + "--output", + output_path.to_str().unwrap(), + ]) + .assert() + .failure() + .stderr(predicate::str::contains("error").or(predicate::str::contains("Error"))); +} + +#[test] +#[ignore = "requires network"] +fn models_sync_integration_smoke_test() { + let dir = tempfile::tempdir().unwrap(); + let output_path = dir.path().join("models.json"); + + arc() + .args([ + "llm", + "models", + "sync", + "--output", + output_path.to_str().unwrap(), + ]) + .assert() + .success(); + + let contents = std::fs::read_to_string(&output_path).unwrap(); + assert!(contents.contains("\"data\"")); +} + +#[test] +fn models_sync_help_mentions_openrouter() { + arc() + .args(["llm", "models", "sync", "--help"]) + .assert() + .success() + .stdout(predicate::str::contains("openrouter").or(predicate::str::contains("OpenRouter"))); +} + +// == LLM: prompt ============================================================== + +#[test] +fn prompt_errors_without_prompt_text() { + arc() + .args(["llm", "prompt"]) + .write_stdin("") + .assert() + .failure() + .stderr(predicate::str::contains("no prompt provided")); +} + +#[test] +fn prompt_reads_from_stdin() { + let result = arc() + .args(["--no-dotenv", "llm", "prompt", "--no-stream", "-m", "test-model"]) + .write_stdin("hello from stdin") + .assert() + .failure(); + + // Should NOT complain about missing prompt + result.stderr(predicate::str::contains("no prompt provided").not()); +} + +#[test] +fn prompt_concatenates_stdin_and_arg() { + let result = arc() + .args(["--no-dotenv", "llm", "prompt", "--no-stream", "-m", "test-model", "summarize this"]) + .write_stdin("some input text") + .assert() + .failure(); + + result.stderr(predicate::str::contains("no prompt provided").not()); +} + +#[test] +fn prompt_rejects_bad_option_format() { + arc() + .args(["llm", "prompt", "-o", "bad_option", "hello"]) + .assert() + .failure() + .stderr(predicate::str::contains("expected key=value")); +} + +#[test] +#[ignore = "requires API key"] +fn prompt_no_stream_generates_response() { + arc() + .args(["llm", "prompt", "--no-stream", "-m", "claude-sonnet-4-5", "Say just the word 'hello'"]) + .assert() + .success() + .stdout(predicate::str::is_empty().not()); +} + +#[test] +#[ignore = "requires API key"] +fn prompt_stream_generates_response() { + arc() + .args(["llm", "prompt", "-m", "claude-sonnet-4-5", "Say just the word 'hello'"]) + .assert() + .success() + .stdout(predicate::str::is_empty().not()); +} + +#[test] +#[ignore = "requires API key"] +fn prompt_usage_shows_tokens() { + arc() + .args(["llm", "prompt", "--no-stream", "-u", "-m", "claude-sonnet-4-5", "Say just the word 'hello'"]) + .assert() + .success() + .stderr(predicate::str::contains("Tokens:")); +} + +#[test] +fn prompt_schema_rejects_invalid_json() { + arc() + .args(["--no-dotenv", "llm", "prompt", "--no-stream", "-m", "test-model", "--schema", "not json", "hello"]) + .assert() + .failure() + .stderr(predicate::str::contains("--schema must be valid JSON")); +} + +#[test] +#[ignore = "requires API key"] +fn prompt_schema_no_stream_generates_json() { + let assert = arc() + .args([ + "llm", "prompt", "--no-stream", "-m", "claude-sonnet-4-5", + "--schema", r#"{"type":"object","properties":{"greeting":{"type":"string"}},"required":["greeting"]}"#, + "Return a JSON object with a greeting field set to hello", + ]) + .assert() + .success(); + + let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap(); + let parsed: serde_json::Value = serde_json::from_str(stdout.trim()).expect("stdout should be valid JSON"); + assert!(parsed.get("greeting").is_some(), "expected 'greeting' key in output"); +} + +#[test] +#[ignore = "requires API key"] +fn prompt_schema_stream_generates_json() { + let assert = arc() + .args([ + "llm", "prompt", "-m", "claude-sonnet-4-5", + "--schema", r#"{"type":"object","properties":{"greeting":{"type":"string"}},"required":["greeting"]}"#, + "Return a JSON object with a greeting field set to hello", + ]) + .assert() + .success(); + + let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap(); + let parsed: serde_json::Value = serde_json::from_str(stdout.trim()).expect("stdout should be valid JSON"); + assert!(parsed.get("greeting").is_some(), "expected 'greeting' key in output"); +} + +// == Agent ==================================================================== + +#[test] +fn agent_no_prompt_prints_usage() { + arc() + .args(["agent"]) + .env_clear() + .assert() + .failure() + .stderr(predicate::str::contains("Usage:")); +} + +#[test] +fn agent_help_flag_prints_help() { + arc() + .args(["agent", "--help"]) + .assert() + .success() + .stdout(predicate::str::contains("Task prompt")); +} + +#[test] +fn agent_missing_api_key_exits_with_error() { + let tmp = std::env::temp_dir(); + arc() + .args(["--no-dotenv", "agent", "test prompt"]) + .env_clear() + .current_dir(&tmp) + .assert() + .failure() + .stderr(predicate::str::contains("API key not set")); +} + +#[test] +fn agent_invalid_permissions_value() { + arc() + .args(["agent", "--permissions", "bogus", "test prompt"]) + .env_clear() + .assert() + .failure() + .stderr(predicate::str::contains("invalid value")); +} + +// == Attractor: validate ====================================================== + +#[test] +fn validate_simple() { + arc() + .args(["validate", "../../test/simple.dot"]) + .assert() + .success() + .stderr(predicate::str::contains("Validation: OK")); +} + +#[test] +fn validate_branching() { + arc() + .args(["validate", "../../test/branching.dot"]) + .assert() + .success() + .stderr(predicate::str::contains("Validation: OK")); +} + +#[test] +fn validate_conditions() { + arc() + .args(["validate", "../../test/conditions.dot"]) + .assert() + .success() + .stderr(predicate::str::contains("Validation: OK")); +} + +#[test] +fn validate_parallel() { + arc() + .args(["validate", "../../test/parallel.dot"]) + .assert() + .success() + .stderr(predicate::str::contains("Validation: OK")); +} + +#[test] +fn validate_styled() { + arc() + .args(["validate", "../../test/styled.dot"]) + .assert() + .success() + .stderr(predicate::str::contains("Validation: OK")); +} + +#[test] +fn validate_legacy_tool() { + arc() + .args(["validate", "../../test/legacy_tool.dot"]) + .assert() + .success() + .stderr(predicate::str::contains("Validation: OK")); +} + +#[test] +fn validate_invalid() { + arc() + .args(["validate", "../../test/invalid.dot"]) + .assert() + .failure(); +} + +// == Attractor: serve ========================================================= + +#[test] +fn serve_help() { + arc() + .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")); +} + +// == Attractor: run --dry-run ================================================= + +#[test] +fn dry_run_simple() { + arc() + .args(["run", "--dry-run", "--auto-approve", "../../test/simple.dot"]) + .assert() + .success(); +} + +#[test] +fn dry_run_branching() { + arc() + .args(["run", "--dry-run", "--auto-approve", "../../test/branching.dot"]) + .assert() + .success(); +} + +#[test] +fn dry_run_conditions() { + arc() + .args(["run", "--dry-run", "--auto-approve", "../../test/conditions.dot"]) + .assert() + .success(); +} + +#[test] +fn dry_run_parallel() { + arc() + .args(["run", "--dry-run", "--auto-approve", "../../test/parallel.dot"]) + .assert() + .success(); +} + +#[test] +fn dry_run_styled() { + arc() + .args(["run", "--dry-run", "--auto-approve", "../../test/styled.dot"]) + .assert() + .success(); +} + +#[test] +fn dry_run_legacy_tool() { + arc() + .args(["run", "--dry-run", "--auto-approve", "../../test/legacy_tool.dot"]) + .assert() + .success(); +} + +// == NDJSON logging =========================================================== + +#[test] +fn dry_run_writes_ndjson_and_live_json() { + let tmp = tempfile::tempdir().unwrap(); + let logs_dir = tmp.path().join("logs"); + + arc() + .args([ + "run", + "--dry-run", + "--auto-approve", + "--logs-dir", + logs_dir.to_str().unwrap(), + "../../test/simple.dot", + ]) + .assert() + .success(); + + // progress.ndjson must exist and contain valid JSON lines + let ndjson_path = logs_dir.join("progress.ndjson"); + assert!(ndjson_path.exists(), "progress.ndjson should exist"); + let ndjson_content = std::fs::read_to_string(&ndjson_path).unwrap(); + let lines: Vec<&str> = ndjson_content.lines().collect(); + assert!(!lines.is_empty(), "progress.ndjson should have at least one line"); + + // Every line must be valid JSON with timestamp, run_id, and event keys + let first_line: serde_json::Value = serde_json::from_str(lines[0]).unwrap(); + assert!(first_line.get("timestamp").is_some(), "line should have timestamp"); + assert!(first_line.get("run_id").is_some(), "line should have run_id"); + assert!(first_line.get("event").is_some(), "line should have event"); + + // Events should contain PipelineStarted (may not be first due to exec env events) + let has_pipeline_started = lines.iter().any(|line| { + let parsed: serde_json::Value = serde_json::from_str(line).unwrap(); + parsed["event"].get("PipelineStarted").is_some() + }); + assert!(has_pipeline_started, "events should contain PipelineStarted"); + + // run_id should be non-empty after PipelineStarted + let last_line: serde_json::Value = serde_json::from_str(lines[lines.len() - 1]).unwrap(); + let run_id = last_line["run_id"].as_str().unwrap(); + assert!(!run_id.is_empty(), "run_id should be non-empty"); + + // live.json must exist and contain valid JSON matching the last NDJSON line + let live_path = logs_dir.join("live.json"); + assert!(live_path.exists(), "live.json should exist"); + let live_content: serde_json::Value = + serde_json::from_str(&std::fs::read_to_string(&live_path).unwrap()).unwrap(); + assert!(live_content.get("timestamp").is_some()); + assert!(live_content.get("run_id").is_some()); + assert!(live_content.get("event").is_some()); +} diff --git a/crates/arc-llm/Cargo.toml b/crates/arc-llm/Cargo.toml index 9dbc83530..5a4296cb9 100644 --- a/crates/arc-llm/Cargo.toml +++ b/crates/arc-llm/Cargo.toml @@ -12,10 +12,6 @@ categories = ["api-bindings"] [lib] doctest = false -[[bin]] -name = "ullm" -path = "src/bin/ullm.rs" - [dependencies] anyhow.workspace = true thiserror.workspace = true @@ -32,12 +28,8 @@ base64.workspace = true bytes.workspace = true tokio-util.workspace = true clap.workspace = true -dotenvy.workspace = true [dev-dependencies] http = "1" tokio = { workspace = true, features = ["test-util", "macros"] } -assert_cmd = "2" -predicates = "3" -httpmock = "0.8" -tempfile = "3" +dotenvy.workspace = true diff --git a/crates/arc-llm/src/bin/ullm.rs b/crates/arc-llm/src/bin/ullm.rs deleted file mode 100644 index 7e845efa7..000000000 --- a/crates/arc-llm/src/bin/ullm.rs +++ /dev/null @@ -1,644 +0,0 @@ -use std::io::{self, IsTerminal, Read}; - -use anyhow::{bail, Context, Result}; -use clap::{Parser, Subcommand}; -use futures::StreamExt; -use arc_llm::catalog; -use arc_llm::generate::{self, GenerateParams}; - -#[derive(Parser)] -#[command(name = "ullm")] -struct Cli { - /// Skip loading .env file - #[arg(long, global = true)] - no_dotenv: bool, - - #[command(subcommand)] - command: Command, -} - -#[derive(Subcommand)] -enum Command { - /// Execute a prompt - Prompt { - /// The prompt text (also accepts stdin) - prompt: Option, - - /// Model to use - #[arg(short, long)] - model: Option, - - /// System prompt - #[arg(short, long)] - system: Option, - - /// Do not stream output - #[arg(long)] - no_stream: bool, - - /// Show token usage - #[arg(short, long)] - usage: bool, - - /// JSON schema for structured output (inline JSON string) - #[arg(short = 'S', long)] - schema: Option, - - /// key=value options (temperature, `max_tokens`, `top_p`) - #[arg(short, long, value_parser = parse_option)] - option: Vec<(String, String)>, - }, - - /// Manage models - Models { - #[command(subcommand)] - command: Option, - }, -} - -#[derive(Subcommand)] -enum ModelsCommand { - /// List available models - List { - /// Filter by provider - #[arg(short, long)] - provider: Option, - - /// Search for models matching this string - #[arg(short, long)] - query: Option, - }, - - /// Download model metadata from OpenRouter - Sync { - /// URL to fetch models from - #[arg(long, default_value = "https://openrouter.ai/api/v1/models")] - url: String, - - /// Output file path - #[arg(short, long, default_value = "openrouter_models.json")] - output: String, - }, -} - -fn parse_option(s: &str) -> Result<(String, String), String> { - let (key, value) = s - .split_once('=') - .ok_or_else(|| format!("expected key=value, got {s}"))?; - Ok((key.to_string(), value.to_string())) -} - -fn print_models_table(models: &[arc_llm::types::ModelInfo]) { - println!( - "{:<30} {:<12} {:<30} {:>14}", - "ID", "PROVIDER", "ALIASES", "CONTEXT" - ); - for model in models { - let aliases = model.aliases.join(", "); - println!( - "{:<30} {:<12} {:<30} {:>14}", - model.id, model.provider, aliases, model.context_window - ); - } -} - -fn read_stdin_prompt() -> Option { - let stdin = io::stdin(); - if stdin.is_terminal() { - return None; - } - let mut buf = String::new(); - stdin.lock().read_to_string(&mut buf).ok()?; - let trimmed = buf.trim(); - if trimmed.is_empty() { - None - } else { - Some(trimmed.to_string()) - } -} - -fn resolve_prompt(arg: Option, stdin: Option) -> Result { - match (stdin, arg) { - (Some(s), Some(a)) => Ok(format!("{s}\n{a}")), - (Some(s), None) => Ok(s), - (None, Some(a)) => Ok(a), - (None, None) => bail!("Error: no prompt provided. Pass a prompt as an argument or pipe text via stdin."), - } -} - -/// Returns (`model_id`, provider) from the catalog, falling back to the first catalog model. -fn resolve_model(model_arg: Option) -> (String, Option) { - let raw = model_arg.unwrap_or_else(|| { - catalog::list_models(None) - .first() - .map_or_else(|| "claude-sonnet-4-5".to_string(), |m| m.id.clone()) - }); - match catalog::get_model_info(&raw) { - Some(info) => (info.id, Some(info.provider)), - None => (raw, None), - } -} - -fn apply_options( - mut params: GenerateParams, - options: &[(String, String)], -) -> Result { - let mut provider_opts = serde_json::Map::new(); - - for (key, value) in options { - match key.as_str() { - "temperature" => { - let v: f64 = value - .parse() - .with_context(|| format!("invalid temperature value: {value}"))?; - params = params.temperature(v); - } - "max_tokens" => { - let v: i64 = value - .parse() - .with_context(|| format!("invalid max_tokens value: {value}"))?; - params = params.max_tokens(v); - } - "top_p" => { - let v: f64 = value - .parse() - .with_context(|| format!("invalid top_p value: {value}"))?; - params = params.top_p(v); - } - _ => { - provider_opts.insert(key.clone(), serde_json::Value::String(value.clone())); - } - } - } - - if !provider_opts.is_empty() { - params = params.provider_options(serde_json::Value::Object(provider_opts)); - } - - Ok(params) -} - -struct PromptArgs { - prompt: Option, - model: Option, - system: Option, - no_stream: bool, - usage: bool, - schema: Option, - option: Vec<(String, String)>, -} - -fn print_usage(usage: &arc_llm::types::Usage) { - eprintln!( - "Tokens: {} input, {} output, {} total", - usage.input_tokens, usage.output_tokens, usage.total_tokens - ); -} - -async fn run_prompt(args: PromptArgs) -> Result<()> { - let stdin_prompt = read_stdin_prompt(); - let prompt_text = resolve_prompt(args.prompt, stdin_prompt)?; - let (model_id, provider) = resolve_model(args.model); - - eprintln!("Using model: {model_id}"); - - let mut params = GenerateParams::new(&model_id).prompt(&prompt_text); - if let Some(p) = provider { - params = params.provider(&p); - } - if let Some(sys) = args.system { - params = params.system(&sys); - } - params = apply_options(params, &args.option)?; - - let schema: Option = match &args.schema { - Some(s) => Some(serde_json::from_str(s).context("--schema must be valid JSON")?), - None => None, - }; - - match (args.no_stream, schema) { - (true, Some(schema)) => { - let result = generate::generate_object(params, schema).await?; - let object = result.output.as_ref().unwrap_or(&serde_json::Value::Null); - println!("{}", serde_json::to_string_pretty(object)?); - if args.usage { - print_usage(&result.usage); - } - } - (true, None) => { - let result = generate::generate(params).await?; - print!("{}", result.text()); - if args.usage { - print_usage(&result.usage); - } - } - (false, Some(schema)) => { - let mut stream_result = generate::stream_object(params, schema).await?; - while let Some(event) = stream_result.next().await { - event?; - } - if let Some(object) = stream_result.object() { - println!("{}", serde_json::to_string_pretty(object)?); - } - } - (false, None) => { - let mut stream_result = generate::stream(params).await?; - while let Some(event) = stream_result.next().await { - if let arc_llm::types::StreamEvent::TextDelta { delta, .. } = event? { - print!("{delta}"); - } - } - println!(); - if args.usage { - if let Some(response) = stream_result.response() { - print_usage(&response.usage); - } - } - } - } - - Ok(()) -} - -async fn sync_models(url: &str, output: &str) -> Result<()> { - let body = reqwest::get(url) - .await - .context("failed to connect to models endpoint")? - .error_for_status() - .context("models endpoint returned an error")? - .text() - .await - .context("failed to read response body")?; - - let json: serde_json::Value = - serde_json::from_str(&body).context("response is not valid JSON")?; - let pretty = - serde_json::to_string_pretty(&json).context("failed to format JSON")?; - - std::fs::write(output, &pretty).with_context(|| format!("failed to write {output}"))?; - - eprintln!("Saved models to {output}"); - Ok(()) -} - -#[tokio::main] -async fn main() -> Result<()> { - let cli = Cli::parse(); - if !cli.no_dotenv { - dotenvy::dotenv().ok(); - } - - match cli.command { - Command::Prompt { - prompt, - model, - system, - no_stream, - usage, - schema, - option, - } => { - run_prompt(PromptArgs { - prompt, - model, - system, - no_stream, - usage, - schema, - option, - }) - .await?; - } - Command::Models { command } => { - let command = command.unwrap_or(ModelsCommand::List { - provider: None, - query: None, - }); - - match command { - ModelsCommand::List { provider, query } => { - let mut models = catalog::list_models(provider.as_deref()); - - if let Some(q) = &query { - let q_lower = q.to_lowercase(); - models.retain(|m| { - m.id.to_lowercase().contains(&q_lower) - || m.display_name.to_lowercase().contains(&q_lower) - || m.aliases - .iter() - .any(|a| a.to_lowercase().contains(&q_lower)) - }); - } - - print_models_table(&models); - } - ModelsCommand::Sync { url, output } => { - sync_models(&url, &output).await?; - } - } - } - } - - Ok(()) -} - -#[cfg(test)] -mod tests { - use assert_cmd::Command; - use predicates::prelude::*; - - #[allow(deprecated)] // assert_cmd deprecated cargo_bin; replacement macro has issues - fn ullm() -> Command { - Command::cargo_bin("ullm").unwrap() - } - - // Step 1: models list prints all catalog models - #[test] - fn models_list_prints_all_models() { - ullm() - .args(["models", "list"]) - .assert() - .success() - .stdout(predicate::str::contains("claude-opus-4-6")) - .stdout(predicate::str::contains("claude-sonnet-4-5")) - .stdout(predicate::str::contains("gpt-5.2")) - .stdout(predicate::str::contains("gemini-3.1-pro-preview")) - .stdout(predicate::str::contains("anthropic")) - .stdout(predicate::str::contains("openai")) - .stdout(predicate::str::contains("gemini")); - } - - // Step 2: models list --provider filters to that provider only - #[test] - fn models_list_filters_by_provider() { - let assert = ullm() - .args(["models", "list", "--provider", "anthropic"]) - .assert() - .success() - .stdout(predicate::str::contains("claude-opus-4-6")) - .stdout(predicate::str::contains("claude-sonnet-4-5")); - - // Should NOT contain other providers - assert - .stdout(predicate::str::contains("gpt-5.2").not()) - .stdout(predicate::str::contains("gemini-3.1-pro-preview").not()); - } - - // Step 3: models list --query does substring match on id/name/aliases - #[test] - fn models_list_filters_by_query() { - ullm() - .args(["models", "list", "--query", "opus"]) - .assert() - .success() - .stdout(predicate::str::contains("claude-opus-4-6")) - .stdout(predicate::str::contains("claude-sonnet-4-5").not()); - } - - #[test] - fn models_list_query_is_case_insensitive() { - ullm() - .args(["models", "list", "--query", "OPUS"]) - .assert() - .success() - .stdout(predicate::str::contains("claude-opus-4-6")); - } - - #[test] - fn models_list_query_matches_aliases() { - ullm() - .args(["models", "list", "--query", "codex"]) - .assert() - .success() - .stdout(predicate::str::contains("gpt-5.2-codex")); - } - - // Step 4: bare "models" defaults to list - #[test] - fn models_bare_defaults_to_list() { - ullm() - .args(["models"]) - .assert() - .success() - .stdout(predicate::str::contains("claude-opus-4-6")) - .stdout(predicate::str::contains("gpt-5.2")) - .stdout(predicate::str::contains("gemini-3.1-pro-preview")); - } - - // models sync downloads and saves pretty-printed JSON - #[test] - fn models_sync_downloads_and_saves() { - let server = httpmock::MockServer::start(); - let mock_response = serde_json::json!({ - "data": [{"id": "test-model", "name": "Test Model"}] - }); - server.mock(|when, then| { - when.method("GET").path("/api/v1/models"); - then.status(200) - .header("content-type", "application/json") - .body(serde_json::to_string(&mock_response).unwrap()); - }); - - let dir = tempfile::tempdir().unwrap(); - let output_path = dir.path().join("models.json"); - - ullm() - .args([ - "models", - "sync", - "--url", - &server.url("/api/v1/models"), - "--output", - output_path.to_str().unwrap(), - ]) - .assert() - .success() - .stderr(predicate::str::contains("Saved models to")); - - let contents = std::fs::read_to_string(&output_path).unwrap(); - let expected = serde_json::to_string_pretty(&mock_response).unwrap(); - assert_eq!(contents, expected); - } - - // models sync reports HTTP errors - #[test] - fn models_sync_reports_http_errors() { - let server = httpmock::MockServer::start(); - server.mock(|when, then| { - when.method("GET").path("/api/v1/models"); - then.status(500); - }); - - let dir = tempfile::tempdir().unwrap(); - let output_path = dir.path().join("models.json"); - - ullm() - .args([ - "models", - "sync", - "--url", - &server.url("/api/v1/models"), - "--output", - output_path.to_str().unwrap(), - ]) - .assert() - .failure() - .stderr(predicate::str::contains("error").or(predicate::str::contains("Error"))); - } - - // models sync with real OpenRouter (requires network) - #[test] - #[ignore = "requires network"] - fn models_sync_integration_smoke_test() { - let dir = tempfile::tempdir().unwrap(); - let output_path = dir.path().join("models.json"); - - ullm() - .args([ - "models", - "sync", - "--output", - output_path.to_str().unwrap(), - ]) - .assert() - .success(); - - let contents = std::fs::read_to_string(&output_path).unwrap(); - assert!(contents.contains("\"data\"")); - } - - // models sync --help succeeds and mentions openrouter - #[test] - fn models_sync_help_mentions_openrouter() { - ullm() - .args(["models", "sync", "--help"]) - .assert() - .success() - .stdout(predicate::str::contains("openrouter").or(predicate::str::contains("OpenRouter"))); - } - - // Step 5: prompt requires prompt text (errors when no prompt and stdin is tty) - #[test] - fn prompt_errors_without_prompt_text() { - // assert_cmd provides no stdin by default (simulating a tty-like "empty pipe") - // We pass an empty stdin to avoid tty detection - ullm() - .args(["prompt"]) - .write_stdin("") - .assert() - .failure() - .stderr(predicate::str::contains("no prompt provided")); - } - - // Step 9: stdin piping — reads from stdin when no prompt arg - #[test] - fn prompt_reads_from_stdin() { - // This test verifies stdin is read, but will fail at the API call stage - // since no API key is set. The error should NOT be "no prompt provided". - let result = ullm() - .args(["--no-dotenv", "prompt", "--no-stream", "-m", "test-model"]) - .write_stdin("hello from stdin") - .assert() - .failure(); - - // Should NOT complain about missing prompt - result.stderr(predicate::str::contains("no prompt provided").not()); - } - - // Step 9b: stdin + arg concatenation - #[test] - fn prompt_concatenates_stdin_and_arg() { - // Same as above — verifies it doesn't error on "no prompt" - let result = ullm() - .args(["--no-dotenv", "prompt", "--no-stream", "-m", "test-model", "summarize this"]) - .write_stdin("some input text") - .assert() - .failure(); - - result.stderr(predicate::str::contains("no prompt provided").not()); - } - - // Step 10: -o option parsing - #[test] - fn prompt_rejects_bad_option_format() { - ullm() - .args(["prompt", "-o", "bad_option", "hello"]) - .assert() - .failure() - .stderr(predicate::str::contains("expected key=value")); - } - - // Step 6/7/8: Integration tests gated behind API key - #[test] - #[ignore = "requires API key"] - fn prompt_no_stream_generates_response() { - ullm() - .args(["prompt", "--no-stream", "-m", "claude-sonnet-4-5", "Say just the word 'hello'"]) - .assert() - .success() - .stdout(predicate::str::is_empty().not()); - } - - #[test] - #[ignore = "requires API key"] - fn prompt_stream_generates_response() { - ullm() - .args(["prompt", "-m", "claude-sonnet-4-5", "Say just the word 'hello'"]) - .assert() - .success() - .stdout(predicate::str::is_empty().not()); - } - - #[test] - #[ignore = "requires API key"] - fn prompt_usage_shows_tokens() { - ullm() - .args(["prompt", "--no-stream", "-u", "-m", "claude-sonnet-4-5", "Say just the word 'hello'"]) - .assert() - .success() - .stderr(predicate::str::contains("Tokens:")); - } - - #[test] - fn prompt_schema_rejects_invalid_json() { - ullm() - .args(["--no-dotenv", "prompt", "--no-stream", "-m", "test-model", "--schema", "not json", "hello"]) - .assert() - .failure() - .stderr(predicate::str::contains("--schema must be valid JSON")); - } - - #[test] - #[ignore = "requires API key"] - fn prompt_schema_no_stream_generates_json() { - let assert = ullm() - .args([ - "prompt", "--no-stream", "-m", "claude-sonnet-4-5", - "--schema", r#"{"type":"object","properties":{"greeting":{"type":"string"}},"required":["greeting"]}"#, - "Return a JSON object with a greeting field set to hello", - ]) - .assert() - .success(); - - let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap(); - let parsed: serde_json::Value = serde_json::from_str(stdout.trim()).expect("stdout should be valid JSON"); - assert!(parsed.get("greeting").is_some(), "expected 'greeting' key in output"); - } - - #[test] - #[ignore = "requires API key"] - fn prompt_schema_stream_generates_json() { - let assert = ullm() - .args([ - "prompt", "-m", "claude-sonnet-4-5", - "--schema", r#"{"type":"object","properties":{"greeting":{"type":"string"}},"required":["greeting"]}"#, - "Return a JSON object with a greeting field set to hello", - ]) - .assert() - .success(); - - let stdout = String::from_utf8(assert.get_output().stdout.clone()).unwrap(); - let parsed: serde_json::Value = serde_json::from_str(stdout.trim()).expect("stdout should be valid JSON"); - assert!(parsed.get("greeting").is_some(), "expected 'greeting' key in output"); - } -} diff --git a/crates/arc-llm/src/cli.rs b/crates/arc-llm/src/cli.rs new file mode 100644 index 000000000..0949f4899 --- /dev/null +++ b/crates/arc-llm/src/cli.rs @@ -0,0 +1,284 @@ +use std::io::{self, IsTerminal, Read}; + +use anyhow::{bail, Context, Result}; +use clap::{Args, Subcommand}; +use futures::StreamExt; + +use crate::catalog; +use crate::generate::{self, GenerateParams}; + +#[derive(Args)] +pub struct PromptArgs { + /// The prompt text (also accepts stdin) + pub prompt: Option, + + /// Model to use + #[arg(short, long)] + pub model: Option, + + /// System prompt + #[arg(short, long)] + pub system: Option, + + /// Do not stream output + #[arg(long)] + pub no_stream: bool, + + /// Show token usage + #[arg(short, long)] + pub usage: bool, + + /// JSON schema for structured output (inline JSON string) + #[arg(short = 'S', long)] + pub schema: Option, + + /// key=value options (temperature, `max_tokens`, `top_p`) + #[arg(short, long, value_parser = parse_option)] + pub option: Vec<(String, String)>, +} + +#[derive(Subcommand)] +pub enum ModelsCommand { + /// List available models + List { + /// Filter by provider + #[arg(short, long)] + provider: Option, + + /// Search for models matching this string + #[arg(short, long)] + query: Option, + }, + + /// Download model metadata from OpenRouter + Sync { + /// URL to fetch models from + #[arg(long, default_value = "https://openrouter.ai/api/v1/models")] + url: String, + + /// Output file path + #[arg(short, long, default_value = "openrouter_models.json")] + output: String, + }, +} + +fn parse_option(s: &str) -> Result<(String, String), String> { + let (key, value) = s + .split_once('=') + .ok_or_else(|| format!("expected key=value, got {s}"))?; + Ok((key.to_string(), value.to_string())) +} + +fn print_models_table(models: &[crate::types::ModelInfo]) { + println!( + "{:<30} {:<12} {:<30} {:>14}", + "ID", "PROVIDER", "ALIASES", "CONTEXT" + ); + for model in models { + let aliases = model.aliases.join(", "); + println!( + "{:<30} {:<12} {:<30} {:>14}", + model.id, model.provider, aliases, model.context_window + ); + } +} + +fn read_stdin_prompt() -> Option { + let stdin = io::stdin(); + if stdin.is_terminal() { + return None; + } + let mut buf = String::new(); + stdin.lock().read_to_string(&mut buf).ok()?; + let trimmed = buf.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.to_string()) + } +} + +fn resolve_prompt(arg: Option, stdin: Option) -> Result { + match (stdin, arg) { + (Some(s), Some(a)) => Ok(format!("{s}\n{a}")), + (Some(s), None) => Ok(s), + (None, Some(a)) => Ok(a), + (None, None) => bail!("Error: no prompt provided. Pass a prompt as an argument or pipe text via stdin."), + } +} + +/// Returns (`model_id`, provider) from the catalog, falling back to the first catalog model. +fn resolve_model(model_arg: Option) -> (String, Option) { + let raw = model_arg.unwrap_or_else(|| { + catalog::list_models(None) + .first() + .map_or_else(|| "claude-sonnet-4-5".to_string(), |m| m.id.clone()) + }); + match catalog::get_model_info(&raw) { + Some(info) => (info.id, Some(info.provider)), + None => (raw, None), + } +} + +fn apply_options( + mut params: GenerateParams, + options: &[(String, String)], +) -> Result { + let mut provider_opts = serde_json::Map::new(); + + for (key, value) in options { + match key.as_str() { + "temperature" => { + let v: f64 = value + .parse() + .with_context(|| format!("invalid temperature value: {value}"))?; + params = params.temperature(v); + } + "max_tokens" => { + let v: i64 = value + .parse() + .with_context(|| format!("invalid max_tokens value: {value}"))?; + params = params.max_tokens(v); + } + "top_p" => { + let v: f64 = value + .parse() + .with_context(|| format!("invalid top_p value: {value}"))?; + params = params.top_p(v); + } + _ => { + provider_opts.insert(key.clone(), serde_json::Value::String(value.clone())); + } + } + } + + if !provider_opts.is_empty() { + params = params.provider_options(serde_json::Value::Object(provider_opts)); + } + + Ok(params) +} + +fn print_usage(usage: &crate::types::Usage) { + eprintln!( + "Tokens: {} input, {} output, {} total", + usage.input_tokens, usage.output_tokens, usage.total_tokens + ); +} + +pub async fn run_prompt(args: PromptArgs) -> Result<()> { + let stdin_prompt = read_stdin_prompt(); + let prompt_text = resolve_prompt(args.prompt, stdin_prompt)?; + let (model_id, provider) = resolve_model(args.model); + + eprintln!("Using model: {model_id}"); + + let mut params = GenerateParams::new(&model_id).prompt(&prompt_text); + if let Some(p) = provider { + params = params.provider(&p); + } + if let Some(sys) = args.system { + params = params.system(&sys); + } + params = apply_options(params, &args.option)?; + + let schema: Option = match &args.schema { + Some(s) => Some(serde_json::from_str(s).context("--schema must be valid JSON")?), + None => None, + }; + + match (args.no_stream, schema) { + (true, Some(schema)) => { + let result = generate::generate_object(params, schema).await?; + let object = result.output.as_ref().unwrap_or(&serde_json::Value::Null); + println!("{}", serde_json::to_string_pretty(object)?); + if args.usage { + print_usage(&result.usage); + } + } + (true, None) => { + let result = generate::generate(params).await?; + print!("{}", result.text()); + if args.usage { + print_usage(&result.usage); + } + } + (false, Some(schema)) => { + let mut stream_result = generate::stream_object(params, schema).await?; + while let Some(event) = stream_result.next().await { + event?; + } + if let Some(object) = stream_result.object() { + println!("{}", serde_json::to_string_pretty(object)?); + } + } + (false, None) => { + let mut stream_result = generate::stream(params).await?; + while let Some(event) = stream_result.next().await { + if let crate::types::StreamEvent::TextDelta { delta, .. } = event? { + print!("{delta}"); + } + } + println!(); + if args.usage { + if let Some(response) = stream_result.response() { + print_usage(&response.usage); + } + } + } + } + + Ok(()) +} + +async fn sync_models(url: &str, output: &str) -> Result<()> { + let body = reqwest::get(url) + .await + .context("failed to connect to models endpoint")? + .error_for_status() + .context("models endpoint returned an error")? + .text() + .await + .context("failed to read response body")?; + + let json: serde_json::Value = + serde_json::from_str(&body).context("response is not valid JSON")?; + let pretty = + serde_json::to_string_pretty(&json).context("failed to format JSON")?; + + std::fs::write(output, &pretty).with_context(|| format!("failed to write {output}"))?; + + eprintln!("Saved models to {output}"); + Ok(()) +} + +pub async fn run_models(command: Option) -> Result<()> { + let command = command.unwrap_or(ModelsCommand::List { + provider: None, + query: None, + }); + + match command { + ModelsCommand::List { provider, query } => { + let mut models = catalog::list_models(provider.as_deref()); + + if let Some(q) = &query { + let q_lower = q.to_lowercase(); + models.retain(|m| { + m.id.to_lowercase().contains(&q_lower) + || m.display_name.to_lowercase().contains(&q_lower) + || m.aliases + .iter() + .any(|a| a.to_lowercase().contains(&q_lower)) + }); + } + + print_models_table(&models); + } + ModelsCommand::Sync { url, output } => { + sync_models(&url, &output).await?; + } + } + + Ok(()) +} diff --git a/crates/arc-llm/src/lib.rs b/crates/arc-llm/src/lib.rs index 747bb09fe..3d68d2495 100644 --- a/crates/arc-llm/src/lib.rs +++ b/crates/arc-llm/src/lib.rs @@ -8,6 +8,7 @@ pub mod retry; pub mod generate; pub mod catalog; pub mod providers; +pub mod cli; // Re-export module-level default client helpers (Section 2.5). pub use generate::set_default_client;