Skip to main content

omni_dev/cli/ai/
chat.rs

1//! Interactive AI chat command.
2
3use 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/// Interactive AI chat session.
15///
16/// Model selection uses the global `--model` flag (propagated as
17/// `OMNI_DEV_MODEL`) and the per-backend env chain; there is no
18/// subcommand-local flag.
19#[derive(Parser)]
20pub struct ChatCommand {}
21
22impl ChatCommand {
23    /// Executes the chat command.
24    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
38/// Sends a single user message to the configured AI and returns the response.
39///
40/// Shared between the MCP `ai_chat` tool and any non-interactive CLI callers.
41/// The function performs the same preflight credential check as the CLI chat
42/// loop — on missing credentials it returns the preflight error verbatim so
43/// MCP tool callers see the same diagnostic message the CLI would print.
44///
45/// `model` selects the AI model; `None` uses the environment default.
46/// `system_prompt` defaults to `"You are a helpful assistant."` (matching the
47/// CLI's default) when `None`.
48pub 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
89/// Reads multiline user input with "> " prompt.
90///
91/// Returns `Ok(Some(text))` on Enter, `Ok(None)` on Ctrl+D/Ctrl+C.
92fn 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    /// Env-isolation lock — tests in this module must serialise because they
147    /// mutate `HOME` and every provider env var the preflight check reads.
148    ///
149    /// Aliases the crate-wide [`crate::test_support::HOME_ENV_MUTEX`] so
150    /// this module's `HOME` mutation also serialises against
151    /// `mcp::ai_tools::tests` (which mutates the same env vars under its own
152    /// separate lock) and every other domain — see that static's doc
153    /// comment (issue #1465).
154    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        // Custom system prompt should not bypass the preflight check.
231        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        // With explicit model override, the same credential check must still run.
249        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    /// Exercises the post-preflight code path (client construction and
258    /// `send_message`) without requiring real AI credentials. Routes through
259    /// Ollama mode, which skips the credential check, and points the client
260    /// at a wiremock server that returns a canned OpenAI-compatible response.
261    #[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    /// As above but with `system_prompt = None`, exercising the
300    /// `.unwrap_or("You are a helpful assistant.")` default branch.
301    #[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}