Skip to main content

navi_core/tool/builtin/
repo_explore.rs

1//! Repository exploration tool — deterministic BM25 + symbol search.
2//!
3//! Fast, non-LLM search over the project index. Combines:
4//! - structured symbol ranking (`ranked_symbol_matches`)
5//! - BM25 text matches over docs/signatures/snippets (`search_text_matches`)
6//!
7//! Returns compact locations (path + line range + snippet + score) so the
8//! parent agent can `read_file` only what matters. No nested model turn.
9
10use std::cmp::Ordering;
11use std::collections::HashMap;
12use std::path::{Path, PathBuf};
13use std::time::Instant;
14
15use anyhow::Result;
16use async_trait::async_trait;
17use serde_json::json;
18
19use super::helpers;
20use crate::repo_intelligence::{
21    RankedSymbolRecord, TextMatchRecord, build_index, ranked_symbol_matches, search_text_matches,
22};
23use crate::tool::{Tool, ToolDefinition, ToolInvocation, ToolKind, ToolResult};
24
25const DEFAULT_MAX_RESULTS: usize = 10;
26const MAX_RESULTS_CAP: usize = 40;
27/// Soft weight so top symbols and BM25 text compete on a shared ranking.
28const SYMBOL_SCORE_SCALE: f64 = 1.0;
29const TEXT_SCORE_SCALE: f64 = 2.5;
30
31pub struct RepoExploreTool {
32    project_dir: PathBuf,
33}
34
35impl RepoExploreTool {
36    pub fn new(project_dir: PathBuf) -> Self {
37        Self { project_dir }
38    }
39}
40
41#[async_trait]
42impl Tool for RepoExploreTool {
43    fn definition(&self) -> ToolDefinition {
44        helpers::definition(
45            "repo_explore",
46            "Fast contextual search over the repository (BM25 + symbol index). \
47             Returns ranked file locations with line ranges and short snippets. \
48             Use this before reading files to find the right places. \
49             Does NOT spawn a subagent or call the model. \
50             Prefer for: \"where is X\", architecture concepts, symbol names, error strings.",
51            ToolKind::Read,
52            json!({
53                "type": "object",
54                "properties": {
55                    "query": {
56                        "type": "string",
57                        "description": "What to find: symbols, paths, concepts, error text, architectural terms."
58                    },
59                    "context": {
60                        "type": "string",
61                        "description": "Optional extra terms to bias ranking (why you need this / related words)."
62                    },
63                    "max_results": {
64                        "type": "integer",
65                        "description": "Maximum locations to return (default 10, max 40)."
66                    },
67                    "kind": {
68                        "type": "string",
69                        "description": "Optional symbol kind filter (function, struct, trait, …). Applied to symbol hits only."
70                    }
71                },
72                "required": ["query"],
73                "additionalProperties": false,
74            }),
75        )
76    }
77
78    async fn invoke(&self, invocation: ToolInvocation) -> Result<ToolResult> {
79        let query = helpers::required_string(&invocation.input, "query")?.to_string();
80        let context = helpers::optional_string(&invocation.input, "context");
81        let kind = helpers::optional_string(&invocation.input, "kind");
82        let max_results = helpers::optional_u64(&invocation.input, "max_results")
83            .map(|v| v as usize)
84            .unwrap_or(DEFAULT_MAX_RESULTS)
85            .clamp(1, MAX_RESULTS_CAP);
86
87        let search_query = match context.as_deref() {
88            Some(ctx) if !ctx.trim().is_empty() => format!("{query} {ctx}"),
89            _ => query.clone(),
90        };
91
92        let project_dir = self.project_dir.clone();
93        let kind_filter = kind.clone();
94        let started = Instant::now();
95
96        // Index + search are CPU-bound; keep the runtime free.
97        let result = tokio::task::spawn_blocking(move || {
98            explore_repo(
99                &project_dir,
100                &search_query,
101                kind_filter.as_deref(),
102                max_results,
103            )
104        })
105        .await;
106
107        let elapsed_ms = started.elapsed().as_millis() as u64;
108
109        match result {
110            Ok(Ok(report)) => Ok(helpers::ok(
111                invocation.id,
112                json!({
113                    "schema_version": helpers::SPECIALIZED_SCHEMA_VERSION,
114                    "query": query,
115                    "context": context,
116                    "locations": report.locations,
117                    "files_indexed": report.files_indexed,
118                    "symbols_considered": report.symbols_considered,
119                    "text_hits": report.text_hits,
120                    "elapsed_ms": elapsed_ms,
121                    "engine": "bm25+symbols",
122                }),
123            )),
124            Ok(Err(err)) => Ok(ToolResult {
125                invocation_id: invocation.id,
126                ok: false,
127                output: json!({
128                    "error": format!("repo_explore failed: {err:#}"),
129                    "elapsed_ms": elapsed_ms,
130                }),
131            }),
132            Err(err) => Ok(ToolResult {
133                invocation_id: invocation.id,
134                ok: false,
135                output: json!({
136                    "error": format!("repo_explore task join error: {err}"),
137                    "elapsed_ms": elapsed_ms,
138                }),
139            }),
140        }
141    }
142}
143
144struct ExploreReport {
145    locations: Vec<serde_json::Value>,
146    files_indexed: usize,
147    symbols_considered: usize,
148    text_hits: usize,
149}
150
151#[derive(Debug, Clone)]
152struct LocationHit {
153    path: PathBuf,
154    start_line: usize,
155    end_line: usize,
156    kind: String,
157    name: Option<String>,
158    snippet: String,
159    score: f64,
160    reasons: Vec<String>,
161}
162
163fn explore_repo(
164    project_dir: &Path,
165    query: &str,
166    kind: Option<&str>,
167    max_results: usize,
168) -> Result<ExploreReport> {
169    let index = build_index(project_dir)?;
170    let symbols = ranked_symbol_matches(&index, query, kind);
171    // Pull a wider BM25 pool so merge can re-rank against symbols.
172    let text_pool = (max_results * 3).clamp(15, 80);
173    let text_matches = search_text_matches(&index, query, text_pool);
174
175    let locations = merge_locations(&symbols, &text_matches, max_results);
176
177    Ok(ExploreReport {
178        locations: locations
179            .into_iter()
180            .map(|hit| {
181                json!({
182                    "path": path_display(&hit.path),
183                    "start_line": hit.start_line,
184                    "end_line": hit.end_line,
185                    "kind": hit.kind,
186                    "name": hit.name,
187                    "snippet": hit.snippet,
188                    "score": hit.score,
189                    "reasons": hit.reasons,
190                    "why": why_summary(&hit),
191                })
192            })
193            .collect(),
194        files_indexed: index.files.len(),
195        symbols_considered: symbols.len(),
196        text_hits: text_matches.len(),
197    })
198}
199
200fn merge_locations(
201    symbols: &[RankedSymbolRecord],
202    text_matches: &[TextMatchRecord],
203    max_results: usize,
204) -> Vec<LocationHit> {
205    // Key: (path, start_line) — keep best score per anchor.
206    let mut by_anchor: HashMap<(String, usize), LocationHit> = HashMap::new();
207
208    for ranked in symbols {
209        let path = ranked.symbol.path.clone();
210        let line = ranked.symbol.line.max(1);
211        let key = (path_display(&path), line);
212        let mut reasons = ranked.reasons.clone();
213        reasons.push("symbol".to_string());
214        let hit = LocationHit {
215            path,
216            start_line: line,
217            // Small window so the agent can read a compact range.
218            end_line: line.saturating_add(12),
219            kind: ranked.symbol.kind.clone(),
220            name: Some(ranked.symbol.name.clone()),
221            snippet: ranked.symbol.signature.clone(),
222            score: ranked.score * SYMBOL_SCORE_SCALE,
223            reasons,
224        };
225        insert_best(&mut by_anchor, key, hit);
226    }
227
228    for text in text_matches {
229        let path = text.path.clone();
230        let line = text.line.max(1);
231        let key = (path_display(&path), line);
232        let hit = LocationHit {
233            path,
234            start_line: line,
235            end_line: line.saturating_add(8),
236            kind: text.kind.clone(),
237            name: None,
238            snippet: text.text.clone(),
239            score: text.score * TEXT_SCORE_SCALE,
240            reasons: vec!["bm25".to_string(), text.kind.clone()],
241        };
242        insert_best(&mut by_anchor, key, hit);
243    }
244
245    let mut hits: Vec<LocationHit> = by_anchor.into_values().collect();
246    hits.sort_by(|a, b| {
247        score_cmp(b.score, a.score)
248            .then_with(|| path_display(&a.path).cmp(&path_display(&b.path)))
249            .then_with(|| a.start_line.cmp(&b.start_line))
250    });
251    hits.truncate(max_results);
252    hits
253}
254
255fn insert_best(
256    map: &mut HashMap<(String, usize), LocationHit>,
257    key: (String, usize),
258    hit: LocationHit,
259) {
260    match map.get(&key) {
261        Some(existing) if existing.score >= hit.score => {
262            // Keep existing; maybe merge reasons for transparency.
263        }
264        Some(existing) => {
265            let mut merged = hit;
266            for reason in &existing.reasons {
267                if !merged.reasons.iter().any(|r| r == reason) {
268                    merged.reasons.push(reason.clone());
269                }
270            }
271            // Prefer a named symbol if the winner was text-only.
272            if merged.name.is_none() {
273                merged.name = existing.name.clone();
274            }
275            if merged.snippet.is_empty() {
276                merged.snippet = existing.snippet.clone();
277            }
278            map.insert(key, merged);
279        }
280        None => {
281            map.insert(key, hit);
282        }
283    }
284}
285
286fn score_cmp(a: f64, b: f64) -> Ordering {
287    a.partial_cmp(&b).unwrap_or(Ordering::Equal)
288}
289
290fn path_display(path: &Path) -> String {
291    path.to_string_lossy().replace('\\', "/")
292}
293
294fn why_summary(hit: &LocationHit) -> String {
295    let mut parts = Vec::new();
296    if let Some(name) = &hit.name {
297        parts.push(format!("{kind} `{name}`", kind = hit.kind));
298    } else {
299        parts.push(hit.kind.clone());
300    }
301    if !hit.reasons.is_empty() {
302        parts.push(format!("matched via {}", hit.reasons.join(", ")));
303    }
304    parts.join(" — ")
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310    use std::fs;
311
312    fn write_src(dir: &Path, rel: &str, body: &str) {
313        let path = dir.join(rel);
314        if let Some(parent) = path.parent() {
315            fs::create_dir_all(parent).unwrap();
316        }
317        fs::write(path, body).unwrap();
318    }
319
320    #[test]
321    fn definition_has_correct_name_and_kind() {
322        let tool = RepoExploreTool::new(PathBuf::from("/tmp"));
323        let def = tool.definition();
324        assert_eq!(def.name, "repo_explore");
325        assert_eq!(def.kind, ToolKind::Read);
326        assert!(def.description.to_lowercase().contains("bm25"));
327        assert!(
328            def.description.to_lowercase().contains("does not spawn"),
329            "should clarify no nested agent turn"
330        );
331    }
332
333    #[test]
334    fn explore_finds_symbol_and_doc_hits() {
335        let dir = tempfile::tempdir().unwrap();
336        write_src(
337            dir.path(),
338            "src/lib.rs",
339            "/// Handles tool approval for guarded commands.\n\
340             pub fn validate_tool_approval() {}\n\
341             fn other() { validate_tool_approval(); }\n",
342        );
343        write_src(
344            dir.path(),
345            "src/security.rs",
346            "pub struct SecurityPolicy;\nimpl SecurityPolicy {\n  pub fn is_guarded_command() {}\n}\n",
347        );
348
349        let report = explore_repo(dir.path(), "tool approval guarded", None, 10).unwrap();
350        assert!(
351            report.files_indexed >= 2,
352            "indexed: {}",
353            report.files_indexed
354        );
355        assert!(!report.locations.is_empty(), "expected locations, got none");
356
357        let blob = serde_json::to_string(&report.locations).unwrap();
358        assert!(
359            blob.contains("validate_tool_approval")
360                || blob.contains("approval")
361                || blob.contains("is_guarded_command")
362                || blob.contains("SecurityPolicy"),
363            "unexpected locations: {blob}"
364        );
365    }
366
367    #[test]
368    fn explore_respects_max_results() {
369        let dir = tempfile::tempdir().unwrap();
370        write_src(
371            dir.path(),
372            "src/a.rs",
373            "pub fn alpha() {}\npub fn alphabet() {}\npub fn alpine() {}\n",
374        );
375        let report = explore_repo(dir.path(), "alp", None, 2).unwrap();
376        assert!(report.locations.len() <= 2);
377    }
378
379    #[tokio::test]
380    async fn invoke_returns_structured_locations() {
381        let dir = tempfile::tempdir().unwrap();
382        write_src(
383            dir.path(),
384            "src/main.rs",
385            "fn main() { println!(\"hello repo explore\"); }\n",
386        );
387        let tool = RepoExploreTool::new(dir.path().to_path_buf());
388        let result = tool
389            .invoke(ToolInvocation {
390                id: "t1".into(),
391                tool_name: "repo_explore".into(),
392                input: json!({ "query": "repo explore hello" }),
393            })
394            .await
395            .unwrap();
396        assert!(result.ok, "{:?}", result.output);
397        assert_eq!(
398            result.output.get("engine").and_then(|v| v.as_str()),
399            Some("bm25+symbols")
400        );
401        assert!(result.output.get("locations").is_some());
402        assert!(result.output.get("elapsed_ms").is_some());
403    }
404}