1use std::io::{self, Write};
4
5use anyhow::Result;
6use clap::Parser;
7use crossterm::{
8 event::{self, Event, KeyCode, KeyModifiers},
9 terminal::enable_raw_mode,
10};
11
12use crate::utils::terminal::RawModeGuard;
13
14#[derive(Parser)]
20pub struct ChatCommand {}
21
22impl ChatCommand {
23 pub async fn execute(self) -> Result<()> {
25 let ai_info = crate::utils::preflight::check_ai_credentials(None)?;
26 eprintln!(
27 "Connected to {} (model: {})",
28 ai_info.provider, ai_info.model
29 );
30 eprintln!("Enter to send, Shift+Enter for newline, Ctrl+D to exit.\n");
31
32 let client = crate::claude::create_default_claude_client(None, None).await?;
33
34 chat_loop(&client).await
35 }
36}
37
38pub async fn run_chat(
49 message: &str,
50 model: Option<String>,
51 system_prompt: Option<String>,
52) -> Result<String> {
53 crate::utils::preflight::check_ai_credentials(model.as_deref())?;
54 let client = crate::claude::create_default_claude_client(model, None).await?;
55 let system = system_prompt
56 .as_deref()
57 .unwrap_or("You are a helpful assistant.");
58 client.send_message(system, message).await
59}
60
61async fn chat_loop(client: &crate::claude::client::ClaudeClient) -> Result<()> {
62 let system_prompt = "You are a helpful assistant.";
63
64 loop {
65 let input = match read_user_input() {
66 Ok(Some(text)) => text,
67 Ok(None) => {
68 eprintln!("\nGoodbye!");
69 break;
70 }
71 Err(e) => {
72 eprintln!("\nInput error: {e}");
73 break;
74 }
75 };
76
77 let trimmed = input.trim();
78 if trimmed.is_empty() {
79 continue;
80 }
81
82 let response = client.send_message(system_prompt, trimmed).await?;
83 println!("{response}\n");
84 }
85
86 Ok(())
87}
88
89fn read_user_input() -> Result<Option<String>> {
93 eprint!("> ");
94 io::stderr().flush()?;
95
96 enable_raw_mode()?;
97 let _guard = RawModeGuard;
98
99 let mut buffer = String::new();
100
101 loop {
102 if let Event::Key(key_event) = event::read()? {
103 match key_event.code {
104 KeyCode::Enter => {
105 if key_event.modifiers.contains(KeyModifiers::SHIFT) {
106 buffer.push('\n');
107 eprint!("\r\n... ");
108 io::stderr().flush()?;
109 } else {
110 eprint!("\r\n");
111 io::stderr().flush()?;
112 return Ok(Some(buffer));
113 }
114 }
115 KeyCode::Char('d') if key_event.modifiers.contains(KeyModifiers::CONTROL) => {
116 if buffer.is_empty() {
117 return Ok(None);
118 }
119 eprint!("\r\n");
120 io::stderr().flush()?;
121 return Ok(Some(buffer));
122 }
123 KeyCode::Char('c') if key_event.modifiers.contains(KeyModifiers::CONTROL) => {
124 return Ok(None);
125 }
126 KeyCode::Char(c) => {
127 buffer.push(c);
128 eprint!("{c}");
129 io::stderr().flush()?;
130 }
131 KeyCode::Backspace if buffer.pop().is_some() => {
132 eprint!("\x08 \x08");
133 io::stderr().flush()?;
134 }
135 _ => {}
136 }
137 }
138 }
139}
140
141#[cfg(test)]
142#[allow(clippy::unwrap_used, clippy::expect_used)]
143mod tests {
144 use super::*;
145
146 static ENV_LOCK: &std::sync::Mutex<()> = &crate::test_support::HOME_ENV_MUTEX;
155
156 const KEYS: &[&str] = &[
157 "OMNI_DEV_AI_BACKEND",
158 "OMNI_DEV_MODEL",
159 "OMNI_DEV_BETA_HEADER",
160 "USE_OPENAI",
161 "USE_OLLAMA",
162 "CLAUDE_CODE_USE_BEDROCK",
163 "CLAUDE_API_KEY",
164 "ANTHROPIC_API_KEY",
165 "ANTHROPIC_AUTH_TOKEN",
166 "ANTHROPIC_BEDROCK_BASE_URL",
167 "OPENAI_API_KEY",
168 "OPENAI_AUTH_TOKEN",
169 "OLLAMA_MODEL",
170 "OLLAMA_BASE_URL",
171 "ANTHROPIC_MODEL",
172 ];
173
174 fn snapshot_env() -> Vec<(&'static str, Option<String>)> {
175 let mut v: Vec<(&'static str, Option<String>)> =
176 KEYS.iter().map(|k| (*k, std::env::var(k).ok())).collect();
177 v.push(("HOME", std::env::var("HOME").ok()));
178 v
179 }
180
181 fn restore_env(snap: Vec<(&'static str, Option<String>)>) {
182 for (k, v) in snap {
183 match v {
184 Some(val) => std::env::set_var(k, val),
185 None => std::env::remove_var(k),
186 }
187 }
188 }
189
190 fn isolate_empty_home() -> tempfile::TempDir {
191 let dir = {
192 std::fs::create_dir_all("tmp").ok();
193 tempfile::TempDir::new_in("tmp").unwrap()
194 };
195 std::env::set_var("HOME", dir.path());
196 for k in KEYS {
197 std::env::remove_var(k);
198 }
199 dir
200 }
201
202 #[allow(clippy::await_holding_lock)]
203 #[tokio::test]
204 async fn run_chat_returns_error_when_credentials_missing() {
205 let _guard = ENV_LOCK
206 .lock()
207 .unwrap_or_else(std::sync::PoisonError::into_inner);
208 let snap = snapshot_env();
209 let _home = isolate_empty_home();
210
211 let err = run_chat("hello", None, None).await.unwrap_err();
212 let msg = format!("{err}");
213 assert!(
214 msg.contains("API key not found") || msg.contains("not found"),
215 "expected credential error, got: {msg}"
216 );
217
218 restore_env(snap);
219 }
220
221 #[allow(clippy::await_holding_lock)]
222 #[tokio::test]
223 async fn run_chat_bubbles_up_credential_error_with_custom_system_prompt() {
224 let _guard = ENV_LOCK
225 .lock()
226 .unwrap_or_else(std::sync::PoisonError::into_inner);
227 let snap = snapshot_env();
228 let _home = isolate_empty_home();
229
230 let err = run_chat("hello", None, Some("be terse".to_string()))
232 .await
233 .unwrap_err();
234 assert!(format!("{err}").contains("not found"));
235
236 restore_env(snap);
237 }
238
239 #[allow(clippy::await_holding_lock)]
240 #[tokio::test]
241 async fn run_chat_propagates_model_override_through_preflight() {
242 let _guard = ENV_LOCK
243 .lock()
244 .unwrap_or_else(std::sync::PoisonError::into_inner);
245 let snap = snapshot_env();
246 let _home = isolate_empty_home();
247
248 let err = run_chat("hello", Some("claude-sonnet-4-6".to_string()), None)
250 .await
251 .unwrap_err();
252 assert!(format!("{err}").contains("not found"));
253
254 restore_env(snap);
255 }
256
257 #[allow(clippy::await_holding_lock)]
262 #[tokio::test]
263 async fn run_chat_happy_path_via_mocked_ollama_returns_response_text() {
264 let _guard = ENV_LOCK
265 .lock()
266 .unwrap_or_else(std::sync::PoisonError::into_inner);
267 let snap = snapshot_env();
268 let _home = isolate_empty_home();
269
270 let server = wiremock::MockServer::start().await;
271 wiremock::Mock::given(wiremock::matchers::method("POST"))
272 .and(wiremock::matchers::path("/v1/chat/completions"))
273 .respond_with(
274 wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({
275 "id": "test",
276 "object": "chat.completion",
277 "choices": [{
278 "index": 0,
279 "message": {"role": "assistant", "content": "canned-response"},
280 "finish_reason": "stop"
281 }]
282 })),
283 )
284 .mount(&server)
285 .await;
286
287 std::env::set_var("USE_OLLAMA", "true");
288 std::env::set_var("OLLAMA_MODEL", "llama2");
289 std::env::set_var("OLLAMA_BASE_URL", server.uri());
290
291 let out = run_chat("hello", None, Some("be terse".to_string()))
292 .await
293 .unwrap();
294 assert_eq!(out, "canned-response");
295
296 restore_env(snap);
297 }
298
299 #[allow(clippy::await_holding_lock)]
302 #[tokio::test]
303 async fn run_chat_default_system_prompt_path_via_mocked_ollama() {
304 let _guard = ENV_LOCK
305 .lock()
306 .unwrap_or_else(std::sync::PoisonError::into_inner);
307 let snap = snapshot_env();
308 let _home = isolate_empty_home();
309
310 let server = wiremock::MockServer::start().await;
311 wiremock::Mock::given(wiremock::matchers::method("POST"))
312 .and(wiremock::matchers::path("/v1/chat/completions"))
313 .respond_with(
314 wiremock::ResponseTemplate::new(200).set_body_json(serde_json::json!({
315 "id": "test",
316 "object": "chat.completion",
317 "choices": [{
318 "index": 0,
319 "message": {"role": "assistant", "content": "ok"},
320 "finish_reason": "stop"
321 }]
322 })),
323 )
324 .mount(&server)
325 .await;
326
327 std::env::set_var("USE_OLLAMA", "true");
328 std::env::set_var("OLLAMA_MODEL", "llama2");
329 std::env::set_var("OLLAMA_BASE_URL", server.uri());
330
331 let out = run_chat("hello", None, None).await.unwrap();
332 assert_eq!(out, "ok");
333
334 restore_env(snap);
335 }
336}