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