Skip to main content

saya_cli/
slash.rs

1use saya_agent::ApprovalPolicy;
2use std::{fmt, str::FromStr};
3
4/// Known slash command names handled by `parse_slash_command`.
5const KNOWN_COMMANDS: &[&str] = &[
6    "connect",
7    "connections",
8    "include",
9    "exclude",
10    "provider",
11    "model",
12    "privacy",
13    "approvals",
14    "schema",
15    "sql",
16    "clear",
17    "history",
18    "sessions",
19    "resume",
20    "help",
21    "exit",
22    "quit",
23];
24
25#[derive(Debug, Clone, PartialEq, Eq)]
26pub enum SlashCommand {
27    Connect(String),
28    Connections,
29    Include(String),
30    Exclude(String),
31    Provider(Option<String>),
32    Model(Option<String>),
33    Privacy(Option<bool>),
34    Approvals(Option<ApprovalPolicy>),
35    Schema(bool),
36    Sql(String),
37    Clear,
38    History,
39    Sessions,
40    Resume(String),
41    Help(Option<String>),
42    Exit,
43}
44
45#[derive(Debug, Clone, PartialEq, Eq)]
46pub struct SlashParseError(pub String);
47
48impl fmt::Display for SlashParseError {
49    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
50        f.write_str(&self.0)
51    }
52}
53impl std::error::Error for SlashParseError {}
54
55pub fn parse_slash_command(input: &str) -> Result<Option<SlashCommand>, SlashParseError> {
56    let trimmed = input.trim();
57    if !trimmed.starts_with('/') {
58        return Ok(None);
59    }
60    let mut parts = trimmed[1..].split_whitespace();
61    let name = parts.next().unwrap_or_default();
62    let arg = parts.collect::<Vec<_>>().join(" ");
63    let required = || {
64        (!arg.is_empty())
65            .then_some(arg.clone())
66            .ok_or_else(|| SlashParseError("command requires an argument".into()))
67    };
68    let command = match name {
69        "connect" => SlashCommand::Connect(required()?),
70        "connections" => SlashCommand::Connections,
71        "include" => SlashCommand::Include(required()?),
72        "exclude" => SlashCommand::Exclude(required()?),
73        "provider" => SlashCommand::Provider((!arg.is_empty()).then_some(arg)),
74        "model" => SlashCommand::Model((!arg.is_empty()).then_some(arg)),
75        "privacy" => SlashCommand::Privacy(parse_bool(&arg)?),
76        "approvals" => SlashCommand::Approvals(parse_approval(&arg)?),
77        "schema" => SlashCommand::Schema(arg == "refresh"),
78        "sql" => {
79            let query = trimmed.strip_prefix("/sql").unwrap_or("").trim();
80            if query.is_empty() {
81                return Err(SlashParseError("sql requires a query".into()));
82            }
83            SlashCommand::Sql(query.to_string())
84        }
85        "clear" => SlashCommand::Clear,
86        "history" => SlashCommand::History,
87        "sessions" => SlashCommand::Sessions,
88        "resume" => SlashCommand::Resume(required()?),
89        "help" => SlashCommand::Help((!arg.is_empty()).then_some(arg)),
90        "exit" | "quit" => SlashCommand::Exit,
91        other => {
92            let msg = match closest_command(other) {
93                Some(sugg) => format!("unknown command: /{other} (did you mean /{sugg}?)"),
94                None => format!("unknown command: /{other}"),
95            };
96            return Err(SlashParseError(msg));
97        }
98    };
99    Ok(Some(command))
100}
101
102/// Calculates the Levenshtein edit distance between two strings using a single rolling row.
103fn levenshtein(a: &str, b: &str) -> usize {
104    let b_chars: Vec<char> = b.chars().collect();
105    let mut row: Vec<usize> = (0..=b_chars.len()).collect();
106
107    for (i, ca) in a.chars().enumerate() {
108        let mut prev = row[0];
109        row[0] = i + 1;
110        for (j, &cb) in b_chars.iter().enumerate() {
111            let old_row_j_plus_1 = row[j + 1];
112            let cost = if ca == cb { 0 } else { 1 };
113            row[j + 1] = (prev + cost).min(row[j] + 1).min(old_row_j_plus_1 + 1);
114            prev = old_row_j_plus_1;
115        }
116    }
117
118    row.last().copied().unwrap_or(0)
119}
120
121/// Returns the known command with the smallest Levenshtein distance to `input` if distance <= 2.
122fn closest_command(input: &str) -> Option<&'static str> {
123    let input_lower = input.to_lowercase();
124    let mut best_cmd = None;
125    let mut min_dist = usize::MAX;
126
127    for &cmd in KNOWN_COMMANDS {
128        let dist = levenshtein(&input_lower, cmd);
129        if dist < min_dist {
130            min_dist = dist;
131            best_cmd = Some(cmd);
132        }
133    }
134
135    if min_dist <= 2 { best_cmd } else { None }
136}
137
138fn parse_bool(value: &str) -> Result<Option<bool>, SlashParseError> {
139    if value.is_empty() {
140        return Ok(None);
141    }
142    match value {
143        "on" | "true" | "enable" => Ok(Some(true)),
144        "off" | "false" | "disable" => Ok(Some(false)),
145        _ => Err(SlashParseError("privacy expects on or off".into())),
146    }
147}
148
149fn parse_approval(value: &str) -> Result<Option<ApprovalPolicy>, SlashParseError> {
150    if value.is_empty() {
151        return Ok(None);
152    }
153    ApprovalPolicy::from_str(value)
154        .map(Some)
155        .map_err(|error| SlashParseError(error.to_string()))
156}
157
158pub fn help_text() -> &'static str {
159    "/connect <profile>  /connections  /include <profile>  /exclude <profile>\n/provider [name]     /model [name]  /privacy [on|off]\n/approvals [ask|read-only|never]  /schema [refresh]  /sql <query>  /clear\n/history  /sessions  /resume <id>  /help  /exit"
160}
161
162/// Returns a short usage and example string for a known slash command, or `None` if unknown.
163pub fn command_help(name: &str) -> Option<&'static str> {
164    let clean_name = name.trim_start_matches('/').to_lowercase();
165    match clean_name.as_str() {
166        "connect" => {
167            Some("connect <profile> — set the active database profile. Example: /connect prod")
168        }
169        "connections" => Some(
170            "connections — list configured database connection profiles. Example: /connections",
171        ),
172        "include" => Some(
173            "include <profile> — include an additional database profile. Example: /include staging",
174        ),
175        "exclude" => {
176            Some("exclude <profile> — exclude a database profile. Example: /exclude staging")
177        }
178        "provider" => {
179            Some("provider [name] — view or set the AI provider. Example: /provider anthropic")
180        }
181        "model" => Some("model [name] — view or set the AI model. Example: /model gpt-4o"),
182        "privacy" => {
183            Some("privacy [on|off] — view or toggle cloud data sharing. Example: /privacy off")
184        }
185        "approvals" => Some(
186            "approvals [ask|read-only|never] — view or set tool execution approval policy. Example: /approvals ask",
187        ),
188        "schema" => Some(
189            "schema [refresh] — display or refresh database schema context. Example: /schema refresh",
190        ),
191        "sql" => Some(
192            "sql <query> — execute a raw SQL query directly. Example: /sql SELECT * FROM users LIMIT 10;",
193        ),
194        "clear" => Some("clear — clear conversation history and context. Example: /clear"),
195        "history" => Some("history — display session history. Example: /history"),
196        "sessions" => Some("sessions — list available interactive sessions. Example: /sessions"),
197        "resume" => Some("resume <id> — resume a previous session by ID. Example: /resume 12345"),
198        "help" => Some(
199            "help [command] — display general help or detailed usage for a command. Example: /help connect",
200        ),
201        "exit" | "quit" => Some("exit — exit the interactive CLI session. Example: /exit"),
202        _ => None,
203    }
204}
205
206/// Returns command-specific help for a topic, or general help text if `topic` is `None`.
207pub fn help_for(topic: Option<&str>) -> String {
208    match topic {
209        Some(name) => {
210            let clean = name.trim_start_matches('/');
211            match command_help(clean) {
212                Some(help) => help.to_string(),
213                None => format!("No help for /{clean}. Type /help for the full list."),
214            }
215        }
216        None => help_text().to_string(),
217    }
218}
219
220#[cfg(test)]
221mod tests {
222    use super::*;
223
224    #[test]
225    fn test_help_command() {
226        assert_eq!(
227            parse_slash_command("/help"),
228            Ok(Some(SlashCommand::Help(None)))
229        );
230        assert_eq!(
231            parse_slash_command("/help connect"),
232            Ok(Some(SlashCommand::Help(Some("connect".into()))))
233        );
234
235        let help_connect = help_for(Some("connect"));
236        assert!(help_connect.contains("connect"));
237        assert!(help_connect.contains("Example"));
238
239        let help_unknown = help_for(Some("nope"));
240        assert!(help_unknown.contains("No help"));
241
242        assert_eq!(help_for(None), help_text().to_string());
243    }
244
245    #[test]
246    fn test_parse_sessions_and_resume() {
247        assert_eq!(
248            parse_slash_command("/sessions"),
249            Ok(Some(SlashCommand::Sessions))
250        );
251        assert_eq!(
252            parse_slash_command("/resume 12345"),
253            Ok(Some(SlashCommand::Resume("12345".into())))
254        );
255        assert_eq!(
256            parse_slash_command("/resume"),
257            Err(SlashParseError("command requires an argument".into()))
258        );
259    }
260
261    #[test]
262    fn test_parse_sql_command() {
263        assert_eq!(
264            parse_slash_command("/sql SELECT * FROM users;"),
265            Ok(Some(SlashCommand::Sql("SELECT * FROM users;".into())))
266        );
267        assert_eq!(
268            parse_slash_command("/sql   SELECT  a,  b  FROM  table  "),
269            Ok(Some(SlashCommand::Sql("SELECT  a,  b  FROM  table".into())))
270        );
271        assert_eq!(
272            parse_slash_command("/sql"),
273            Err(SlashParseError("sql requires a query".into()))
274        );
275        assert_eq!(
276            parse_slash_command("/sql   "),
277            Err(SlashParseError("sql requires a query".into()))
278        );
279    }
280
281    #[test]
282    fn test_unknown_command_suggestion() {
283        let err = parse_slash_command("/conect prod").unwrap_err();
284        assert!(
285            err.0.contains("did you mean /connect"),
286            "expected suggestion in error message, got: {}",
287            err.0
288        );
289
290        let err = parse_slash_command("/zzzzzzzz").unwrap_err();
291        assert!(
292            !err.0.contains("did you mean"),
293            "unexpected suggestion in error message, got: {}",
294            err.0
295        );
296
297        assert_eq!(
298            parse_slash_command("/connect prod"),
299            Ok(Some(SlashCommand::Connect("prod".into())))
300        );
301    }
302}