lc_tools/extended/
file.rs1use std::path::PathBuf;
4
5use async_trait::async_trait;
6use serde_json::Value;
7use tokio::fs;
8use tokio::io::AsyncWriteExt;
9
10use lc_core::tools::ToolError;
11use lc_core::BaseTool;
12
13pub struct FileTool {
15 base_path: PathBuf,
16 allowed_extensions: Vec<String>,
17 max_size: usize,
18}
19
20impl FileTool {
21 pub fn new(base_path: PathBuf) -> Self {
23 Self {
24 base_path,
25 allowed_extensions: vec![
26 "txt".to_string(),
27 "md".to_string(),
28 "json".to_string(),
29 "csv".to_string(),
30 ],
31 max_size: 10 * 1024 * 1024,
32 }
33 }
34
35 pub fn with_allowed_extensions(mut self, exts: Vec<String>) -> Self {
37 self.allowed_extensions = exts;
38 self
39 }
40
41 pub fn with_max_size(mut self, size: usize) -> Self {
43 self.max_size = size;
44 self
45 }
46
47 fn safe_path(&self, relative: &str) -> Result<PathBuf, ToolError> {
48 let base = self
49 .base_path
50 .canonicalize()
51 .map_err(|e| ToolError::InvalidInput(format!("invalid base_path: {}", e)))?;
52 let target = base.join(relative);
53
54 if target.extension().is_none() {
56 return Err(ToolError::InvalidInput(
57 "file must have an extension (files without an extension are not in the whitelist)"
58 .to_string(),
59 ));
60 }
61
62 if let Some(ext) = target.extension().and_then(|e| e.to_str()) {
64 if !self.allowed_extensions.iter().any(|a| a == ext) {
65 return Err(ToolError::InvalidInput(format!(
66 "extension not allowed: {} (allowed: {:?})",
67 ext, self.allowed_extensions
68 )));
69 }
70 }
71
72 let canon = target
74 .canonicalize()
75 .or_else(|_| {
76 let file_name = target.file_name().ok_or_else(|| {
77 std::io::Error::new(std::io::ErrorKind::InvalidInput, "empty file name")
78 })?;
79 target
80 .parent()
81 .and_then(|p| p.canonicalize().ok())
82 .map(|p| p.join(file_name))
83 .ok_or_else(|| std::io::Error::new(std::io::ErrorKind::NotFound, "parent"))
84 })
85 .map_err(|e| ToolError::InvalidInput(format!("invalid path: {}", e)))?;
86
87 if !canon.starts_with(&base) {
88 return Err(ToolError::InvalidInput(
89 "path escapes base_path sandbox".to_string(),
90 ));
91 }
92 Ok(canon)
93 }
94
95 pub async fn read(&self, path: &str) -> Result<String, ToolError> {
97 let p = self.safe_path(path)?;
98 let metadata = fs::metadata(&p)
99 .await
100 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
101 if metadata.len() as usize > self.max_size {
102 return Err(ToolError::InvalidInput(format!(
103 "file exceeds {} bytes",
104 self.max_size
105 )));
106 }
107 fs::read_to_string(p)
108 .await
109 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))
110 }
111
112 pub async fn write(&self, path: &str, content: &str) -> Result<(), ToolError> {
114 if content.len() > self.max_size {
115 return Err(ToolError::InvalidInput(format!(
116 "content exceeds {} bytes",
117 self.max_size
118 )));
119 }
120 let p = self.safe_path(path)?;
121 if let Some(parent) = p.parent() {
122 fs::create_dir_all(parent)
123 .await
124 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
125 }
126 let mut f = fs::File::create(&p)
127 .await
128 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
129 f.write_all(content.as_bytes())
130 .await
131 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
132 Ok(())
133 }
134
135 pub async fn list(&self, dir: &str) -> Result<Vec<String>, ToolError> {
137 let base = self
138 .base_path
139 .canonicalize()
140 .map_err(|e| ToolError::InvalidInput(format!("invalid base_path: {}", e)))?;
141 let target = base.join(dir);
142 let canon = target
143 .canonicalize()
144 .map_err(|e| ToolError::InvalidInput(format!("invalid path: {}", e)))?;
145 if !canon.starts_with(&base) {
146 return Err(ToolError::InvalidInput(
147 "path escapes base_path".to_string(),
148 ));
149 }
150 let mut entries = fs::read_dir(&canon)
151 .await
152 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?;
153 let mut result = Vec::new();
154 while let Some(entry) = entries
155 .next_entry()
156 .await
157 .map_err(|e| ToolError::ExecutionFailed(e.to_string()))?
158 {
159 result.push(entry.file_name().to_string_lossy().to_string());
160 }
161 Ok(result)
162 }
163}
164
165#[async_trait]
166impl BaseTool for FileTool {
167 fn name(&self) -> &str {
168 "file_operation"
169 }
170
171 fn description(&self) -> &str {
172 "文件操作。输入 JSON: {\"op\": \"read|write|list\", \"path\": \"...\", \"content\": \"...\"}"
173 }
174
175 async fn run(&self, input: String) -> Result<String, ToolError> {
176 let v: Value =
177 serde_json::from_str(&input).map_err(|e| ToolError::InvalidInput(e.to_string()))?;
178 let op = v
179 .get("op")
180 .and_then(|x| x.as_str())
181 .ok_or_else(|| ToolError::InvalidInput("missing op".to_string()))?;
182 let path = v
183 .get("path")
184 .and_then(|x| x.as_str())
185 .ok_or_else(|| ToolError::InvalidInput("missing path".to_string()))?;
186 match op {
187 "read" => self.read(path).await,
188 "write" => {
189 let content = v.get("content").and_then(|x| x.as_str()).unwrap_or("");
190 self.write(path, content).await?;
191 Ok("写入成功".to_string())
192 }
193 "list" => {
194 let list = self.list(path).await?;
195 serde_json::to_string(&list).map_err(|e| ToolError::ExecutionFailed(e.to_string()))
196 }
197 other => Err(ToolError::InvalidInput(format!("unknown op: {}", other))),
198 }
199 }
200}
201
202#[cfg(test)]
203mod tests {
204 use super::*;
205 use tempfile::TempDir;
206
207 fn tool() -> (FileTool, TempDir) {
208 let dir = TempDir::new().unwrap();
209 let tool = FileTool::new(dir.path().to_path_buf());
210 (tool, dir)
211 }
212
213 #[tokio::test]
214 async fn test_write_and_read() {
215 let (tool, _dir) = tool();
216 tool.write("test.txt", "hello").await.unwrap();
217 let content = tool.read("test.txt").await.unwrap();
218 assert_eq!(content, "hello");
219 }
220
221 #[tokio::test]
222 async fn test_extension_not_allowed() {
223 let (tool, _dir) = tool();
224 assert!(tool.write("file.exe", "x").await.is_err());
225 }
226
227 #[tokio::test]
228 async fn test_path_traversal_blocked() {
229 let (tool, _dir) = tool();
230 assert!(tool.read("../../../../etc/passwd").await.is_err());
231 }
232
233 #[tokio::test]
234 async fn test_list() {
235 let (tool, _dir) = tool();
236 tool.write("a.txt", "a").await.unwrap();
237 tool.write("b.txt", "b").await.unwrap();
238 let list = tool.list(".").await.unwrap();
239 assert!(list.len() >= 2);
240 }
241
242 #[tokio::test]
243 async fn test_max_size_exceeded() {
244 let (tool, _dir) = tool();
245 let tool = tool.with_max_size(5);
246 assert!(tool.write("big.txt", "12345678").await.is_err());
247 }
248
249 #[tokio::test]
250 async fn test_run_write_read_via_base_tool() {
251 let (tool, _dir) = tool();
252 let write_input = r#"{"op":"write","path":"x.txt","content":"hi"}"#;
253 assert!(tool.run(write_input.to_string()).await.is_ok());
254 let read_input = r#"{"op":"read","path":"x.txt"}"#;
255 let result = tool.run(read_input.to_string()).await.unwrap();
256 assert_eq!(result, "hi");
257 }
258}