Add tests for Rich markup escape safety in formatter
- Introduced a new test file `test_rich_escape.py` to validate the safety of Rich markup escaping in the ToolResultFormatter. - Added tests to ensure that tool names and error messages containing brackets do not cause crashes during formatting.
This commit is contained in:
File diff suppressed because one or more lines are too long
|
After Width: | Height: | Size: 100 KiB |
File diff suppressed because one or more lines are too long
|
After Width: | Height: | Size: 100 KiB |
@@ -9,6 +9,7 @@ dist/
|
||||
build/
|
||||
*.egg
|
||||
*.pytest_cache/
|
||||
.coverage
|
||||
|
||||
# Environment
|
||||
.env
|
||||
|
||||
@@ -64,8 +64,6 @@ SUBAGENTS_CONFIG = Path(__file__).parent / "subagent.yaml"
|
||||
# Initialization
|
||||
# =============================================================================
|
||||
|
||||
# Get current date
|
||||
current_date = datetime.now().strftime("%Y-%m-%d")
|
||||
|
||||
# Generate system prompt with limits
|
||||
SYSTEM_PROMPT = get_system_prompt(
|
||||
@@ -170,7 +168,7 @@ def _build_base_kwargs(base_backend, base_middleware):
|
||||
subs = load_subagents(
|
||||
SUBAGENTS_CONFIG,
|
||||
tool_registry=tool_registry,
|
||||
prompt_refs=prompt_refs,
|
||||
prompt_refs=_build_prompt_refs(),
|
||||
)
|
||||
_inject_subagent_middleware(subs)
|
||||
return dict(
|
||||
@@ -206,7 +204,7 @@ def load_mcp_and_build_kwargs(base_backend, base_middleware):
|
||||
subs = load_subagents(
|
||||
SUBAGENTS_CONFIG,
|
||||
tool_registry=registry,
|
||||
prompt_refs=prompt_refs,
|
||||
prompt_refs=_build_prompt_refs(),
|
||||
)
|
||||
|
||||
_inject_subagent_middleware(subs)
|
||||
@@ -228,9 +226,13 @@ def load_mcp_and_build_kwargs(base_backend, base_middleware):
|
||||
)
|
||||
|
||||
|
||||
prompt_refs = {
|
||||
"RESEARCHER_INSTRUCTIONS": RESEARCHER_INSTRUCTIONS.format(date=current_date),
|
||||
}
|
||||
def _build_prompt_refs() -> dict:
|
||||
"""Build prompt references with the current date (not frozen at import)."""
|
||||
return {
|
||||
"RESEARCHER_INSTRUCTIONS": RESEARCHER_INSTRUCTIONS.format(
|
||||
date=datetime.now().strftime("%Y-%m-%d"),
|
||||
),
|
||||
}
|
||||
|
||||
base_middleware = [
|
||||
ToolErrorHandlerMiddleware(),
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import os
|
||||
import re
|
||||
import shlex
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
@@ -24,7 +25,6 @@ _SYSTEM_PATH_PREFIXES = (
|
||||
|
||||
# Dangerous patterns that could escape the workspace
|
||||
BLOCKED_PATTERNS = [
|
||||
r'\.\.', # ../ directory traversal
|
||||
r'~/', # home directory
|
||||
r'\bcd\s+/', # cd to absolute path
|
||||
r'\brm\s+-rf\s+/', # rm -rf with absolute path
|
||||
@@ -42,6 +42,38 @@ BLOCKED_COMMANDS = [
|
||||
]
|
||||
|
||||
|
||||
def _split_shell_commands(command: str) -> list[str]:
|
||||
"""Split a compound shell command into individual base commands.
|
||||
|
||||
Handles &&, ||, ;, and | operators. Returns base command names.
|
||||
"""
|
||||
base_commands: list[str] = []
|
||||
# Split by sequential operators first
|
||||
for segment in re.split(r'\s*(?:&&|\|\||;)\s*', command):
|
||||
# Then split by pipe
|
||||
for pipe_seg in segment.split("|"):
|
||||
pipe_seg = pipe_seg.strip()
|
||||
if not pipe_seg:
|
||||
continue
|
||||
try:
|
||||
tokens = shlex.split(pipe_seg)
|
||||
except ValueError:
|
||||
tokens = pipe_seg.split()
|
||||
if tokens:
|
||||
base_commands.append(tokens[0])
|
||||
return base_commands
|
||||
|
||||
|
||||
def _has_traversal_component(command: str) -> bool:
|
||||
"""Check if command contains '..' as a path component (not substring)."""
|
||||
from pathlib import PurePosixPath
|
||||
|
||||
for token in command.split():
|
||||
if ".." in PurePosixPath(token).parts:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def validate_command(command: str) -> str | None:
|
||||
"""
|
||||
Validate a shell command for safety.
|
||||
@@ -49,7 +81,15 @@ def validate_command(command: str) -> str | None:
|
||||
Returns:
|
||||
None if command is safe, error message string if blocked.
|
||||
"""
|
||||
# Check for directory traversal and dangerous patterns
|
||||
# Check for '..' path traversal as a path component
|
||||
if _has_traversal_component(command):
|
||||
return (
|
||||
"Command blocked: contains '..' path traversal. "
|
||||
"All commands must operate within the workspace directory. "
|
||||
"Use relative paths (e.g., './file.py') instead."
|
||||
)
|
||||
|
||||
# Check for dangerous patterns
|
||||
for pattern in BLOCKED_PATTERNS:
|
||||
if re.search(pattern, command):
|
||||
return (
|
||||
@@ -58,11 +98,11 @@ def validate_command(command: str) -> str | None:
|
||||
f"Use relative paths (e.g., './file.py') instead."
|
||||
)
|
||||
|
||||
# Check for dangerous commands
|
||||
for cmd in BLOCKED_COMMANDS:
|
||||
if re.search(rf'\b{cmd}\b', command):
|
||||
# Check for dangerous commands (pipeline-aware)
|
||||
for base_cmd in _split_shell_commands(command):
|
||||
if base_cmd in BLOCKED_COMMANDS:
|
||||
return (
|
||||
f"Command blocked: '{cmd}' is not allowed in sandbox mode. "
|
||||
f"Command blocked: '{base_cmd}' is not allowed in sandbox mode. "
|
||||
f"Only standard development commands are permitted."
|
||||
)
|
||||
|
||||
@@ -328,7 +368,7 @@ class CustomSandboxBackend(LocalShellBackend):
|
||||
|
||||
return super()._resolve_path(key)
|
||||
|
||||
def execute(self, command: str) -> ExecuteResponse:
|
||||
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||
"""
|
||||
Execute shell command in sandbox environment.
|
||||
|
||||
@@ -362,4 +402,4 @@ class CustomSandboxBackend(LocalShellBackend):
|
||||
)
|
||||
|
||||
# Delegate to parent for subprocess execution
|
||||
return super().execute(command)
|
||||
return super().execute(command, timeout=timeout)
|
||||
|
||||
@@ -10,6 +10,8 @@ from typing import Any, Optional
|
||||
import typer # type: ignore[import-untyped]
|
||||
from rich.table import Table
|
||||
|
||||
from rich.markup import escape
|
||||
|
||||
from ..stream.display import console
|
||||
from ..paths import ensure_dirs, set_workspace_root
|
||||
from ._app import app, config_app, mcp_app, channel_app
|
||||
@@ -213,9 +215,9 @@ def config_set(
|
||||
from ..config import set_config_value
|
||||
|
||||
if set_config_value(key, value):
|
||||
console.print(f"[green]Set {key}[/green]")
|
||||
console.print(f"[green]Set {escape(key)}[/green]")
|
||||
else:
|
||||
console.print(f"[red]Invalid key: {key}[/red]")
|
||||
console.print(f"[red]Invalid key: {escape(key)}[/red]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
|
||||
@@ -572,7 +574,7 @@ def _configure_logging():
|
||||
if record.levelno == logging.WARNING:
|
||||
# Use Rich console to print dim warning
|
||||
msg = record.getMessage()
|
||||
console.print(f"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] [dim]{msg}[/dim]")
|
||||
console.print(f"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] [dim]{escape(msg)}[/dim]")
|
||||
else:
|
||||
super().emit(record)
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ from prompt_toolkit.auto_suggest import AutoSuggestFromHistory # type: ignore[i
|
||||
from prompt_toolkit.formatted_text import HTML # type: ignore[import-untyped]
|
||||
from prompt_toolkit.shortcuts import CompleteStyle # type: ignore[import-untyped]
|
||||
from prompt_toolkit.styles import Style as PtStyle # type: ignore[import-untyped]
|
||||
from rich.markup import escape
|
||||
from rich.table import Table
|
||||
from rich.text import Text
|
||||
|
||||
@@ -236,11 +237,11 @@ def cmd_interactive(
|
||||
if len(similar) == 1:
|
||||
return similar[0]
|
||||
if len(similar) > 1:
|
||||
console.print(f"[yellow]Ambiguous thread ID '{tid}'. Matches:[/yellow]")
|
||||
console.print(f"[yellow]Ambiguous thread ID '{escape(tid)}'. Matches:[/yellow]")
|
||||
for s in similar:
|
||||
console.print(f" [cyan]{s}[/cyan]")
|
||||
return None
|
||||
console.print(f"[red]Thread '{tid}' not found.[/red]")
|
||||
console.print(f"[red]Thread '{escape(tid)}' not found.[/red]")
|
||||
return None
|
||||
|
||||
async def _cmd_threads():
|
||||
@@ -701,7 +702,7 @@ def cmd_interactive(
|
||||
state["running"] = False
|
||||
break
|
||||
else:
|
||||
console.print(f"[red]Error: {e}[/red]")
|
||||
console.print(f"[red]Error: {escape(str(e))}[/red]")
|
||||
finally:
|
||||
queue_task.cancel()
|
||||
try:
|
||||
|
||||
@@ -4,6 +4,7 @@ Async generator that streams events from an agent graph,
|
||||
plus helpers for processing AI message chunks and tool results.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import mimetypes
|
||||
import os
|
||||
@@ -260,15 +261,19 @@ async def stream_agent_events(
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
if message:
|
||||
content_blocks.append({"type": "text", "text": message})
|
||||
def _read_file_b64(path: str) -> str:
|
||||
with open(path, "rb") as fh:
|
||||
return base64.b64encode(fh.read()).decode("ascii")
|
||||
|
||||
file_refs: list[str] = []
|
||||
for path in media:
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
if ext in _IMAGE_EXTS and os.path.isfile(path):
|
||||
fsize = os.path.getsize(path)
|
||||
is_image = ext in _IMAGE_EXTS and await asyncio.to_thread(os.path.isfile, path)
|
||||
if is_image:
|
||||
fsize = await asyncio.to_thread(os.path.getsize, path)
|
||||
if fsize <= _MAX_INLINE_SIZE:
|
||||
mime = mimetypes.guess_type(path)[0] or "image/png"
|
||||
with open(path, "rb") as fh:
|
||||
b64 = base64.b64encode(fh.read()).decode("ascii")
|
||||
b64 = await asyncio.to_thread(_read_file_b64, path)
|
||||
content_blocks.append({"type": "image_url", "image_url": {
|
||||
"url": f"data:{mime};base64,{b64}",
|
||||
}})
|
||||
|
||||
@@ -9,6 +9,7 @@ from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any, List
|
||||
|
||||
from rich.markup import escape
|
||||
from rich.panel import Panel
|
||||
from rich.syntax import Syntax
|
||||
from rich.text import Text
|
||||
@@ -123,7 +124,7 @@ class ToolResultFormatter:
|
||||
display = truncate(content, max_length)
|
||||
return [Panel(
|
||||
Text(display, style="green"),
|
||||
title=f"{name} OK",
|
||||
title=f"{escape(name)} OK",
|
||||
border_style="green",
|
||||
)]
|
||||
|
||||
@@ -131,7 +132,7 @@ class ToolResultFormatter:
|
||||
display = truncate(content, max_length)
|
||||
return [Panel(
|
||||
Text(display, style="red"),
|
||||
title=f"{name} FAILED",
|
||||
title=f"{escape(name)} FAILED",
|
||||
border_style="red",
|
||||
)]
|
||||
|
||||
@@ -155,7 +156,7 @@ class ToolResultFormatter:
|
||||
display = truncate(content, max_length)
|
||||
return [Panel(
|
||||
Markdown(display),
|
||||
title=f"{name}",
|
||||
title=escape(name),
|
||||
border_style="cyan dim",
|
||||
)]
|
||||
|
||||
|
||||
@@ -1,14 +1,15 @@
|
||||
<!-- Add logo here -->
|
||||
<h1 align="center">
|
||||
<img src="./assets/EvoScientist_logo.png" alt="EvoScientist Logo" height="27" style="position: relative; top: 1px;"/>
|
||||
<strong>EvoScientist</strong>
|
||||
</h1>
|
||||
|
||||
<div align="center">
|
||||
<picture>
|
||||
<source media="(prefers-color-scheme: light)" srcset=".github/images/logo-dark.svg">
|
||||
<source media="(prefers-color-scheme: dark)" srcset=".github/images/logo-light.svg">
|
||||
<img alt="EvoScientist Logo" src=".github/images/logo-dark.svg" width="80%">
|
||||
</picture>
|
||||
</a>
|
||||
</div>
|
||||
|
||||
<div align="center">
|
||||
|
||||
<a href="https://git.io/typing-svg"><img src="https://readme-typing-svg.demolab.com?font=Fira+Code&pause=1000&width=435&lines=Towards+Self-Evolving+AI+Scientists+for+End-to-End+Scientific+Discovery" alt="Typing SVG" /></a>
|
||||
|
||||
[](https://pypi.org/project/EvoScientist/)
|
||||
[]()
|
||||
[]()
|
||||
@@ -16,6 +17,8 @@
|
||||
<!-- []()
|
||||
[]() -->
|
||||
|
||||
<a href="https://git.io/typing-svg"><img src="https://readme-typing-svg.demolab.com?font=Fira+Code&pause=1000&width=435&lines=Towards+Self-Evolving+AI+Scientists+for+End-to-End+Scientific+Discovery" alt="Typing SVG" /></a>
|
||||
|
||||
</div>
|
||||
|
||||
## 🔥 News
|
||||
|
||||
+3
-3
@@ -16,9 +16,9 @@ classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
]
|
||||
dependencies = [
|
||||
"deepagents>=0.4.1",
|
||||
"langchain>=1.2.7",
|
||||
"langchain-anthropic>=1.3.1",
|
||||
"deepagents>=0.4.3",
|
||||
"langchain>=1.2.10",
|
||||
"langchain-anthropic>=1.3.3",
|
||||
"langchain-openai>=0.3",
|
||||
"langchain-nvidia-ai-endpoints>=0.3",
|
||||
"langchain-google-genai>=4.2",
|
||||
|
||||
@@ -297,3 +297,73 @@ class TestExecuteStderr:
|
||||
resp = backend.execute("echo ok")
|
||||
assert resp.exit_code == 0
|
||||
assert "Exit code:" not in resp.output
|
||||
|
||||
|
||||
# === execute() timeout kwarg ===
|
||||
|
||||
class TestExecuteTimeout:
|
||||
def test_execute_accepts_timeout_kwarg(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True)
|
||||
resp = backend.execute("echo hello", timeout=60)
|
||||
assert resp.exit_code == 0
|
||||
assert "hello" in resp.output
|
||||
|
||||
def test_execute_timeout_none_uses_default(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True)
|
||||
resp = backend.execute("echo ok", timeout=None)
|
||||
assert resp.exit_code == 0
|
||||
|
||||
def test_execute_accepts_timeout_introspection(self):
|
||||
from deepagents.backends.protocol import execute_accepts_timeout
|
||||
execute_accepts_timeout.cache_clear()
|
||||
assert execute_accepts_timeout(CustomSandboxBackend) is True
|
||||
|
||||
|
||||
# === '..' traversal false-positive fix ===
|
||||
|
||||
class TestTraversalFalsePositiveFix:
|
||||
def test_dotdot_in_filename_allowed(self):
|
||||
assert validate_command("echo foo..bar.txt") is None
|
||||
|
||||
def test_dotdot_path_component_still_blocked(self):
|
||||
result = validate_command("cat ../secret")
|
||||
assert result is not None
|
||||
assert "blocked" in result.lower()
|
||||
|
||||
def test_dotdot_nested_still_blocked(self):
|
||||
result = validate_command("cat foo/../../etc/passwd")
|
||||
assert result is not None
|
||||
|
||||
|
||||
# === Pipeline command validation ===
|
||||
|
||||
class TestPipelineCommandValidation:
|
||||
def test_pipe_blocked_command(self):
|
||||
"""sudo after pipe should be caught."""
|
||||
result = validate_command("echo hi | sudo tee /etc/passwd")
|
||||
assert result is not None
|
||||
assert "sudo" in result
|
||||
|
||||
def test_chained_blocked_command(self):
|
||||
"""chmod after && should be caught."""
|
||||
result = validate_command("echo ok && chmod 777 file")
|
||||
assert result is not None
|
||||
assert "chmod" in result
|
||||
|
||||
def test_semicolon_blocked_command(self):
|
||||
"""dd after ; should be caught."""
|
||||
result = validate_command("echo start ; dd if=/dev/zero of=disk")
|
||||
assert result is not None
|
||||
assert "dd" in result
|
||||
|
||||
def test_safe_pipe_allowed(self):
|
||||
"""Normal pipes should be fine."""
|
||||
assert validate_command("cat file.txt | grep pattern") is None
|
||||
|
||||
def test_safe_chain_allowed(self):
|
||||
"""Normal && chains should be fine."""
|
||||
assert validate_command("mkdir build && cd build") is None
|
||||
|
||||
def test_quoted_pipe_not_split(self):
|
||||
"""Pipe inside quotes is not a shell operator."""
|
||||
assert validate_command("echo 'hello | world'") is None
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
"""Tests for Rich markup escape safety in formatter."""
|
||||
|
||||
from EvoScientist.stream.formatter import ToolResultFormatter
|
||||
|
||||
|
||||
class TestRichMarkupEscape:
|
||||
def test_panel_title_with_brackets(self):
|
||||
"""Tool name with brackets shouldn't crash formatter."""
|
||||
formatter = ToolResultFormatter()
|
||||
result = formatter.format("arr[0]", "[OK]\nDone", max_length=800)
|
||||
# Should not raise — content is safely rendered
|
||||
assert result.elements
|
||||
|
||||
def test_error_format_with_brackets(self):
|
||||
"""Error content with markup-like text shouldn't crash."""
|
||||
formatter = ToolResultFormatter()
|
||||
result = formatter.format("test", "Error: expected [int] got [str]", max_length=800)
|
||||
assert result.elements
|
||||
Reference in New Issue
Block a user