Skip to main content

lc_tools/extended/
file.rs

1//! File 工具(沙箱安全)
2
3use 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
13/// 文件读写工具(限制在 base_path 沙箱内,扩展名白名单,大小限制)
14pub 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        // Reject files without an extension (bypasses whitelist)
52        if target.extension().is_none() {
53            return Err(ToolError::InvalidInput(
54                "文件必须包含扩展名(无扩展名文件不在白名单中)".to_string(),
55            ));
56        }
57
58        // 扩展名检查
59        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        // 路径越界检查
69        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}