Skip to main content

telos_agent/tools/builtin/
code_index.rs

1use async_trait::async_trait;
2use serde_json::{Value, json};
3
4use crate::error::AgentError;
5use crate::knowledge::code_index::CodeIndex;
6use crate::tools::api::{Tool, ToolContext, ToolDefinition, ToolOutput};
7
8pub struct CodeSearchTool;
9pub struct CodeContextTool;
10pub struct CodeIndexRefreshTool;
11
12#[async_trait]
13impl Tool for CodeSearchTool {
14    fn definition(&self) -> ToolDefinition {
15        ToolDefinition {
16            name: "CodeSearch".into(),
17            description:
18                "Search the local code index and return exact file paths and line numbers.".into(),
19            input_schema: json!({"type":"object","properties":{
20                "query":{"type":"string"},
21                "path_prefix":{"type":"string","description":"Optional path substring filter."},
22                "max_results":{"type":"integer","default":50},
23                "case_sensitive":{"type":"boolean","default":false}
24            },"required":["query"]}),
25        }
26    }
27
28    fn prompt_text(&self) -> Option<&'static str> {
29        Some("Use CodeSearch for indexed repository search before broad filesystem reads.")
30    }
31
32    fn is_concurrency_safe(&self, _arguments: &Value) -> bool {
33        true
34    }
35
36    async fn invoke(
37        &self,
38        arguments: Value,
39        context: ToolContext,
40    ) -> Result<ToolOutput, AgentError> {
41        let query = required_string(&arguments, "query")?.to_string();
42        let path_prefix = arguments.get("path_prefix").and_then(Value::as_str).map(str::to_string);
43        let max_results =
44            arguments.get("max_results").and_then(Value::as_u64).unwrap_or(50).min(500) as usize;
45        let case_sensitive =
46            arguments.get("case_sensitive").and_then(Value::as_bool).unwrap_or(false);
47        let root = context.cwd.clone();
48        tokio::task::spawn_blocking(move || {
49            let index = CodeIndex::load_or_refresh(root).map_err(|err| {
50                AgentError::ToolExecution { tool: "CodeSearch".into(), message: err.to_string() }
51            })?;
52            let matches = index.search(&query, path_prefix.as_deref(), max_results, case_sensitive);
53            Ok(ToolOutput::json(json!({
54                "index_path": CodeIndex::index_path(&index.root),
55                "count": matches.len(),
56                "matches": matches,
57            })))
58        })
59        .await
60        .map_err(|err| AgentError::ToolExecution {
61            tool: "CodeSearch".into(),
62            message: format!("code index task panicked: {err}"),
63        })?
64    }
65}
66
67#[async_trait]
68impl Tool for CodeContextTool {
69    fn definition(&self) -> ToolDefinition {
70        ToolDefinition {
71            name: "CodeContext".into(),
72            description: "Return nearby indexed code lines for a path and line number.".into(),
73            input_schema: json!({"type":"object","properties":{
74                "path":{"type":"string"},
75                "line":{"type":"integer"},
76                "before":{"type":"integer","default":5},
77                "after":{"type":"integer","default":5}
78            },"required":["path","line"]}),
79        }
80    }
81
82    fn is_concurrency_safe(&self, _arguments: &Value) -> bool {
83        true
84    }
85
86    async fn invoke(
87        &self,
88        arguments: Value,
89        context: ToolContext,
90    ) -> Result<ToolOutput, AgentError> {
91        let path = required_string(&arguments, "path")?.to_string();
92        let line = arguments.get("line").and_then(Value::as_u64).unwrap_or(0) as usize;
93        let before = arguments.get("before").and_then(Value::as_u64).unwrap_or(5).min(100) as usize;
94        let after = arguments.get("after").and_then(Value::as_u64).unwrap_or(5).min(100) as usize;
95        let root = context.cwd.clone();
96        tokio::task::spawn_blocking(move || {
97            let index = CodeIndex::load_or_refresh(root).map_err(|err| {
98                AgentError::ToolExecution { tool: "CodeContext".into(), message: err.to_string() }
99            })?;
100            let lines = index.context(&path, line, before, after).ok_or_else(|| {
101                AgentError::ToolExecution {
102                    tool: "CodeContext".into(),
103                    message: format!("path not found in code index: {path}"),
104                }
105            })?;
106            Ok(ToolOutput::json(json!({"path": path, "line": line, "lines": lines})))
107        })
108        .await
109        .map_err(|err| AgentError::ToolExecution {
110            tool: "CodeContext".into(),
111            message: format!("code index task panicked: {err}"),
112        })?
113    }
114}
115
116#[async_trait]
117impl Tool for CodeIndexRefreshTool {
118    fn definition(&self) -> ToolDefinition {
119        ToolDefinition {
120            name: "CodeIndexRefresh".into(),
121            description: "Refresh the local code index under .telos/index/code_index.json.".into(),
122            input_schema: json!({"type":"object","properties":{}}),
123        }
124    }
125
126    async fn invoke(
127        &self,
128        _arguments: Value,
129        context: ToolContext,
130    ) -> Result<ToolOutput, AgentError> {
131        let root = context.cwd.clone();
132        tokio::task::spawn_blocking(move || {
133            let index = CodeIndex::refresh(root).map_err(|err| AgentError::ToolExecution {
134                tool: "CodeIndexRefresh".into(),
135                message: err.to_string(),
136            })?;
137            Ok(ToolOutput::json(json!({
138                "index_path": CodeIndex::index_path(&index.root),
139                "files": index.files.len(),
140            })))
141        })
142        .await
143        .map_err(|err| AgentError::ToolExecution {
144            tool: "CodeIndexRefresh".into(),
145            message: format!("code index task panicked: {err}"),
146        })?
147    }
148}
149
150fn required_string<'a>(arguments: &'a Value, key: &str) -> Result<&'a str, AgentError> {
151    arguments
152        .get(key)
153        .and_then(Value::as_str)
154        .filter(|value| !value.trim().is_empty())
155        .ok_or_else(|| AgentError::Validation(format!("missing `{key}`")))
156}
157
158#[cfg(test)]
159mod tests {
160    use std::sync::Arc;
161
162    use serde_json::json;
163
164    use super::{CodeContextTool, CodeIndexRefreshTool, CodeSearchTool};
165    use crate::knowledge::code_index::CodeIndex;
166    use crate::tools::api::{Tool, ToolContext};
167
168    fn test_context(cwd: std::path::PathBuf) -> ToolContext {
169        ToolContext {
170            session_id: "test".into(),
171            turn_id: 1,
172            tool_call_id: None,
173            cwd,
174            env: Default::default(),
175            messages: Arc::new(vec![]),
176            progress: None,
177            read_file_state: Arc::new(tokio::sync::Mutex::new(Default::default())),
178            timeout: None,
179            max_file_read_bytes: 50 * 1024 * 1024,
180        }
181    }
182
183    #[tokio::test]
184    async fn refresh_and_search_normalize_nested_paths_to_forward_slashes() {
185        let dir = tempfile::tempdir().unwrap();
186        let nested = dir.path().join("src").join("windows");
187        std::fs::create_dir_all(&nested).unwrap();
188        std::fs::write(nested.join("mod.rs"), "fn windows_path() {}\n").unwrap();
189        let ctx = test_context(dir.path().to_path_buf());
190
191        let refresh = CodeIndexRefreshTool.invoke(json!({}), ctx.clone()).await.unwrap().content;
192        assert_eq!(refresh["index_path"], json!(CodeIndex::index_path(dir.path())));
193
194        let search = CodeSearchTool
195            .invoke(json!({"query": "windows_path", "path_prefix": "src/windows"}), ctx.clone())
196            .await
197            .unwrap()
198            .content;
199        assert_eq!(search["count"], 1);
200        assert_eq!(search["matches"][0]["path"], "src/windows/mod.rs");
201        assert!(!search["matches"][0]["path"].as_str().unwrap().contains('\\'));
202
203        let context = CodeContextTool
204            .invoke(json!({"path": "src/windows/mod.rs", "line": 1, "before": 0, "after": 0}), ctx)
205            .await
206            .unwrap()
207            .content;
208        assert_eq!(context["lines"][0]["text"], "fn windows_path() {}");
209    }
210}