1use std::collections::{BTreeMap, BTreeSet};
17use std::sync::OnceLock;
18
19use sha2::{Digest, Sha256};
20
21use crate::cst::{self, Cmd, Script, SimpleCmd, Word, WordPart};
22use crate::registry;
23
24const DEFAULT_LEVEL: &str = "SafeWrite";
28
29#[derive(Debug, Clone, PartialEq, Eq)]
31pub struct GeneratedEntry {
32 pub name: String,
33 pub standalone: Vec<String>,
35 pub max_positional: usize,
37 pub level: String,
38}
39
40#[derive(Debug, Clone, PartialEq, Eq)]
42pub enum Outcome {
43 AlreadyAllowed,
45 Unparseable,
47 RecognizedButDenied { names: Vec<String> },
50 Generated { entries: Vec<GeneratedEntry>, also_recognized: Vec<String> },
53}
54
55pub fn analyze(command: &str) -> Outcome {
57 if cst::command_verdict(command).is_allowed() {
58 return Outcome::AlreadyAllowed;
59 }
60 let Some(script) = cst::parse(command) else {
61 return Outcome::Unparseable;
62 };
63
64 let mut simples: Vec<&SimpleCmd> = Vec::new();
65 collect_script(&script, &mut simples);
66
67 let mut unknown: BTreeMap<String, (BTreeSet<String>, usize)> = BTreeMap::new();
69 let mut recognized: BTreeSet<String> = BTreeSet::new();
70 for sc in simples {
71 let Some(name) = command_basename(sc) else {
72 continue;
73 };
74 if is_known(&name) {
75 recognized.insert(name);
76 continue;
77 }
78 let (flags, positionals) = observed_shape(sc);
79 let entry = unknown.entry(name).or_default();
80 entry.0.extend(flags);
81 entry.1 = entry.1.max(positionals);
82 }
83
84 if unknown.is_empty() {
85 return Outcome::RecognizedButDenied { names: recognized.into_iter().collect() };
86 }
87 let entries = unknown
88 .into_iter()
89 .map(|(name, (flags, max_positional))| GeneratedEntry {
90 name,
91 standalone: flags.into_iter().collect(),
92 max_positional,
93 level: DEFAULT_LEVEL.to_string(),
94 })
95 .collect();
96 Outcome::Generated { entries, also_recognized: recognized.into_iter().collect() }
97}
98
99fn command_basename(sc: &SimpleCmd) -> Option<String> {
102 let raw = sc.words.first()?.eval();
103 if raw.is_empty() {
104 return None;
105 }
106 Some(crate::parse::Token::from_raw(raw).command_name().to_string())
107}
108
109fn observed_shape(sc: &SimpleCmd) -> (Vec<String>, usize) {
113 let mut flags = Vec::new();
114 let mut positionals = 0;
115 for word in sc.words.iter().skip(1) {
116 let s = word.eval();
117 if s.starts_with('-') && s != "-" {
118 flags.push(s);
119 } else {
120 positionals += 1;
121 }
122 }
123 (flags, positionals)
124}
125
126fn collect_script<'a>(script: &'a Script, out: &mut Vec<&'a SimpleCmd>) {
127 for stmt in &script.0 {
128 for cmd in &stmt.pipeline.commands {
129 collect_cmd(cmd, out);
130 }
131 }
132}
133
134fn collect_cmd<'a>(cmd: &'a Cmd, out: &mut Vec<&'a SimpleCmd>) {
135 match cmd {
136 Cmd::Simple(sc) => {
137 out.push(sc);
138 for word in &sc.words {
139 collect_word(word, out);
140 }
141 }
142 Cmd::Subshell { body, .. } | Cmd::BraceGroup { body, .. } => collect_script(body, out),
143 Cmd::For { items, body, .. } => {
144 for word in items {
145 collect_word(word, out);
146 }
147 collect_script(body, out);
148 }
149 Cmd::While { cond, body, .. } | Cmd::Until { cond, body, .. } => {
150 collect_script(cond, out);
151 collect_script(body, out);
152 }
153 Cmd::If { branches, else_body, .. } => {
154 for branch in branches {
155 collect_script(&branch.cond, out);
156 collect_script(&branch.body, out);
157 }
158 if let Some(body) = else_body {
159 collect_script(body, out);
160 }
161 }
162 Cmd::DoubleBracket { words, .. } => {
163 for word in words {
164 collect_word(word, out);
165 }
166 }
167 Cmd::Case { subject, arms, .. } => {
168 collect_word(subject, out);
169 for arm in arms {
170 collect_script(&arm.body, out);
171 }
172 }
173 Cmd::FunctionDef { body, .. } => collect_script(body, out),
174 }
175}
176
177fn collect_word<'a>(word: &'a Word, out: &mut Vec<&'a SimpleCmd>) {
183 for part in &word.0 {
184 match part {
185 WordPart::CmdSub(s) | WordPart::ProcSub(s) => collect_script(s, out),
186 WordPart::DQuote(w) => collect_word(w, out),
187 _ => {}
188 }
189 }
190}
191
192fn is_known(name: &str) -> bool {
196 known_names().contains(registry::canonical_name(name))
197}
198
199fn known_names() -> &'static BTreeSet<String> {
200 static KNOWN: OnceLock<BTreeSet<String>> = OnceLock::new();
201 KNOWN.get_or_init(|| {
202 let mut set: BTreeSet<String> = crate::docs::all_command_docs().into_iter().map(|d| d.name).collect();
203 for name in registry::toml_command_names() {
204 set.insert(name.to_string());
205 }
206 set
207 })
208}
209
210pub fn render_toml(entries: &[GeneratedEntry]) -> String {
213 let mut out = String::new();
214 for (i, entry) in entries.iter().enumerate() {
215 if i > 0 {
216 out.push('\n');
217 }
218 out.push_str("[[command]]\n");
219 out.push_str(&format!("name = {}\n", toml_str(&entry.name)));
220 if !entry.standalone.is_empty() {
221 let items: Vec<String> = entry.standalone.iter().map(|f| toml_str(f)).collect();
222 out.push_str(&format!("standalone = [{}]\n", items.join(", ")));
223 }
224 out.push_str(&format!("max_positional = {}\n", entry.max_positional));
225 out.push_str(&format!("level = {}\n", toml_str(&entry.level)));
226 }
227 out
228}
229
230fn toml_str(s: &str) -> String {
233 let mut out = String::from("\"");
234 for c in s.chars() {
235 match c {
236 '"' => out.push_str("\\\""),
237 '\\' => out.push_str("\\\\"),
238 '\n' => out.push_str("\\n"),
239 '\r' => out.push_str("\\r"),
240 '\t' => out.push_str("\\t"),
241 c if (c as u32) < 0x20 || c == '\u{7f}' => {
242 out.push_str(&format!("\\u{:04X}", c as u32));
243 }
244 c => out.push(c),
245 }
246 }
247 out.push('"');
248 out
249}
250
251pub fn config_hash(bytes: &[u8]) -> String {
254 Sha256::digest(bytes).iter().map(|b| format!("{b:02x}")).collect()
255}
256
257pub fn merged_content(existing: &str, entries: &[GeneratedEntry]) -> String {
260 let block = render_toml(entries);
261 if existing.trim().is_empty() {
262 return block;
263 }
264 let mut content = existing.to_string();
265 if !content.ends_with('\n') {
266 content.push('\n');
267 }
268 content.push('\n');
269 content.push_str(&block);
270 content
271}
272
273pub fn pin_block(canonical_dir: &str, hash: &str) -> String {
275 format!("[[trusted]]\npath = {}\nsha256 = {}\n", toml_str(canonical_dir), toml_str(hash),)
276}
277
278#[cfg(test)]
279mod tests;