Skip to content
9 changes: 8 additions & 1 deletion tests/tools/safety/test_scanner.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,14 @@ def test_unknown_language_scans_python_and_bash_rules():
assert report.language == "unknown"
assert report.decision == Decision.DENY
assert "BASH_RECURSIVE_DELETE" in _rule_ids(report)
assert "PY_PARSE_ERROR_REVIEW" in _rule_ids(report)
assert "PY_PARSE_ERROR_REVIEW" not in _rule_ids(report)


def test_unknown_language_keeps_python_findings_when_parse_succeeds():
report = _scanner().scan_script("while True:\n pass\n", "unknown")

assert "PY_INFINITE_LOOP" in _rule_ids(report)
assert "PY_PARSE_ERROR_REVIEW" not in _rule_ids(report)


def test_language_aliases_are_normalized():
Expand Down
131 changes: 131 additions & 0 deletions tests/tools/safety/test_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,137 @@ async def test_filter_extracts_command_as_bash():
assert result.rsp["language"] == "bash"


@pytest.mark.asyncio
async def test_filter_scans_all_script_like_fields():
safety_filter = ToolSafetyFilter()
result = FilterResult()

await safety_filter._before(
None,
{
"script": "echo ok",
"command": "rm -rf /",
"language": "bash",
"tool_name": "shell_tool",
},
result,
)

assert result.is_continue is False
assert result.rsp["decision"] == "deny"
assert any(finding["rule_id"] == "BASH_RECURSIVE_DELETE" for finding in result.rsp["findings"])


@pytest.mark.asyncio
async def test_filter_scans_mixed_language_fields_without_language():
safety_filter = ToolSafetyFilter()
result = FilterResult()

await safety_filter._before(
None,
{
"code": "print('ok')",
"command": "rm -rf /",
"tool_name": "mixed_tool",
},
result,
)

assert result.is_continue is False
assert result.rsp["language"] == "mixed"
assert result.rsp["decision"] == "deny"
assert any(finding["rule_id"] == "BASH_RECURSIVE_DELETE" for finding in result.rsp["findings"])
assert all(finding["rule_id"] != "PY_PARSE_ERROR_REVIEW" for finding in result.rsp["findings"])


@pytest.mark.asyncio
async def test_filter_allows_safe_mixed_language_fields_in_strict_mode():
safety_filter = ToolSafetyFilter(block_on_review=True)
result = FilterResult()

await safety_filter._before(
None,
{
"code": "print('ok')",
"command": "echo ok",
"tool_name": "mixed_tool",
},
result,
)

assert result.is_continue is True
assert result.rsp["language"] == "mixed"
assert result.rsp["decision"] == "allow"
assert all(finding["rule_id"] != "PY_PARSE_ERROR_REVIEW" for finding in result.rsp["findings"])


@pytest.mark.asyncio
async def test_filter_scans_execution_context_once_for_mixed_fields():
safety_filter = ToolSafetyFilter()
result = FilterResult()

await safety_filter._before(
None,
{
"python_code": "print('ok')",
"command": "echo ok",
"tool_metadata": {
"timeout": 301
},
"tool_name": "mixed_tool",
},
result,
)

matching_findings = [
finding for finding in result.rsp["findings"] if finding["rule_id"] == "RESOURCE_TIMEOUT_LIMIT_EXCEEDED"
]
assert result.rsp["language"] == "mixed"
assert result.rsp["decision"] == "needs_human_review"
assert len(matching_findings) == 1
assert result.rsp["telemetry_attributes"]["tool.safety.scan_id"] == result.rsp["scan_id"]


@pytest.mark.asyncio
async def test_filter_deduplicates_findings_across_segments():
safety_filter = ToolSafetyFilter()
result = FilterResult()

await safety_filter._before(
None,
{
"command": "rm -rf /",
"script": "rm -rf /",
"tool_name": "custom",
},
result,
)

matching_findings = [finding for finding in result.rsp["findings"] if finding["rule_id"] == "BASH_RECURSIVE_DELETE"]
assert result.rsp["decision"] == "deny"
assert len(matching_findings) == 1
assert all(finding["rule_id"] != "PY_PARSE_ERROR_REVIEW" for finding in result.rsp["findings"])


@pytest.mark.asyncio
async def test_filter_allows_safe_unknown_bash_script_in_strict_mode():
safety_filter = ToolSafetyFilter(block_on_review=True)
result = FilterResult()

await safety_filter._before(
None,
{
"script": "git status",
"tool_name": "custom",
},
result,
)

assert result.is_continue is True
assert result.rsp["decision"] == "allow"
assert all(finding["rule_id"] != "PY_PARSE_ERROR_REVIEW" for finding in result.rsp["findings"])


@pytest.mark.asyncio
async def test_filter_extracts_python_code_language():
safety_filter = ToolSafetyFilter()
Expand Down
100 changes: 69 additions & 31 deletions trpc_agent_sdk/tools/safety/_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,9 @@
from ._types import SafetyReport
from ._types import ToolScriptScanRequest

_SCRIPT_ARG_KEYS = ("script", "code", "command", "cmd", "python_code", "bash_code")
_PYTHON_ARG_KEYS = ("python_code", )
_BASH_ARG_KEYS = ("command", "cmd", "bash_code")
_GENERIC_ARG_KEYS = ("script", )
_LANGUAGE_ARG_KEYS = ("language", "lang")
_COMMAND_ARGS_KEYS = ("command_args", "args", "argv")

Expand Down Expand Up @@ -60,20 +62,11 @@ async def _before(self, ctx: AgentContext, req: Any, rsp: FilterResult):
self._current_report.set(None)
if not isinstance(req, dict):
return None
script = _extract_script(req)
if not script:
return None
tool_name = _extract_tool_name(req)
request = ToolScriptScanRequest(
script=script,
language=_extract_language(req, tool_name),
command_args=_extract_command_args(req),
cwd=str(req.get("cwd", "")),
env=dict(req.get("env", {}) or {}),
tool_name=tool_name,
tool_metadata=dict(req.get("tool_metadata", {}) or {}),
)
report = self.scanner.scan(request)
requests = _extract_scan_requests(req, tool_name)
if not requests:
return None
report = self.scanner.scan_segments(requests)
should_block = report.decision == Decision.DENY or (self.block_on_review
and report.decision == Decision.NEEDS_HUMAN_REVIEW)
report.set_blocked(should_block)
Expand All @@ -98,25 +91,58 @@ async def _after(self, ctx: AgentContext, req: Any, rsp: FilterResult):
return None


def _extract_script(req: dict[str, Any]) -> str:
for key in _SCRIPT_ARG_KEYS:
value = req.get(key)
if isinstance(value, str) and value.strip():
return value
def _extract_scan_requests(req: dict[str, Any], tool_name: str) -> list[ToolScriptScanRequest]:
grouped_parts: dict[str, list[str]] = {}

for key in _PYTHON_ARG_KEYS:
_add_script_part(grouped_parts, "python", req.get(key))
for key in _BASH_ARG_KEYS:
_add_script_part(grouped_parts, "bash", req.get(key))

generic_language = _extract_language(req, tool_name)
for key in _GENERIC_ARG_KEYS:
_add_script_part(grouped_parts, generic_language, req.get(key))
code_language = _extract_explicit_language(req) or "python"
_add_script_part(grouped_parts, code_language, req.get("code"))

code_blocks = req.get("code_blocks")
if isinstance(code_blocks, list):
parts: list[str] = []
for block in code_blocks:
if isinstance(block, dict):
code = block.get("code", "")
language = block.get("language", "")
else:
code = getattr(block, "code", "")
if isinstance(code, str) and code:
parts.append(code)
if parts:
return "\n".join(parts)
return ""
language = getattr(block, "language", "")
block_language = _canonical_language(language) if isinstance(language,
str) and language.strip() else generic_language
_add_script_part(grouped_parts, block_language, code)

command_args = _extract_command_args(req)
cwd = str(req.get("cwd", ""))
env = dict(req.get("env", {}) or {})
tool_metadata = dict(req.get("tool_metadata", {}) or {})
requests: list[ToolScriptScanRequest] = []
for language, parts in grouped_parts.items():
requests.append(
ToolScriptScanRequest(
script="\n".join(parts),
language=language,
command_args=command_args,
cwd=cwd,
env=env,
tool_name=tool_name,
tool_metadata=tool_metadata,
))
return requests


def _add_script_part(grouped_parts: dict[str, list[str]], language: str, value: Any) -> None:
if not isinstance(value, str) or not value.strip():
return
parts = grouped_parts.setdefault(_canonical_language(language), [])
if value not in parts:
parts.append(value)


def _extract_tool_name(req: dict[str, Any]) -> str:
Expand All @@ -130,15 +156,18 @@ def _extract_tool_name(req: dict[str, Any]) -> str:
return "unknown_tool"


def _extract_language(req: dict[str, Any], tool_name: str) -> str:
def _extract_explicit_language(req: dict[str, Any]) -> str:
for key in _LANGUAGE_ARG_KEYS:
value = req.get(key)
if isinstance(value, str) and value.strip():
return value.strip().lower()
if isinstance(req.get("python_code"), str) or "code" in req:
return "python"
if isinstance(req.get("bash_code"), str) or "command" in req or "cmd" in req:
return "bash"
return _canonical_language(value)
return ""


def _extract_language(req: dict[str, Any], tool_name: str) -> str:
explicit_language = _extract_explicit_language(req)
if explicit_language:
return explicit_language
lowered_tool_name = tool_name.lower()
if "python" in lowered_tool_name:
return "python"
Expand All @@ -147,6 +176,15 @@ def _extract_language(req: dict[str, Any], tool_name: str) -> str:
return "unknown"


def _canonical_language(language: str) -> str:
normalized = (language or "unknown").strip().lower()
if normalized in {"py", "python3"}:
return "python"
if normalized in {"shell", "sh"}:
return "bash"
return normalized


def _extract_command_args(req: dict[str, Any]) -> list[str]:
for key in _COMMAND_ARGS_KEYS:
value = req.get(key)
Expand Down
55 changes: 40 additions & 15 deletions trpc_agent_sdk/tools/safety/_scanner.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from pathlib import Path

from ._policy import ToolSafetyPolicy
from ._rules import _dedupe_findings
from ._rules import _finding
from ._rules import SENSITIVE_NAME_RE
from ._rules import scan_bash_script
Expand Down Expand Up @@ -50,21 +51,43 @@ def register_rule(self, rule: SafetyRule) -> None:
self.custom_rules.append(rule)

def scan(self, request: ToolScriptScanRequest) -> SafetyReport:
return self.scan_segments([request])

def scan_segments(self, requests: Iterable[ToolScriptScanRequest]) -> SafetyReport:
"""Scan language-specific script segments as one tool invocation."""
request_list = list(requests)
if not request_list:
raise ValueError("At least one script segment is required.")

started = time.perf_counter()
language = self._normalize_language(request.language)
sanitized = self._env_contains_sensitive_keys(request.env)
_, script_sanitized = sanitize_text(request.script, limit=max(len(request.script), 1))
sanitized = sanitized or script_sanitized
base_request = request_list[0]
languages: list[str] = []
findings: list[RiskFinding] = []
sanitized = False

for request in request_list:
language = self._normalize_language(request.language)
languages.append(language)
sanitized = sanitized or self._env_contains_sensitive_keys(request.env)
_, script_sanitized = sanitize_text(request.script, limit=max(len(request.script), 1))
sanitized = sanitized or script_sanitized

if language == "python":
findings.extend(scan_python_script(request.script, self.policy))
elif language in {"bash", "sh", "shell"}:
findings.extend(scan_bash_script(request.script, self.policy))
else:
findings.extend(scan_bash_script(request.script, self.policy))
python_findings = scan_python_script(request.script, self.policy)
# An unknown segment may be valid non-Python input; keep Python
# findings only when AST parsing succeeds.
if not any(finding.rule_id == "PY_PARSE_ERROR_REVIEW" for finding in python_findings):
findings.extend(python_findings)

if language == "python":
findings = scan_python_script(request.script, self.policy)
elif language in {"bash", "sh", "shell"}:
findings = scan_bash_script(request.script, self.policy)
else:
findings = scan_bash_script(request.script, self.policy)
findings.extend(scan_python_script(request.script, self.policy))
findings.extend(self._scan_execution_context(request))
findings.extend(self._scan_custom_rules(request))
findings.extend(self._scan_execution_context(base_request))
for request in request_list:
findings.extend(self._scan_custom_rules(request))
findings = _dedupe_findings(findings)

decision = aggregate_decision(findings)
risk_level = max_risk_level(findings)
Expand All @@ -74,14 +97,16 @@ def scan(self, request: ToolScriptScanRequest) -> SafetyReport:
summary = self._build_summary(decision.value, risk_level.value, rule_ids)
scan_id = str(uuid.uuid4())
timestamp = datetime.now(timezone.utc).isoformat()
normalized_languages = list(dict.fromkeys(languages))
language = normalized_languages[0] if len(normalized_languages) == 1 else "mixed"
telemetry_attributes = {
"tool.safety.scan_id": scan_id,
"tool.safety.decision": decision.value,
"tool.safety.risk_level": risk_level.value,
"tool.safety.rule_id": ",".join(rule_ids[:10]),
"tool.safety.blocked": blocked,
"tool.safety.sanitized": sanitized,
"tool.safety.tool_name": request.tool_name,
"tool.safety.tool_name": base_request.tool_name,
"tool.safety.duration_ms": elapsed_ms,
}
return SafetyReport(
Expand All @@ -90,7 +115,7 @@ def scan(self, request: ToolScriptScanRequest) -> SafetyReport:
decision=decision,
risk_level=risk_level,
findings=findings,
tool_name=request.tool_name,
tool_name=base_request.tool_name,
language=language,
elapsed_ms=elapsed_ms,
sanitized=sanitized,
Expand Down
Loading