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    /// 创建文件工具(沙箱根目录为 `base_path`)。
22    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    /// 设置扩展名白名单(builder 风格)。
36    pub fn with_allowed_extensions(mut self, exts: Vec<String>) -> Self {
37        self.allowed_extensions = exts;
38        self
39    }
40
41    /// 设置最大文件大小(字节,builder 风格)。
42    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        // Reject files without an extension (bypasses whitelist)
55        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        // 扩展名检查
63        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        // 路径越界检查
73        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    /// 读取沙箱内文件内容(受路径越界与大小限制)。
96    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    /// 写入内容到沙箱内文件(自动创建父目录)。
113    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    /// 列出沙箱内目录下的条目名。
136    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}