mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-06 02:48:25 +00:00
Unify ullm, arc-agent, arc-attractor into single arc binary
Three separate binaries are replaced by a single `arc` CLI with subcommands: arc llm prompt/models, arc agent, arc run, arc validate, arc serve Extract public CLI modules (arc_llm::cli, arc_agent::cli::AgentArgs/run_with_args) so the new arc-cli crate can dispatch to each library. Integration tests migrate to crates/arc-cli/tests/cli.rs. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
1b32e1592c
commit
2d79362750
15 changed files with 956 additions and 994 deletions
23
Cargo.lock
generated
23
Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<String>,
|
||||
pub model: Option<String>,
|
||||
|
||||
/// 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<String>,
|
||||
pub skills_dir: Option<String>,
|
||||
|
||||
/// 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<dyn crate::ExecutionEnvironment> = 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::*;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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}"
|
||||
);
|
||||
}
|
||||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
27
crates/arc-cli/Cargo.toml
Normal file
27
crates/arc-cli/Cargo.toml
Normal file
|
|
@ -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"
|
||||
73
crates/arc-cli/src/main.rs
Normal file
73
crates/arc-cli/src/main.rs
Normal file
|
|
@ -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<arc_llm::cli::ModelsCommand>,
|
||||
},
|
||||
}
|
||||
|
||||
#[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(())
|
||||
}
|
||||
512
crates/arc-cli/tests/cli.rs
Normal file
512
crates/arc-cli/tests/cli.rs
Normal file
|
|
@ -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());
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<String>,
|
||||
|
||||
/// Model to use
|
||||
#[arg(short, long)]
|
||||
model: Option<String>,
|
||||
|
||||
/// System prompt
|
||||
#[arg(short, long)]
|
||||
system: Option<String>,
|
||||
|
||||
/// 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<String>,
|
||||
|
||||
/// 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<ModelsCommand>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Subcommand)]
|
||||
enum ModelsCommand {
|
||||
/// List available models
|
||||
List {
|
||||
/// Filter by provider
|
||||
#[arg(short, long)]
|
||||
provider: Option<String>,
|
||||
|
||||
/// Search for models matching this string
|
||||
#[arg(short, long)]
|
||||
query: Option<String>,
|
||||
},
|
||||
|
||||
/// 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<String> {
|
||||
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<String>, stdin: Option<String>) -> Result<String> {
|
||||
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>) -> (String, Option<String>) {
|
||||
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<GenerateParams> {
|
||||
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<String>,
|
||||
model: Option<String>,
|
||||
system: Option<String>,
|
||||
no_stream: bool,
|
||||
usage: bool,
|
||||
schema: Option<String>,
|
||||
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<serde_json::Value> = 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");
|
||||
}
|
||||
}
|
||||
284
crates/arc-llm/src/cli.rs
Normal file
284
crates/arc-llm/src/cli.rs
Normal file
|
|
@ -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<String>,
|
||||
|
||||
/// Model to use
|
||||
#[arg(short, long)]
|
||||
pub model: Option<String>,
|
||||
|
||||
/// System prompt
|
||||
#[arg(short, long)]
|
||||
pub system: Option<String>,
|
||||
|
||||
/// 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<String>,
|
||||
|
||||
/// 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<String>,
|
||||
|
||||
/// Search for models matching this string
|
||||
#[arg(short, long)]
|
||||
query: Option<String>,
|
||||
},
|
||||
|
||||
/// 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<String> {
|
||||
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<String>, stdin: Option<String>) -> Result<String> {
|
||||
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>) -> (String, Option<String>) {
|
||||
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<GenerateParams> {
|
||||
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<serde_json::Value> = 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<ModelsCommand>) -> 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(())
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue