Skip to main content

aether_cli/generate_command/
mod.rs

1use crate::output::OutputFormat;
2use futures::StreamExt;
3use llm::catalog::{ReasoningEffortError, validate_reasoning_effort};
4use llm::parser::ModelProviderParser;
5use llm::{ChatMessage, Context, LlmError, LlmResponse, ModelSettings, ReasoningEffort, StreamingModelProvider};
6use std::io::{Read, stdin};
7use std::process::ExitCode;
8use thiserror::Error;
9
10#[derive(clap::Args)]
11pub struct GenerateArgs {
12    /// Model to call, as `provider:model` (e.g. `anthropic:claude-sonnet-4-5`).
13    #[arg(long)]
14    pub model: String,
15
16    /// Prompt text to send.
17    #[arg(long, conflicts_with = "prompt_file", required_unless_present = "prompt_file")]
18    pub prompt: Option<String>,
19
20    /// Prompt to send from a file path, or `-` to read from stdin.
21    #[arg(long, value_name = "PATH_OR_DASH", conflicts_with = "prompt", required_unless_present = "prompt")]
22    pub prompt_file: Option<String>,
23
24    /// Optional system prompt.
25    #[arg(long)]
26    pub system: Option<String>,
27
28    /// Sampling controls (temperature, top-p, max-tokens) for the model call.
29    #[command(flatten)]
30    pub model_settings: ModelSettingsArgs,
31
32    /// Reasoning effort for models that support extended thinking
33    /// (`minimal`, `low`, `medium`, `high`, `xhigh`, `max`).
34    #[arg(long)]
35    pub reasoning_effort: Option<ReasoningEffort>,
36
37    /// Output format. `text` prints the raw response; `json`/`pretty` wrap it as
38    /// `{ "text": <response>, "model": <model> }`.
39    #[arg(long, default_value = "text")]
40    pub output: OutputFormat,
41}
42
43#[derive(Debug, Clone, Default, clap::Args)]
44pub struct ModelSettingsArgs {
45    /// Sampling temperature. Lower is more deterministic (e.g. `0` for grading).
46    #[arg(long)]
47    pub temperature: Option<f32>,
48    /// Nucleus sampling: the probability mass to sample from.
49    #[arg(long)]
50    pub top_p: Option<f32>,
51    /// Upper bound on the number of tokens generated in the response.
52    #[arg(long)]
53    pub max_tokens: Option<u32>,
54}
55
56impl From<ModelSettingsArgs> for ModelSettings {
57    fn from(args: ModelSettingsArgs) -> Self {
58        ModelSettings { temperature: args.temperature, top_p: args.top_p, max_tokens: args.max_tokens }
59    }
60}
61
62#[derive(Debug, Error)]
63pub enum GenerateCommandError {
64    #[error("provide exactly one of --prompt or --prompt-file")]
65    PromptSource,
66
67    #[error("failed to read prompt from {path}: {source}")]
68    ReadPrompt { path: String, source: std::io::Error },
69
70    #[error("invalid reasoning effort: {0}")]
71    ReasoningEffort(#[from] ReasoningEffortError),
72
73    #[error("failed to initialize model `{model}`: {source}")]
74    Model { model: String, source: LlmError },
75
76    #[error("model stream error: {0}")]
77    Stream(LlmError),
78}
79
80/// Call a model with a single prompt and print its response. The judge is a special case of this:
81/// the caller supplies a grading prompt and parses the structured verdict out of the response.
82pub async fn run(args: GenerateArgs) -> Result<ExitCode, GenerateCommandError> {
83    let prompt = resolve_prompt(args.prompt.as_deref(), args.prompt_file.as_deref())?;
84    validate_reasoning_effort(&args.model, args.reasoning_effort)?;
85    let (provider, _) = ModelProviderParser::default()
86        .parse(&args.model)
87        .await
88        .map_err(|source| GenerateCommandError::Model { model: args.model.clone(), source })?;
89
90    let messages = {
91        let mut messages = Vec::new();
92        if let Some(system) = args.system.as_deref() {
93            messages.push(ChatMessage::system(system));
94        }
95        messages.push(ChatMessage::user(&prompt));
96        messages
97    };
98
99    let mut context = Context::new(messages, vec![]);
100    context.set_model_settings(args.model_settings.clone().into());
101    context.set_reasoning_effort(args.reasoning_effort);
102
103    let mut stream = provider.stream_response(&context);
104    let mut text = String::new();
105    while let Some(result) = stream.next().await {
106        match result {
107            Ok(LlmResponse::Text { chunk }) => text.push_str(&chunk),
108            Err(error) => return Err(GenerateCommandError::Stream(error)),
109            _ => {}
110        }
111    }
112    println!("{}", format_output(&text, &args.model, args.output));
113    Ok(ExitCode::SUCCESS)
114}
115
116fn format_output(text: &str, model: &str, format: OutputFormat) -> String {
117    match format {
118        OutputFormat::Text => text.to_string(),
119        OutputFormat::Json => serde_json::json!({ "text": text, "model": model }).to_string(),
120        OutputFormat::Pretty => serde_json::to_string_pretty(&serde_json::json!({ "text": text, "model": model }))
121            .expect("generate response serializes to JSON"),
122    }
123}
124
125fn resolve_prompt(prompt: Option<&str>, prompt_file: Option<&str>) -> Result<String, GenerateCommandError> {
126    match (prompt, prompt_file) {
127        (Some(prompt), None) => Ok(prompt.to_string()),
128        (None, Some(prompt_file)) => read_prompt_file(prompt_file),
129        _ => Err(GenerateCommandError::PromptSource),
130    }
131}
132
133fn read_prompt_file(source: &str) -> Result<String, GenerateCommandError> {
134    if source == "-" {
135        let mut prompt = String::new();
136        stdin()
137            .read_to_string(&mut prompt)
138            .map_err(|error| GenerateCommandError::ReadPrompt { path: "-".to_string(), source: error })?;
139        return Ok(prompt);
140    }
141    std::fs::read_to_string(source)
142        .map_err(|error| GenerateCommandError::ReadPrompt { path: source.to_string(), source: error })
143}
144
145#[cfg(test)]
146mod tests {
147    use super::*;
148    use clap::Parser;
149
150    #[test]
151    fn format_output_wraps_json_and_passes_text_through() {
152        assert_eq!(format_output("hi", "anthropic:m", OutputFormat::Text), "hi");
153
154        let json: serde_json::Value =
155            serde_json::from_str(&format_output("hi", "anthropic:m", OutputFormat::Json)).unwrap();
156        assert_eq!(json["text"], "hi");
157        assert_eq!(json["model"], "anthropic:m");
158    }
159
160    #[tokio::test]
161    async fn run_rejects_reasoning_effort_unsupported_by_model_before_initializing_provider() {
162        let args = GenerateArgs {
163            model: "anthropic:claude-opus-4-6".to_string(),
164            prompt: Some("the prompt".to_string()),
165            prompt_file: None,
166            system: None,
167            model_settings: ModelSettingsArgs::default(),
168            reasoning_effort: Some(ReasoningEffort::Xhigh),
169            output: OutputFormat::Text,
170        };
171
172        let error = run(args).await.unwrap_err();
173
174        assert!(matches!(error, GenerateCommandError::ReasoningEffort(_)));
175    }
176
177    #[tokio::test]
178    async fn run_errors_on_unknown_provider() {
179        let dir = tempfile::tempdir().unwrap();
180        let prompt_path = dir.path().join("prompt.txt");
181        std::fs::write(&prompt_path, "the prompt").unwrap();
182        let args = GenerateArgs {
183            model: "definitely-not-a-provider:nope".to_string(),
184            prompt: None,
185            prompt_file: Some(prompt_path.to_string_lossy().into_owned()),
186            system: None,
187            model_settings: ModelSettingsArgs::default(),
188            reasoning_effort: None,
189            output: OutputFormat::Json,
190        };
191
192        let error = run(args).await.unwrap_err();
193
194        assert!(matches!(error, GenerateCommandError::Model { .. }), "got: {error:?}");
195    }
196
197    #[derive(clap::Parser)]
198    struct TestCli {
199        #[command(flatten)]
200        args: GenerateArgs,
201    }
202
203    fn parse_args(argv: &[&str]) -> GenerateArgs {
204        TestCli::try_parse_from(argv).unwrap().args
205    }
206
207    #[test]
208    fn model_settings_and_reasoning_flags_parse_convert_and_validate() {
209        let set = parse_args(&[
210            "gen",
211            "--model",
212            "anthropic:m",
213            "--prompt",
214            "hi",
215            "--temperature",
216            "0",
217            "--top-p",
218            "0.5",
219            "--max-tokens",
220            "64",
221            "--reasoning-effort",
222            "high",
223        ]);
224        assert_eq!(
225            ModelSettings::from(set.model_settings),
226            ModelSettings { temperature: Some(0.0), top_p: Some(0.5), max_tokens: Some(64) }
227        );
228        assert_eq!(set.reasoning_effort, Some(ReasoningEffort::High));
229
230        let absent = parse_args(&["gen", "--model", "anthropic:m", "--prompt", "hi"]);
231        assert!(ModelSettings::from(absent.model_settings).is_empty());
232        assert_eq!(absent.reasoning_effort, None);
233
234        let bad =
235            TestCli::try_parse_from(["gen", "--model", "anthropic:m", "--prompt", "hi", "--reasoning-effort", "nope"]);
236        assert!(bad.is_err());
237    }
238
239    #[test]
240    fn resolve_prompt_uses_inline_prompt() {
241        assert_eq!(resolve_prompt(Some("say hi"), None).unwrap(), "say hi");
242    }
243
244    #[test]
245    fn resolve_prompt_reads_a_file() {
246        let dir = tempfile::tempdir().unwrap();
247        let path = dir.path().join("prompt.txt");
248        std::fs::write(&path, "graded prompt").unwrap();
249
250        assert_eq!(resolve_prompt(None, Some(&path.to_string_lossy())).unwrap(), "graded prompt");
251    }
252
253    #[test]
254    fn resolve_prompt_reports_missing_file() {
255        let error = resolve_prompt(None, Some("/nonexistent/prompt.txt")).unwrap_err();
256
257        assert!(matches!(error, GenerateCommandError::ReadPrompt { .. }), "got: {error:?}");
258    }
259
260    #[test]
261    fn resolve_prompt_rejects_missing_prompt_source() {
262        let error = resolve_prompt(None, None).unwrap_err();
263
264        assert!(matches!(error, GenerateCommandError::PromptSource), "got: {error:?}");
265    }
266
267    #[test]
268    fn resolve_prompt_rejects_multiple_prompt_sources() {
269        let error = resolve_prompt(Some("inline"), Some("prompt.txt")).unwrap_err();
270
271        assert!(matches!(error, GenerateCommandError::PromptSource), "got: {error:?}");
272    }
273}