1use 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;
27const 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 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 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 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 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 }
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 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}