aether_cli/generate_command/
mod.rs1use 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 #[arg(long)]
14 pub model: String,
15
16 #[arg(long, conflicts_with = "prompt_file", required_unless_present = "prompt_file")]
18 pub prompt: Option<String>,
19
20 #[arg(long, value_name = "PATH_OR_DASH", conflicts_with = "prompt", required_unless_present = "prompt")]
22 pub prompt_file: Option<String>,
23
24 #[arg(long)]
26 pub system: Option<String>,
27
28 #[command(flatten)]
30 pub model_settings: ModelSettingsArgs,
31
32 #[arg(long)]
35 pub reasoning_effort: Option<ReasoningEffort>,
36
37 #[arg(long, default_value = "text")]
40 pub output: OutputFormat,
41}
42
43#[derive(Debug, Clone, Default, clap::Args)]
44pub struct ModelSettingsArgs {
45 #[arg(long)]
47 pub temperature: Option<f32>,
48 #[arg(long)]
50 pub top_p: Option<f32>,
51 #[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
80pub 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}