1use saya_agent::ApprovalPolicy;
2use std::{fmt, str::FromStr};
3
4const 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
102fn 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
121fn 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
162pub 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
206pub 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}