From 0bbba2b936fd92f8a0e6c439b8379ed0b13a9413 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=AD=A6=E9=98=B3?= Date: Fri, 8 May 2026 17:09:07 +0800 Subject: [PATCH] Enforce session workspace boundaries --- backend/api/server.py | 14 +- frontend/app/app/api/chat/route.ts | 15 +- frontend/app/lib/claw-auth.ts | 16 ++ skills/data-factory-sql/SKILL.md | 6 +- skills/data-factory-sql/run_sql.py | 13 +- skills/eval-repair/SKILL.md | 13 +- src/agent_prompting.py | 23 ++- src/agent_tools.py | 232 ++++++++++++++++++++++++++--- src/data_agent_inputs.py | 13 +- tests/test_agent_prompting.py | 25 ++-- tests/test_data_agent_inputs.py | 11 ++ tests/test_extended_tools.py | 160 ++++++++++++++++++++ tests/test_python_exec_tool.py | 20 +++ 13 files changed, 509 insertions(+), 52 deletions(-) diff --git a/backend/api/server.py b/backend/api/server.py index 963c06d..64b7d89 100644 --- a/backend/api/server.py +++ b/backend/api/server.py @@ -48,6 +48,7 @@ from src.token_budget import calculate_token_budget STATIC_DIR = Path(__file__).resolve().parents[2] / 'frontend' / 'legacy-static' +API_TOOL_CONTENT_MAX_CHARS = 20000 VALIDATED_CHAT_MODEL_PROVIDERS = { # 这些 provider 已验证可以在当前 WebUI 中使用。 @@ -1293,9 +1294,12 @@ def _serialize_run_result(result: Any) -> dict[str, Any]: def _normalize_transcript_entry(entry: dict[str, Any]) -> dict[str, Any]: + content = entry.get('content', '') + if entry.get('role') == 'tool' and isinstance(content, str): + content = _truncate_api_text(content, API_TOOL_CONTENT_MAX_CHARS) out: dict[str, Any] = { 'role': entry.get('role', ''), - 'content': entry.get('content', ''), + 'content': content, } for key in ('name', 'tool_call_id', 'tool_calls', 'metadata', 'message_id'): if key in entry and entry[key] not in (None, '', [], {}): @@ -1306,6 +1310,14 @@ def _normalize_transcript_entry(entry: dict[str, Any]) -> dict[str, Any]: return out +def _truncate_api_text(text: str, limit: int) -> str: + if len(text) <= limit: + return text + head = text[: limit // 2] + tail = text[-(limit // 2) :] + return f'{head}\n...[truncated for api response]...\n{tail}' + + def _serialize_stored_session(stored: StoredAgentSession) -> dict[str, Any]: return { 'session_id': stored.session_id, diff --git a/frontend/app/app/api/chat/route.ts b/frontend/app/app/api/chat/route.ts index 4d239c5..2f5f0e2 100644 --- a/frontend/app/app/api/chat/route.ts +++ b/frontend/app/app/api/chat/route.ts @@ -9,6 +9,8 @@ import { import { accountSessionInputRoot, accountSessionOutputRoot, + accountSessionRoot, + accountSessionScratchpadRoot, getCurrentAccount, } from "@/lib/claw-auth"; @@ -274,17 +276,26 @@ async function getLastUserText( async function ensureSessionDirectories(accountId: string, sessionId: string) { await Promise.all([ + mkdir(accountSessionRoot(accountId, sessionId), { recursive: true }), mkdir(accountSessionInputRoot(accountId, sessionId), { recursive: true }), mkdir(accountSessionOutputRoot(accountId, sessionId), { recursive: true }), + mkdir(accountSessionScratchpadRoot(accountId, sessionId), { + recursive: true, + }), ]); } function renderSessionRuntimeContext(accountId: string, sessionId: string) { return [ - "[当前会话目录]", + "[当前 session 工作区]", + `- session_root: ${accountSessionRoot(accountId, sessionId)}`, `- 输入目录: ${accountSessionInputRoot(accountId, sessionId)}`, `- 输出目录: ${accountSessionOutputRoot(accountId, sessionId)}`, - "如需保存本轮任务产物,请优先写入输出目录。", + `- 临时目录: ${accountSessionScratchpadRoot(accountId, sessionId)}`, + "当前 session 目录是默认可写工作区。", + "交付产物必须优先写入输出目录。", + "临时脚本、缓存和中间结果必须写入临时目录。", + "可以读取完成任务所需的外部资料,但不要修改平台服务代码目录。", ].join("\n\n"); } diff --git a/frontend/app/lib/claw-auth.ts b/frontend/app/lib/claw-auth.ts index a975c76..fdb0c6a 100644 --- a/frontend/app/lib/claw-auth.ts +++ b/frontend/app/lib/claw-auth.ts @@ -110,6 +110,22 @@ export function accountSessionOutputRoot(accountId: string, sessionId: string) { return path.join(accountBaseRoot(accountId), "sessions", sessionId, "output"); } +export function accountSessionScratchpadRoot( + accountId: string, + sessionId: string, +) { + return path.join( + accountBaseRoot(accountId), + "sessions", + sessionId, + "scratchpad", + ); +} + +export function accountSessionRoot(accountId: string, sessionId: string) { + return path.join(accountBaseRoot(accountId), "sessions", sessionId); +} + async function createSession(account: UserRecord) { const token = randomBytes(32).toString("hex"); const sessionsFile = await readSessions(); diff --git a/skills/data-factory-sql/SKILL.md b/skills/data-factory-sql/SKILL.md index 1db4fd1..20cf5dc 100644 --- a/skills/data-factory-sql/SKILL.md +++ b/skills/data-factory-sql/SKILL.md @@ -49,7 +49,7 @@ skills/data-factory-sql/ } ``` -多行 SQL 或复杂 SQL 不要通过命令行字符串硬塞。优先把 SQL 写到当前会话 scratchpad 或用户指定的任务目录,再用 `-f` 执行: +多行 SQL 或复杂 SQL 不要通过命令行字符串硬塞。优先把 SQL 写到当前 session/scratchpad;只有用户明确指定外部目标文件时,才写到用户指定目录,再用 `-f` 执行: ```json { @@ -86,7 +86,7 @@ https://data.mioffice.cn/workspace/?wid=#/workspace/.csv` | +| 输出 | 当前 session/output/data_factory_<时间戳>.csv;如果没有 session 环境,则退回 `~/Downloads/data_factory_<时间戳>.csv` | | 轮询间隔 | 2.0s | | 查询超时 | 600s | @@ -166,7 +166,7 @@ https://data.mioffice.cn/workspace/?wid=#/workspace/ str: def default_output_path(query_id: str) -> Path: - """Default: ~/Downloads/data_factory_.csv""" + """Default to current session output when python_exec exposes it.""" ts = time.strftime("%Y%m%d_%H%M%S") + scratchpad = os.environ.get("PYTHON_EXEC_SCRATCHPAD") + if scratchpad: + output_dir = Path(scratchpad).resolve().parent / "output" + output_dir.mkdir(parents=True, exist_ok=True) + return output_dir / f"data_factory_{ts}.csv" downloads = Path.home() / "Downloads" downloads.mkdir(parents=True, exist_ok=True) return downloads / f"data_factory_{ts}.csv" @@ -136,7 +141,11 @@ def build_parser() -> argparse.ArgumentParser: p.add_argument("--catalog", help="default catalog") p.add_argument("--schema", help="default schema") p.add_argument( - "--output", help="output CSV path (default: ~/Downloads/data_factory_.csv)" + "--output", + help=( + "output CSV path (default: session output when PYTHON_EXEC_SCRATCHPAD " + "is set, otherwise ~/Downloads/data_factory_.csv)" + ), ) p.add_argument( "--no-save", action="store_true", help="don't save CSV, only print to stdout" diff --git a/skills/eval-repair/SKILL.md b/skills/eval-repair/SKILL.md index 193e93a..3f2355e 100644 --- a/skills/eval-repair/SKILL.md +++ b/skills/eval-repair/SKILL.md @@ -57,15 +57,18 @@ allowed_tools: read_file, write_file, edit_file, grep_search, glob_search, ask_u ## 推荐产物 ```text -tasks/{task_id}/context/eval_errors.* -tasks/{task_id}/artifacts/error_analysis.md -tasks/{task_id}/artifacts/generation_plan.md -tasks/{task_id}/artifacts/generated_candidates.jsonl -tasks/{task_id}/memory/open_questions.md +output/eval_errors_context.md +output/error_analysis.md +output/generation_plan.md +output/generated_candidates.jsonl +output/open_questions.md ``` +上面的 `output/...` 是逻辑文件名,实际路径必须位于当前 session/output。临时脚本、抽样缓存和中间文件必须位于当前 session/scratchpad。 + ## 约束 - 在 canonical records 达成一致并通过校验前,不要直接导出最终训练/评测格式。 - 不要把自动聚类当成最终事实,聚类和类别命名必须可 review。 - 面向人的总结要简洁,并尽量用样例支撑。 +- 不要在项目根目录、源码目录或其他非 session 目录创建本任务产物;用户明确指定外部目标文件时除外。 diff --git a/src/agent_prompting.py b/src/agent_prompting.py index c4be4c2..f5ee0ad 100644 --- a/src/agent_prompting.py +++ b/src/agent_prompting.py @@ -98,6 +98,7 @@ def build_system_prompt_parts( get_system_section(), get_doing_tasks_section(), get_actions_section(), + get_workspace_boundary_section(), get_using_your_tools_section(enabled_tool_names), get_skill_guidance_section(prompt_context, runtime_config, enabled_tool_names), get_agent_guidance_section(enabled_tool_names, available_agents), @@ -163,7 +164,7 @@ def get_doing_tasks_section() -> str: '对于文件解析、表格转换、JSON 或 JSONL 处理、日志分析、批量校验、数据抽样等任务,如果小型 Python 脚本比手工文本处理更可靠,可以编写小型 Python 脚本。', '小脚本应短小、可读,并明确输入和输出。需要执行 Python 代码或 Python 脚本时,使用 python_exec;不要通过 bash 执行 python、python3 或 .venv/bin/python。', '一次性 Python 分析优先直接传给 python_exec.code,不要为了临时分析创建项目文件。', - '如果确实需要临时脚本、缓存或中间产物,必须写入当前会话 scratchpad 目录,或写入明确的任务产物目录;禁止在项目根目录创建 analyze_*.py、tmp_*.py、scratch_*.py 等临时脚本。', + '如果确实需要临时脚本、缓存或中间产物,必须写入当前 session/scratchpad;交付产物必须优先写入当前 session/output。', '只有用户明确要求长期复用或该脚本属于产品代码时,才把 Python 脚本加入项目源码目录。', '当事情失败时,先诊断原因再改变方向。不要对同一个失败动作反复循环。', '注意不要引入命令注入、SQL 注入、XSS 或不安全 shell 行为等安全问题。', @@ -192,6 +193,20 @@ def get_actions_section() -> str: 遇到意外状态时,先调查清楚,再删除或覆盖。""" +def get_workspace_boundary_section() -> str: + return """# 工作空间边界 + +当前 session 目录是本轮任务的默认工作区,也是默认可写区域。 + +你可以读取完成任务所需的外部资料,包括用户指定路径、标签定义、产品文档、历史样例、线上数据或其他必要文件。读取外部目录时应尽量限定路径和目的,避免无边界扫描依赖目录、构建产物、历史 session 或大型数据目录。 + +默认情况下,所有交付产物必须写入当前 session/output;所有临时脚本、缓存、中间结果必须写入当前 session/scratchpad。 + +除非用户明确要求修改某个具体的非 session 文件或目录,否则不要在当前 session 之外创建、编辑、删除或覆盖文件。 + +平台服务代码目录始终只读,不能修改。不要修改 src、backend、frontend、scripts、部署脚本或平台运行配置等平台代码文件。""" + + def get_using_your_tools_section(enabled_tool_names: set[str]) -> str: items: list[str | list[str]] = [ '当有更具体的专用工具可用时,不要使用 bash 工具。这对可审查性和安全执行很重要。', @@ -201,7 +216,7 @@ def get_using_your_tools_section(enabled_tool_names: set[str]) -> str: if 'edit_file' in enabled_tool_names: items.append('编辑文件时,优先使用 edit_file,而不是 shell 文本替换。') if 'write_file' in enabled_tool_names: - items.append('创建文件时,优先使用 write_file,而不是 heredoc 或 echo 重定向。') + items.append('创建文件时,优先使用 write_file,而不是 heredoc 或 echo 重定向;默认写入当前 session 目录。') if 'glob_search' in enabled_tool_names: items.append('搜索文件时,优先使用 glob_search,而不是 find 或 ls。') if 'grep_search' in enabled_tool_names: @@ -214,14 +229,14 @@ def get_using_your_tools_section(enabled_tool_names: set[str]) -> str: 'python_exec 默认使用当前用户独立 Python venv,不使用项目 .venv;不要用 bash 执行 python 或 pip。' ) items.append( - 'python_exec 会注入 PYTHON_EXEC_SCRATCHPAD 环境变量,指向当前用户当前会话隔离的 scratchpad。一次性脚本、缓存和中间输出都应写入这里。' + 'python_exec 会注入 PYTHON_EXEC_SCRATCHPAD 环境变量,指向当前用户当前会话隔离的 scratchpad。一次性脚本、缓存和中间输出都应写入这里,交付产物写入 session/output。' ) if 'python_package' in enabled_tool_names: items.append( '当 python_exec 因缺少 pandas、pyarrow、openpyxl 等 Python 包失败时,使用 python_package 在当前用户独立 venv 中安装缺失包,然后重试;不要安装到系统 Python 或项目 .venv。' ) items.append( - '不要用 write_file 在项目根目录生成临时 Python 脚本;如果用户明确要求保留脚本,才写入合适的项目路径,并用 python_exec 的 script_path 执行。' + '不要用 write_file 在项目根目录生成临时 Python 脚本;如果用户明确要求保留脚本,仍应优先写入当前 session,除非用户指定具体非 session 目标文件。' ) if 'bash' in enabled_tool_names: items.append( diff --git a/src/agent_tools.py b/src/agent_tools.py index 35dc63d..363cdb1 100644 --- a/src/agent_tools.py +++ b/src/agent_tools.py @@ -1079,6 +1079,53 @@ def _snapshot_text(text: str, limit: int = 240) -> str: return normalized[: limit - 3] + '...' +_GREP_DEFAULT_SKIPPED_DIRS = { + '.git', + '.hg', + '.mypy_cache', + '.next', + '.port_sessions', + '.pytest_cache', + '.ruff_cache', + '.svn', + '.venv', + '__pycache__', + 'build', + 'coverage', + 'dist', + 'node_modules', + 'router_session_parquet', + 'venv', +} +_GREP_BINARY_SUFFIXES = { + '.7z', + '.avif', + '.bin', + '.doc', + '.docx', + '.gz', + '.jpeg', + '.jpg', + '.parquet', + '.pdf', + '.png', + '.pyc', + '.snappy', + '.tar', + '.webp', + '.xls', + '.xlsx', + '.zip', +} +_GREP_MAX_LINE_CHARS = 800 +_PLATFORM_READONLY_DIRS = {'src', 'backend', 'frontend', 'scripts'} +_PLATFORM_ROOT_MARKERS = ( + Path('src') / 'agent_tools.py', + Path('backend') / 'api' / 'server.py', + Path('frontend') / 'app', +) + + def _require_string(arguments: dict[str, Any], key: str) -> str: value = arguments.get(key) if not isinstance(value, str) or not value: @@ -1543,19 +1590,90 @@ def _data_agent_session_output_root(context: ToolExecutionContext) -> Path: return context.root / '.port_sessions' / 'data_agent_output' -def _resolve_path(raw_path: str, context: ToolExecutionContext, *, allow_missing: bool = True) -> Path: +def _resolve_path( + raw_path: str, + context: ToolExecutionContext, + *, + allow_missing: bool = True, + allow_outside_root: bool = False, +) -> Path: expanded = Path(raw_path).expanduser() - candidate = expanded if expanded.is_absolute() else context.root / expanded + session_candidate = _session_logical_path(expanded, context) + candidate = ( + expanded + if expanded.is_absolute() + else session_candidate + if session_candidate is not None + else context.root / expanded + ) resolved = candidate.resolve(strict=not allow_missing) try: resolved.relative_to(context.root) except ValueError as exc: + if allow_outside_root: + return resolved raise ToolExecutionError( f'Path {raw_path!r} escapes the workspace root {context.root}' ) from exc return resolved +def _session_logical_path(path: Path, context: ToolExecutionContext) -> Path | None: + """把 output/scratchpad/input 这类逻辑路径路由到当前 session。""" + + if path.is_absolute() or context.scratchpad_directory is None: + return None + if not path.parts: + return None + session_root = context.scratchpad_directory.parent + head, *tail = path.parts + tail_path = Path(*tail) if tail else Path() + if head in {'output', 'outputs'}: + return session_root / 'output' / tail_path + if head in {'scratchpad', 'scratch'}: + return context.scratchpad_directory / tail_path + if head in {'input', 'inputs'}: + return session_root / 'input' / tail_path + return None + + +def _display_path(path: Path, context: ToolExecutionContext) -> str: + resolved = path.resolve() + try: + return resolved.relative_to(context.root.resolve()).as_posix() + except ValueError: + return resolved.as_posix() + + +def _execution_cwd(context: ToolExecutionContext) -> Path: + if context.scratchpad_directory is not None: + context.scratchpad_directory.mkdir(parents=True, exist_ok=True) + return context.scratchpad_directory + return context.root + + +def _is_platform_app_root(root: Path) -> bool: + return all((root / marker).exists() for marker in _PLATFORM_ROOT_MARKERS) + + +def _is_platform_code_path(path: Path, context: ToolExecutionContext) -> bool: + if not _is_platform_app_root(context.root): + return False + try: + rel = path.resolve().relative_to(context.root.resolve()) + except ValueError: + return False + return bool(rel.parts) and rel.parts[0] in _PLATFORM_READONLY_DIRS + + +def _ensure_not_platform_code_write(path: Path, context: ToolExecutionContext) -> None: + if _is_platform_code_path(path, context): + raise ToolPermissionError( + 'Platform code is read-only in data-agent sessions. ' + 'Do not modify src, backend, frontend, or scripts from this agent.' + ) + + def _ensure_write_allowed(context: ToolExecutionContext) -> None: if not context.permissions.allow_file_write: raise ToolPermissionError( @@ -1572,6 +1690,11 @@ def _ensure_process_execution_allowed(context: ToolExecutionContext, tool_label: def _ensure_shell_allowed(command: str, context: ToolExecutionContext) -> None: _ensure_process_execution_allowed(context, 'Shell commands') + if _looks_like_platform_code_shell_write(command, context): + raise ToolPermissionError( + 'Shell command appears to write platform code. ' + 'Platform code directories are read-only in data-agent sessions.' + ) if context.permissions.allow_destructive_shell_commands: return destructive_patterns = [ @@ -1594,12 +1717,34 @@ def _ensure_shell_allowed(command: str, context: ToolExecutionContext) -> None: ) +def _looks_like_platform_code_shell_write( + command: str, + context: ToolExecutionContext, +) -> bool: + if not _is_platform_app_root(context.root): + return False + lowered = command.lower() + write_marker = re.search( + r'>|(^|[\s;&|])(tee|touch|python|python3|sed\s+-i|perl\s+-pi)\b', + lowered, + ) + if write_marker is None: + return False + platform_fragments = [ + f'{name}/' for name in _PLATFORM_READONLY_DIRS + ] + [ + f'{context.root.as_posix().lower()}/{name}/' + for name in _PLATFORM_READONLY_DIRS + ] + return any(fragment in lowered for fragment in platform_fragments) + + def _list_dir(arguments: dict[str, Any], context: ToolExecutionContext) -> str: raw_path = arguments.get('path', '.') if not isinstance(raw_path, str): raise ToolExecutionError('path must be a string') max_entries = _coerce_int(arguments, 'max_entries', 200) - target = _resolve_path(raw_path, context) + target = _resolve_path(raw_path, context, allow_outside_root=True) if not target.exists(): raise ToolExecutionError(f'Path not found: {raw_path}') if not target.is_dir(): @@ -1608,7 +1753,7 @@ def _list_dir(arguments: dict[str, Any], context: ToolExecutionContext) -> str: lines: list[str] = [] for entry in entries[:max_entries]: kind = 'dir' if entry.is_dir() else 'file' - rel = entry.relative_to(context.root) + rel = _display_path(entry, context) lines.append(f'{kind}\t{rel}') if len(entries) > max_entries: lines.append(f'... truncated at {max_entries} entries ...') @@ -1616,7 +1761,12 @@ def _list_dir(arguments: dict[str, Any], context: ToolExecutionContext) -> str: def _read_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str: - target = _resolve_path(_require_string(arguments, 'path'), context, allow_missing=False) + target = _resolve_path( + _require_string(arguments, 'path'), + context, + allow_missing=False, + allow_outside_root=True, + ) if not target.is_file(): raise ToolExecutionError(f'Path is not a file: {target}') text = target.read_text(encoding='utf-8', errors='replace') @@ -1639,6 +1789,7 @@ def _read_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str: def _write_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str: _ensure_write_allowed(context) target = _resolve_path(_require_string(arguments, 'path'), context) + _ensure_not_platform_code_write(target, context) content = arguments.get('content') if not isinstance(content, str): raise ToolExecutionError('content must be a string') @@ -1649,13 +1800,13 @@ def _write_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str previous_sha256 = hashlib.sha256(previous_text.encode('utf-8')).hexdigest() target.parent.mkdir(parents=True, exist_ok=True) target.write_text(content, encoding='utf-8') - rel = target.relative_to(context.root) + rel = _display_path(target, context) new_sha256 = hashlib.sha256(content.encode('utf-8')).hexdigest() return ( f'wrote {rel} ({len(content)} chars)', { 'action': 'write_file', - 'path': str(rel), + 'path': rel, 'before_exists': previous_text is not None, 'before_sha256': previous_sha256, 'before_size': len(previous_text) if previous_text is not None else 0, @@ -1675,6 +1826,7 @@ def _write_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str def _edit_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str: _ensure_write_allowed(context) target = _resolve_path(_require_string(arguments, 'path'), context, allow_missing=False) + _ensure_not_platform_code_write(target, context) if not target.is_file(): raise ToolExecutionError(f'Path is not a file: {target}') old_text = arguments.get('old_text') @@ -1697,14 +1849,14 @@ def _edit_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str: before_sha256 = hashlib.sha256(current.encode('utf-8')).hexdigest() updated = current.replace(old_text, new_text) if replace_all else current.replace(old_text, new_text, 1) target.write_text(updated, encoding='utf-8') - rel = target.relative_to(context.root) + rel = _display_path(target, context) replaced = occurrences if replace_all else 1 after_sha256 = hashlib.sha256(updated.encode('utf-8')).hexdigest() return ( f'edited {rel}; replaced {replaced} occurrence(s)', { 'action': 'edit_file', - 'path': str(rel), + 'path': rel, 'before_sha256': before_sha256, 'after_sha256': after_sha256, 'before_size': len(current), @@ -1721,6 +1873,7 @@ def _edit_file(arguments: dict[str, Any], context: ToolExecutionContext) -> str: def _notebook_edit(arguments: dict[str, Any], context: ToolExecutionContext) -> str: _ensure_write_allowed(context) target = _resolve_path(_require_string(arguments, 'path'), context, allow_missing=False) + _ensure_not_platform_code_write(target, context) if target.suffix != '.ipynb': raise ToolExecutionError('notebook_edit requires a .ipynb target') if not target.is_file(): @@ -1781,12 +1934,12 @@ def _notebook_edit(arguments: dict[str, Any], context: ToolExecutionContext) -> updated = json.dumps(notebook, ensure_ascii=True, indent=1) + '\n' target.write_text(updated, encoding='utf-8') after_sha256 = hashlib.sha256(updated.encode('utf-8')).hexdigest() - rel = target.relative_to(context.root) + rel = _display_path(target, context) return ( f'updated notebook cell {cell_index} in {rel}', { 'action': 'notebook_edit', - 'path': str(rel), + 'path': rel, 'cell_index': cell_index, 'cell_type': cell['cell_type'], 'before_sha256': before_sha256, @@ -1824,7 +1977,7 @@ def _grep_search(arguments: dict[str, Any], context: ToolExecutionContext) -> st if not isinstance(literal, bool): raise ToolExecutionError('literal must be a boolean') max_matches = _coerce_int(arguments, 'max_matches', 100) - root = _resolve_path(raw_path, context) + root = _resolve_path(raw_path, context, allow_outside_root=True) if not root.exists(): raise ToolExecutionError(f'Path not found: {raw_path}') try: @@ -1833,20 +1986,61 @@ def _grep_search(arguments: dict[str, Any], context: ToolExecutionContext) -> st raise ToolExecutionError(f'Invalid regex pattern: {exc}') from exc hits: list[str] = [] file_iter = root.rglob('*') if root.is_dir() else [root] + root_parts = _relative_parts(root, context) for file_path in file_iter: if not file_path.is_file(): continue + if _grep_should_skip_path(file_path, context, root_parts): + continue try: text = file_path.read_text(encoding='utf-8', errors='replace') except OSError: continue for line_no, line in enumerate(text.splitlines(), start=1): if regex.search(line): - rel = file_path.relative_to(context.root) - hits.append(f'{rel}:{line_no}: {line}') + rel = _display_path(file_path, context) + hits.append(f'{rel}:{line_no}: {_truncate_grep_line(line)}') if len(hits) >= max_matches: - return '\n'.join(hits + [f'... truncated at {max_matches} matches ...']) - return '\n'.join(hits) if hits else '(no matches)' + return _truncate_output( + '\n'.join( + hits + + [f'... truncated at {max_matches} matches ...'] + ), + context.max_output_chars, + ) + if not hits: + return '(no matches)' + return _truncate_output('\n'.join(hits), context.max_output_chars) + + +def _relative_parts(path: Path, context: ToolExecutionContext) -> tuple[str, ...]: + try: + return path.resolve().relative_to(context.root.resolve()).parts + except ValueError: + return () + + +def _grep_should_skip_path( + file_path: Path, + context: ToolExecutionContext, + root_parts: tuple[str, ...], +) -> bool: + rel_parts = _relative_parts(file_path, context) + if file_path.suffix.lower() in _GREP_BINARY_SUFFIXES: + return True + explicit_roots = set(root_parts) + path_parts = rel_parts[:-1] if rel_parts else file_path.parts[:-1] + for part in path_parts: + if part in _GREP_DEFAULT_SKIPPED_DIRS and part not in explicit_roots: + return True + return False + + +def _truncate_grep_line(line: str) -> str: + if len(line) <= _GREP_MAX_LINE_CHARS: + return line + omitted = len(line) - _GREP_MAX_LINE_CHARS + return f'{line[:_GREP_MAX_LINE_CHARS]}... [line truncated, {omitted} chars omitted]' def _run_python_exec(arguments: dict[str, Any], context: ToolExecutionContext) -> str: @@ -1887,7 +2081,7 @@ def _run_python_exec(arguments: dict[str, Any], context: ToolExecutionContext) - try: process = subprocess.Popen( command, - cwd=context.root, + cwd=_execution_cwd(context), stdin=subprocess.PIPE if stdin else None, stdout=subprocess.PIPE, stderr=subprocess.PIPE, @@ -2045,7 +2239,7 @@ def _run_python_package(arguments: dict[str, Any], context: ToolExecutionContext try: completed = subprocess.run( command, - cwd=context.root, + cwd=_execution_cwd(context), capture_output=True, text=True, timeout=timeout_seconds, @@ -2130,7 +2324,7 @@ def _run_bash(arguments: dict[str, Any], context: ToolExecutionContext) -> str: command, shell=True, executable='/bin/bash', - cwd=context.root, + cwd=_execution_cwd(context), capture_output=True, text=True, timeout=context.command_timeout_seconds, diff --git a/src/data_agent_inputs.py b/src/data_agent_inputs.py index 3bba184..70d9f04 100644 --- a/src/data_agent_inputs.py +++ b/src/data_agent_inputs.py @@ -182,10 +182,6 @@ def _resolve_input_paths(root: str | Path, paths: list[str], *, max_files: int) path = Path(raw).expanduser() candidate = path if path.is_absolute() else root_path / path candidate = candidate.resolve() - try: - candidate.relative_to(root_path) - except ValueError as exc: - raise DataAgentInputError(f'path escapes workspace root: {raw}') from exc if not candidate.exists(): raise DataAgentInputError(f'path not found: {raw}') if candidate.is_dir(): @@ -210,7 +206,7 @@ def _load_one_source( max_cell_chars: int, ) -> dict[str, Any]: suffix = path.suffix.lower() - rel_path = path.relative_to(root).as_posix() + rel_path = _display_input_path(path, root) source: dict[str, Any] = { 'path': rel_path, 'kind': suffix.lstrip('.'), @@ -238,6 +234,13 @@ def _load_one_source( return source +def _display_input_path(path: Path, root: Path) -> str: + try: + return path.relative_to(root).as_posix() + except ValueError: + return path.as_posix() + + def _load_xlsx(path: Path, source: dict[str, Any], *, max_tables: int, max_rows: int, max_cell_chars: int) -> None: try: import openpyxl # type: ignore[import-not-found] diff --git a/tests/test_agent_prompting.py b/tests/test_agent_prompting.py index e6cd969..97e8682 100644 --- a/tests/test_agent_prompting.py +++ b/tests/test_agent_prompting.py @@ -34,6 +34,9 @@ class AgentPromptingTests(unittest.TestCase): prompt = render_system_prompt(parts) self.assertIn('# 系统规则', prompt) self.assertIn('# 处理任务', prompt) + self.assertIn('# 工作空间边界', prompt) + self.assertIn('当前 session 目录是本轮任务的默认工作区', prompt) + self.assertIn('平台服务代码目录始终只读', prompt) self.assertIn('# 使用工具', prompt) self.assertIn('# Skills', prompt) self.assertIn('product-data', prompt) @@ -100,7 +103,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Plugins', prompt) + self.assertIn('# 插件', prompt) def test_prompt_builder_mentions_hook_policy_when_manifest_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -119,7 +122,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Hook Policy', prompt) + self.assertIn('# Hook 策略', prompt) def test_prompt_builder_mentions_mcp_when_manifest_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -162,7 +165,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Search', prompt) + self.assertIn('# 搜索', prompt) self.assertIn('web_search', prompt) def test_prompt_builder_mentions_remote_when_manifest_is_loaded(self) -> None: @@ -185,7 +188,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Remote', prompt) + self.assertIn('# 远程环境', prompt) def test_prompt_builder_mentions_account_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -204,7 +207,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Account', prompt) + self.assertIn('# 账号', prompt) def test_prompt_builder_mentions_ask_user_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -223,7 +226,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Ask User', prompt) + self.assertIn('# 询问用户', prompt) def test_prompt_builder_mentions_config_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -244,7 +247,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Config', prompt) + self.assertIn('# 配置', prompt) def test_prompt_builder_mentions_lsp_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -261,7 +264,7 @@ class AgentPromptingTests(unittest.TestCase): prompt = render_system_prompt(parts) self.assertIn('# LSP', prompt) - self.assertIn('Use the LSP tool', prompt) + self.assertIn('使用 LSP 工具', prompt) def test_prompt_builder_mentions_tasks_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -278,7 +281,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Tasks', prompt) + self.assertIn('# 任务', prompt) def test_prompt_builder_mentions_teams_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -297,7 +300,7 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Teams', prompt) + self.assertIn('# 团队', prompt) def test_prompt_builder_mentions_planning_when_runtime_is_loaded(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: @@ -317,4 +320,4 @@ class AgentPromptingTests(unittest.TestCase): ) prompt = render_system_prompt(parts) - self.assertIn('# Planning', prompt) + self.assertIn('# 计划', prompt) diff --git a/tests/test_data_agent_inputs.py b/tests/test_data_agent_inputs.py index 12aab24..23eb7bd 100644 --- a/tests/test_data_agent_inputs.py +++ b/tests/test_data_agent_inputs.py @@ -30,6 +30,17 @@ class DataAgentInputTests(unittest.TestCase): self.assertEqual(table['title'], 'Sheet') self.assertEqual(table['rows'][0], ['query', '预期domain', '0106-prev-domain', 'type', '备注']) + def test_load_input_sources_reads_explicit_external_file(self) -> None: + with tempfile.TemporaryDirectory() as root_dir, tempfile.TemporaryDirectory() as external_dir: + path = Path(external_dir) / 'cases.xlsx' + _write_xlsx(path) + + payload = load_input_sources(root_dir, [str(path)]) + + self.assertEqual(payload['source_count'], 1) + self.assertEqual(payload['sources'][0]['path'], path.resolve().as_posix()) + self.assertEqual(payload['sources'][0]['tables'][0]['rows'][1][0], '怎么开启查找设备') + def test_extract_case_evidence_profiles_and_extracts_rows(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: path = Path(tmp_dir) / 'cases.xlsx' diff --git a/tests/test_extended_tools.py b/tests/test_extended_tools.py index 0a00ca0..f58516d 100644 --- a/tests/test_extended_tools.py +++ b/tests/test_extended_tools.py @@ -49,6 +49,166 @@ class ExtendedToolTests(unittest.TestCase): self.assertIn('read_file', result.content) self.assertIn('write_file', result.content) + def test_grep_search_skips_generated_dirs_by_default(self) -> None: + registry = default_tool_registry() + with tempfile.TemporaryDirectory() as tmp_dir: + workspace = Path(tmp_dir) + (workspace / '.next' / 'static').mkdir(parents=True) + (workspace / '.next' / 'static' / 'bundle.js').write_text( + 'router_session_parquet should not be searched by default\n', + encoding='utf-8', + ) + (workspace / 'README.md').write_text( + 'router_session_parquet is documented here\n', + encoding='utf-8', + ) + context = build_tool_context( + AgentRuntimeConfig(cwd=workspace), + tool_registry=registry, + ) + result = execute_tool( + registry, + 'grep_search', + {'pattern': 'router_session_parquet'}, + context, + ) + + self.assertTrue(result.ok) + self.assertIn('README.md:1:', result.content) + self.assertNotIn('.next/static/bundle.js', result.content) + + def test_grep_search_truncates_very_long_lines(self) -> None: + registry = default_tool_registry() + with tempfile.TemporaryDirectory() as tmp_dir: + workspace = Path(tmp_dir) + long_line = 'prefix needle ' + ('x' * 5000) + (workspace / 'large.txt').write_text(long_line, encoding='utf-8') + context = build_tool_context( + AgentRuntimeConfig(cwd=workspace, max_output_chars=1200), + tool_registry=registry, + ) + result = execute_tool( + registry, + 'grep_search', + {'pattern': 'needle'}, + context, + ) + + self.assertTrue(result.ok) + self.assertIn('large.txt:1:', result.content) + self.assertIn('[line truncated,', result.content) + self.assertLess(len(result.content), 1200) + + def test_read_file_can_read_explicit_external_path(self) -> None: + registry = default_tool_registry() + with tempfile.TemporaryDirectory() as workspace_dir, tempfile.TemporaryDirectory() as external_dir: + external = Path(external_dir) / 'reference.txt' + external.write_text('external reference\n', encoding='utf-8') + context = build_tool_context( + AgentRuntimeConfig(cwd=Path(workspace_dir)), + tool_registry=registry, + ) + result = execute_tool( + registry, + 'read_file', + {'path': str(external)}, + context, + ) + + self.assertTrue(result.ok) + self.assertIn('external reference', result.content) + + def test_write_file_blocks_platform_code_paths(self) -> None: + registry = default_tool_registry() + with tempfile.TemporaryDirectory() as tmp_dir: + workspace = Path(tmp_dir) + (workspace / 'src').mkdir() + (workspace / 'src' / 'agent_tools.py').write_text('', encoding='utf-8') + (workspace / 'backend' / 'api').mkdir(parents=True) + (workspace / 'backend' / 'api' / 'server.py').write_text('', encoding='utf-8') + (workspace / 'frontend' / 'app').mkdir(parents=True) + scratchpad = ( + workspace + / '.port_sessions' + / 'accounts' + / 'user' + / 'sessions' + / 'thread' + / 'scratchpad' + ) + scratchpad.mkdir(parents=True) + context = build_tool_context( + AgentRuntimeConfig( + cwd=workspace, + permissions=AgentPermissions(allow_file_write=True), + ), + scratchpad_directory=scratchpad, + tool_registry=registry, + ) + + result = execute_tool( + registry, + 'write_file', + {'path': 'src/new_file.py', 'content': 'print(1)\n'}, + context, + ) + output_result = execute_tool( + registry, + 'write_file', + { + 'path': '.port_sessions/accounts/user/sessions/thread/output/report.txt', + 'content': 'ok\n', + }, + context, + ) + + self.assertFalse(result.ok) + self.assertEqual(result.metadata.get('error_kind'), 'permission_denied') + self.assertTrue(output_result.ok) + + def test_logical_session_paths_route_to_current_session(self) -> None: + registry = default_tool_registry() + with tempfile.TemporaryDirectory() as tmp_dir: + workspace = Path(tmp_dir) + scratchpad = ( + workspace + / '.port_sessions' + / 'accounts' + / 'user' + / 'sessions' + / 'thread' + / 'scratchpad' + ) + scratchpad.mkdir(parents=True) + context = build_tool_context( + AgentRuntimeConfig( + cwd=workspace, + permissions=AgentPermissions(allow_file_write=True), + ), + scratchpad_directory=scratchpad, + tool_registry=registry, + ) + write_result = execute_tool( + registry, + 'write_file', + {'path': 'output/report.md', 'content': 'session report\n'}, + context, + ) + read_result = execute_tool( + registry, + 'read_file', + {'path': 'output/report.md'}, + context, + ) + session_file_exists = (scratchpad.parent / 'output' / 'report.md').is_file() + root_file_exists = (workspace / 'output' / 'report.md').exists() + + self.assertTrue(write_result.ok, write_result.content) + self.assertTrue(read_result.ok, read_result.content) + self.assertIn('session report', read_result.content) + self.assertTrue(session_file_exists) + self.assertFalse(root_file_exists) + def test_sleep_tool_waits_briefly_and_returns_metadata(self) -> None: registry = default_tool_registry() with tempfile.TemporaryDirectory() as tmp_dir: diff --git a/tests/test_python_exec_tool.py b/tests/test_python_exec_tool.py index 0c44a6b..c7ce5e9 100644 --- a/tests/test_python_exec_tool.py +++ b/tests/test_python_exec_tool.py @@ -51,6 +51,26 @@ class PythonExecToolTests(TestCase): self.assertIn(str(scratchpad), result.content) self.assertEqual(result.metadata.get('scratchpad_directory'), str(scratchpad)) + def test_python_exec_runs_from_session_scratchpad(self) -> None: + with tempfile.TemporaryDirectory() as tmp_dir: + root = Path(tmp_dir) + scratchpad = (root / 'session' / 'scratchpad').resolve() + scratchpad.mkdir(parents=True) + config = AgentRuntimeConfig( + cwd=root, + permissions=AgentPermissions(allow_shell_commands=True), + ) + context = build_tool_context(config, scratchpad_directory=scratchpad) + result = execute_tool( + default_tool_registry(), + 'python_exec', + {'code': 'import os\nprint(os.getcwd())'}, + context, + ) + + self.assertTrue(result.ok) + self.assertIn(str(scratchpad), result.content) + def test_python_exec_prefers_user_python_env(self) -> None: with tempfile.TemporaryDirectory() as tmp_dir: root = Path(tmp_dir)