diff --git a/tests/tools/safety/test_scanner.py b/tests/tools/safety/test_scanner.py index d7e1af97..67dfc7ad 100644 --- a/tests/tools/safety/test_scanner.py +++ b/tests/tools/safety/test_scanner.py @@ -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(): diff --git a/tests/tools/safety/test_wrapper.py b/tests/tools/safety/test_wrapper.py index 54e00646..1bd17f20 100644 --- a/tests/tools/safety/test_wrapper.py +++ b/tests/tools/safety/test_wrapper.py @@ -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() diff --git a/trpc_agent_sdk/tools/safety/_filter.py b/trpc_agent_sdk/tools/safety/_filter.py index 3d2ef66f..4c2f39e3 100644 --- a/trpc_agent_sdk/tools/safety/_filter.py +++ b/trpc_agent_sdk/tools/safety/_filter.py @@ -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") @@ -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) @@ -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: @@ -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" @@ -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) diff --git a/trpc_agent_sdk/tools/safety/_scanner.py b/trpc_agent_sdk/tools/safety/_scanner.py index 5fde3002..2fa7006f 100644 --- a/trpc_agent_sdk/tools/safety/_scanner.py +++ b/trpc_agent_sdk/tools/safety/_scanner.py @@ -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 @@ -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) @@ -74,6 +97,8 @@ 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, @@ -81,7 +106,7 @@ def scan(self, request: ToolScriptScanRequest) -> SafetyReport: "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( @@ -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,