Skip to main content

kiss_coding/tools/
grep.rs

1//! Grep tool: gitignore-aware content search using ripgrep's libraries
2//! in-process (ignore + grep-searcher), matching pi's grep tool surface.
3
4use grep_matcher::Matcher;
5use grep_regex::RegexMatcherBuilder;
6use grep_searcher::SearcherBuilder;
7use grep_searcher::sinks::UTF8;
8use kiss_agent::tool::{AgentTool, ToolResult, ToolUpdateSink};
9use kiss_agent::tools::path::resolve;
10use kiss_agent::tools::truncate::{
11    DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, GREP_MAX_LINE_LENGTH, truncate_head, truncate_line,
12};
13use serde_json::{Value, json};
14use std::path::PathBuf;
15use std::sync::Mutex;
16use tokio_util::sync::CancellationToken;
17
18const DEFAULT_LIMIT: usize = 100;
19const PARALLEL_GREP_MIN_FILES: usize = 256;
20const MAX_GREP_WORKERS: usize = 4;
21
22struct GrepRecord {
23    path: String,
24    line_number: u64,
25    text: String,
26    is_match: bool,
27    was_truncated: bool,
28}
29
30struct GrepChunk {
31    records: Vec<GrepRecord>,
32}
33
34pub struct GrepTool {
35    pub cwd: PathBuf,
36}
37
38#[async_trait::async_trait]
39impl AgentTool for GrepTool {
40    fn name(&self) -> &str {
41        "grep"
42    }
43
44    fn description(&self) -> String {
45        "Search file contents for a pattern (regex or literal string). Respects .gitignore. Returns matching lines with file paths and line numbers.".to_string()
46    }
47
48    fn parameters(&self) -> Value {
49        json!({
50            "type": "object",
51            "properties": {
52                "pattern": {"type": "string", "description": "Search pattern (regex or literal string)"},
53                "path": {"type": "string", "description": "Directory or file to search (default: current directory)"},
54                "glob": {"type": "string", "description": "Filter files by glob pattern, e.g. '*.rs' or '**/*.spec.ts'"},
55                "ignoreCase": {"type": "boolean", "description": "Case-insensitive search (default: false)"},
56                "literal": {"type": "boolean", "description": "Treat pattern as literal string instead of regex (default: false)"},
57                "context": {"type": "number", "description": "Number of lines to show before and after each match (default: 0)"},
58                "limit": {"type": "number", "description": "Maximum number of matches to return (default: 100)"},
59            },
60            "required": ["pattern"],
61        })
62    }
63
64    async fn execute(
65        &self,
66        _id: &str,
67        args: Value,
68        cancel: CancellationToken,
69        _on_update: Option<ToolUpdateSink>,
70    ) -> anyhow::Result<ToolResult> {
71        let pattern = args["pattern"].as_str().unwrap_or_default().to_string();
72        let search_path = resolve(&self.cwd, args["path"].as_str().unwrap_or("."));
73        let glob = args["glob"].as_str().map(String::from);
74        let ignore_case = args["ignoreCase"].as_bool().unwrap_or(false);
75        let literal = args["literal"].as_bool().unwrap_or(false);
76        let context = args["context"].as_f64().unwrap_or(0.0) as usize;
77        let limit = args["limit"]
78            .as_f64()
79            .map(|v| v as usize)
80            .unwrap_or(DEFAULT_LIMIT);
81        let cwd = self.cwd.clone();
82
83        // File walking + searching is sync CPU/IO work. Run it off the async
84        // executor so long searches never stall the event loop.
85        let result = tokio::task::spawn_blocking(move || {
86            run_grep(
87                &cwd,
88                &search_path,
89                &pattern,
90                glob.as_deref(),
91                ignore_case,
92                literal,
93                context,
94                limit,
95                cancel,
96            )
97        })
98        .await??;
99        Ok(result)
100    }
101}
102
103#[allow(clippy::too_many_arguments)]
104fn run_grep(
105    cwd: &std::path::Path,
106    search_path: &std::path::Path,
107    pattern: &str,
108    glob: Option<&str>,
109    ignore_case: bool,
110    literal: bool,
111    context: usize,
112    limit: usize,
113    cancel: CancellationToken,
114) -> anyhow::Result<ToolResult> {
115    if !search_path.exists() {
116        anyhow::bail!("Path not found: {}", search_path.display());
117    }
118    let matcher = RegexMatcherBuilder::new()
119        .case_insensitive(ignore_case)
120        .fixed_strings(literal)
121        .build(pattern)
122        .map_err(|e| anyhow::anyhow!("Invalid pattern: {e}"))?;
123
124    let glob_matcher = match glob {
125        Some(g) => {
126            let mut builder = globset::GlobSetBuilder::new();
127            let normalized = if g.contains('/') {
128                g.to_string()
129            } else {
130                format!("**/{g}")
131            };
132            builder.add(
133                globset::Glob::new(&normalized)
134                    .map_err(|e| anyhow::anyhow!("Invalid glob: {e}"))?,
135            );
136            Some(builder.build()?)
137        }
138        None => None,
139    };
140
141    let paths = Mutex::new(Vec::new());
142    let mut walker = ignore::WalkBuilder::new(search_path);
143    walker
144        .hidden(true)
145        .git_ignore(true)
146        .git_global(true)
147        .require_git(false);
148    walker.build_parallel().run(|| {
149        let cancel = cancel.clone();
150        let paths = &paths;
151        let glob_matcher = &glob_matcher;
152        Box::new(move |entry| {
153            if cancel.is_cancelled() {
154                return ignore::WalkState::Quit;
155            }
156            let Ok(entry) = entry else {
157                return ignore::WalkState::Continue;
158            };
159            let path = entry.path();
160            if !entry.file_type().is_some_and(|kind| kind.is_file()) {
161                return ignore::WalkState::Continue;
162            }
163            if let Some(glob_matcher) = glob_matcher {
164                let relative = path.strip_prefix(search_path).unwrap_or(path);
165                if !glob_matcher.is_match(relative) && !glob_matcher.is_match(path) {
166                    return ignore::WalkState::Continue;
167                }
168            }
169            paths.lock().unwrap().push(path.to_path_buf());
170            ignore::WalkState::Continue
171        })
172    });
173
174    let mut paths = paths.into_inner().unwrap();
175    paths.sort_unstable();
176    let worker_count = if paths.len() >= PARALLEL_GREP_MIN_FILES {
177        std::thread::available_parallelism()
178            .map(usize::from)
179            .unwrap_or(1)
180            .min(MAX_GREP_WORKERS)
181            .min(paths.len())
182    } else {
183        1
184    };
185    let chunk_size = paths.len().div_ceil(worker_count.max(1));
186    let chunks = if worker_count <= 1 {
187        vec![search_grep_chunk(
188            &matcher, &paths, cwd, context, limit, &cancel,
189        )]
190    } else {
191        std::thread::scope(|scope| {
192            paths
193                .chunks(chunk_size)
194                .map(|paths| {
195                    scope.spawn(|| search_grep_chunk(&matcher, paths, cwd, context, limit, &cancel))
196                })
197                .collect::<Vec<_>>()
198                .into_iter()
199                .map(|handle| handle.join().expect("grep worker panicked"))
200                .collect::<Vec<_>>()
201        })
202    };
203
204    let mut output_lines = Vec::new();
205    let mut selected_matches = 0usize;
206    let mut lines_truncated = false;
207    'chunks: for chunk in chunks {
208        for record in chunk.records {
209            if record.is_match && selected_matches >= limit {
210                break 'chunks;
211            }
212            let separator = if record.is_match { ':' } else { '-' };
213            output_lines.push(format!(
214                "{}{separator}{}{separator}{}",
215                record.path, record.line_number, record.text
216            ));
217            lines_truncated |= record.was_truncated;
218            if record.is_match {
219                selected_matches += 1;
220                if selected_matches >= limit {
221                    break 'chunks;
222                }
223            }
224        }
225    }
226    if output_lines.is_empty() {
227        return Ok(ToolResult::text("No matches found"));
228    }
229    let joined = output_lines.join("\n");
230    let truncation = truncate_head(&joined, DEFAULT_MAX_LINES, DEFAULT_MAX_BYTES);
231    let mut output = truncation.content.clone();
232    if selected_matches >= limit {
233        output.push_str(&format!(
234            "\n\n[Match limit of {limit} reached. Narrow the pattern or raise limit.]"
235        ));
236    }
237    if truncation.truncated {
238        output.push_str("\n\n[Output truncated. Narrow the search or use a more specific path.]");
239    }
240    let details = json!({
241        "matchLimitReached": if selected_matches >= limit { Some(limit) } else { None },
242        "linesTruncated": lines_truncated,
243        "truncation": if truncation.truncated { Some(&truncation) } else { None },
244    });
245    Ok(ToolResult {
246        content: vec![kiss_ai::ContentBlock::text(output)],
247        details,
248        ..Default::default()
249    })
250}
251
252fn search_grep_chunk(
253    matcher: &grep_regex::RegexMatcher,
254    paths: &[PathBuf],
255    cwd: &std::path::Path,
256    context: usize,
257    limit: usize,
258    cancel: &CancellationToken,
259) -> GrepChunk {
260    let mut records = Vec::new();
261    let mut match_count = 0usize;
262    let mut searcher = SearcherBuilder::new()
263        .line_number(true)
264        .before_context(context)
265        .after_context(context)
266        .build();
267    for path in paths {
268        if cancel.is_cancelled() || match_count >= limit {
269            break;
270        }
271        let display_path = path.strip_prefix(cwd).unwrap_or(path).display().to_string();
272        let _ = searcher.search_path(
273            matcher,
274            path,
275            UTF8(|line_number, line| {
276                if cancel.is_cancelled() || match_count >= limit {
277                    return Ok(false);
278                }
279                let is_match = matcher.is_match(line.as_bytes()).unwrap_or(false);
280                let (text, was_truncated) =
281                    truncate_line(line.trim_end_matches('\n'), GREP_MAX_LINE_LENGTH);
282                records.push(GrepRecord {
283                    path: display_path.clone(),
284                    line_number,
285                    text,
286                    is_match,
287                    was_truncated,
288                });
289                if is_match {
290                    match_count += 1;
291                }
292                Ok(match_count < limit)
293            }),
294        );
295    }
296    GrepChunk { records }
297}
298
299#[cfg(test)]
300mod tests {
301    use super::*;
302
303    fn setup() -> tempfile::TempDir {
304        let dir = tempfile::tempdir().unwrap();
305        std::fs::write(dir.path().join("a.rs"), "fn main() {}\nlet needle = 1;\n").unwrap();
306        std::fs::write(dir.path().join("b.txt"), "needle here too\n").unwrap();
307        std::fs::create_dir_all(dir.path().join("skip")).unwrap();
308        std::fs::write(dir.path().join(".gitignore"), "skip/\n").unwrap();
309        std::fs::write(dir.path().join("skip/c.rs"), "needle ignored\n").unwrap();
310        dir
311    }
312
313    #[tokio::test]
314    async fn finds_matches_respecting_gitignore() {
315        let dir = setup();
316        let tool = GrepTool {
317            cwd: dir.path().to_path_buf(),
318        };
319        let r = tool
320            .execute(
321                "1",
322                json!({"pattern": "needle"}),
323                CancellationToken::new(),
324                None,
325            )
326            .await
327            .unwrap();
328        let text = r.output_text();
329        assert!(text.contains("a.rs:2:"));
330        assert!(text.contains("b.txt:1:"));
331        assert!(!text.contains("ignored"));
332    }
333
334    #[tokio::test]
335    async fn glob_filter_and_literal() {
336        let dir = setup();
337        let tool = GrepTool {
338            cwd: dir.path().to_path_buf(),
339        };
340        let r = tool
341            .execute(
342                "1",
343                json!({"pattern": "needle", "glob": "*.rs", "literal": true}),
344                CancellationToken::new(),
345                None,
346            )
347            .await
348            .unwrap();
349        let text = r.output_text();
350        assert!(text.contains("a.rs"));
351        assert!(!text.contains("b.txt"));
352    }
353
354    #[tokio::test]
355    async fn no_matches_message() {
356        let dir = setup();
357        let tool = GrepTool {
358            cwd: dir.path().to_path_buf(),
359        };
360        let r = tool
361            .execute(
362                "1",
363                json!({"pattern": "zzz_absent"}),
364                CancellationToken::new(),
365                None,
366            )
367            .await
368            .unwrap();
369        assert_eq!(r.output_text(), "No matches found");
370    }
371
372    #[test]
373    fn parallel_limit_selects_the_same_sorted_files() {
374        let dir = tempfile::tempdir().unwrap();
375        for index in (0..300).rev() {
376            std::fs::write(dir.path().join(format!("file_{index:03}.rs")), "needle\n").unwrap();
377        }
378        let mut outputs = Vec::new();
379        for _ in 0..4 {
380            outputs.push(
381                run_grep(
382                    dir.path(),
383                    dir.path(),
384                    "needle",
385                    Some("*.rs"),
386                    false,
387                    true,
388                    0,
389                    10,
390                    CancellationToken::new(),
391                )
392                .unwrap()
393                .output_text(),
394            );
395        }
396
397        assert!(outputs.windows(2).all(|pair| pair[0] == pair[1]));
398        for index in 0..10 {
399            assert!(outputs[0].contains(&format!("file_{index:03}.rs:1:needle")));
400        }
401        assert!(!outputs[0].contains("file_010.rs:1:needle"));
402    }
403
404    #[test]
405    #[ignore = "release-mode performance benchmark"]
406    fn benchmark_performance_grep_tree() {
407        let dir = tempfile::tempdir().unwrap();
408        for directory in 0..20 {
409            let path = dir.path().join(format!("src/module_{directory:02}"));
410            std::fs::create_dir_all(&path).unwrap();
411            for file in 0..50 {
412                let marker = if file % 5 == 0 { "needle" } else { "ordinary" };
413                std::fs::write(
414                    path.join(format!("file_{file:03}.rs")),
415                    format!("fn item_{file}() {{}}\nlet value = \"{marker}\";\n"),
416                )
417                .unwrap();
418            }
419        }
420        kiss_bench::measure("grep_tree_1000", 11, 1, "1000_files_200_matches", || {
421            run_grep(
422                dir.path(),
423                dir.path(),
424                "needle",
425                Some("*.rs"),
426                false,
427                true,
428                0,
429                10_000,
430                CancellationToken::new(),
431            )
432            .unwrap()
433            .output_text()
434            .len()
435        });
436    }
437}