Skip to content

Commit fa9d32d

Browse files
committed
fix: align tool safety CI behavior
1 parent c903c6d commit fa9d32d

12 files changed

Lines changed: 355 additions & 556 deletions

File tree

examples/tool_safety_guard/run_safety_scan.py

Lines changed: 12 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -10,24 +10,18 @@
1010
import sys
1111
from pathlib import Path
1212

13-
try:
14-
from trpc_agent_sdk.tools.safety._cli_helpers import FIXTURE_GENERATED_AT
15-
from trpc_agent_sdk.tools.safety._cli_helpers import format_mismatches
16-
from trpc_agent_sdk.tools.safety._cli_helpers import load_policy
17-
from trpc_agent_sdk.tools.safety._cli_helpers import load_samples
18-
from trpc_agent_sdk.tools.safety._cli_helpers import scan_samples
19-
from trpc_agent_sdk.tools.safety._cli_helpers import write_audit_log
20-
from trpc_agent_sdk.tools.safety._cli_helpers import write_json_report
21-
except ModuleNotFoundError: # pragma: no cover - exercised by direct script execution in fresh checkouts.
22-
repo_root = Path(__file__).resolve().parents[2]
23-
sys.path.insert(0, str(repo_root))
24-
from trpc_agent_sdk.tools.safety._cli_helpers import FIXTURE_GENERATED_AT
25-
from trpc_agent_sdk.tools.safety._cli_helpers import format_mismatches
26-
from trpc_agent_sdk.tools.safety._cli_helpers import load_policy
27-
from trpc_agent_sdk.tools.safety._cli_helpers import load_samples
28-
from trpc_agent_sdk.tools.safety._cli_helpers import scan_samples
29-
from trpc_agent_sdk.tools.safety._cli_helpers import write_audit_log
30-
from trpc_agent_sdk.tools.safety._cli_helpers import write_json_report
13+
REPO_ROOT = Path(__file__).resolve().parents[2]
14+
repo_root = str(REPO_ROOT)
15+
sys.path[:] = [path for path in sys.path if path != repo_root]
16+
sys.path.insert(0, repo_root)
17+
18+
from trpc_agent_sdk.tools.safety._cli_helpers import FIXTURE_GENERATED_AT
19+
from trpc_agent_sdk.tools.safety._cli_helpers import format_mismatches
20+
from trpc_agent_sdk.tools.safety._cli_helpers import load_policy
21+
from trpc_agent_sdk.tools.safety._cli_helpers import load_samples
22+
from trpc_agent_sdk.tools.safety._cli_helpers import scan_samples
23+
from trpc_agent_sdk.tools.safety._cli_helpers import write_audit_log
24+
from trpc_agent_sdk.tools.safety._cli_helpers import write_json_report
3125

3226
EXAMPLE_DIR = Path(__file__).resolve().parent
3327
POLICY_PATH = EXAMPLE_DIR / "tool_safety_policy.yaml"

tests/tools/safety/test_code_executor_wrapper.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,6 @@
88

99
import asyncio
1010
import logging
11-
from typing import Any
1211
from unittest.mock import Mock
1312
from unittest.mock import patch
1413

@@ -53,6 +52,12 @@ async def execute_code(
5352
return self.result
5453

5554

55+
def _enable_caplog_logger(name: str) -> None:
56+
target_logger = logging.getLogger(name)
57+
target_logger.disabled = False
58+
target_logger.propagate = True
59+
60+
5661
class RecordingAuditLogger:
5762
def __init__(self):
5863
self.events: list[SafetyAuditEvent] = []
@@ -452,6 +457,7 @@ def test_fail_closed_blocks_and_logs_without_exception_detail(self, caplog):
452457
scanner=RaisingScanner(policy),
453458
audit_logger=audit_logger,
454459
)
460+
_enable_caplog_logger("trpc_agent_sdk.tools.safety._code_executor")
455461
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._code_executor")
456462

457463
with patch("trpc_agent_sdk.tools.safety._code_executor.set_safety_span_attributes") as mock_span:
@@ -475,6 +481,7 @@ def test_fail_open_delegates_and_skips_audit_span(self, caplog):
475481
scanner=RaisingScanner(policy),
476482
audit_logger=audit_logger,
477483
)
484+
_enable_caplog_logger("trpc_agent_sdk.tools.safety._code_executor")
478485
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._code_executor")
479486

480487
with patch("trpc_agent_sdk.tools.safety._code_executor.set_safety_span_attributes") as mock_span:

tests/tools/safety/test_filter_and_audit.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from trpc_agent_sdk.context import new_agent_context
1616
from trpc_agent_sdk.filter import BaseFilter
1717
from trpc_agent_sdk.filter import get_tool_filter
18+
from trpc_agent_sdk.filter import register_tool_filter
1819
from trpc_agent_sdk.filter import run_filters
1920
from trpc_agent_sdk.tools._context_var import reset_tool_var
2021
from trpc_agent_sdk.tools._context_var import set_tool_var
@@ -42,6 +43,17 @@ def emit(self, event: SafetyAuditEvent) -> None:
4243
self.events.append(event)
4344

4445

46+
def _ensure_safety_filter_registered() -> None:
47+
if get_tool_filter("tool_safety_guard") is None:
48+
register_tool_filter("tool_safety_guard")(ToolSafetyFilter)
49+
50+
51+
def _enable_caplog_logger(name: str) -> None:
52+
target_logger = logging.getLogger(name)
53+
target_logger.disabled = False
54+
target_logger.propagate = True
55+
56+
4557
class StaticScanner:
4658
def __init__(
4759
self,
@@ -190,6 +202,8 @@ def test_review_blocks_by_default_and_can_be_nonblocking(self):
190202
assert audit_logger.events[0].blocked is False
191203

192204
def test_registry_name_uses_default_policy_instance(self):
205+
_ensure_safety_filter_registered()
206+
193207
registered = get_tool_filter("tool_safety_guard")
194208

195209
assert isinstance(registered, ToolSafetyFilter)
@@ -257,6 +271,7 @@ def test_shell_tools_default_to_shell_language(self):
257271
def test_fail_closed_blocks_and_logs_without_exception_detail(self, caplog):
258272
policy = SafetyPolicy(fail_closed=True)
259273
filter_ = ToolSafetyFilter(scanner=RaisingScanner(policy), audit_logger=RecordingAuditLogger())
274+
_enable_caplog_logger("trpc_agent_sdk.tools.safety._filter")
260275
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._filter")
261276

262277
result, _, calls = _run_filter(filter_, {"command": "echo ok"})
@@ -272,6 +287,7 @@ def test_fail_open_allows_and_logs_without_exception_detail(self, caplog):
272287
policy = SafetyPolicy(fail_closed=False)
273288
audit_logger = RecordingAuditLogger()
274289
filter_ = ToolSafetyFilter(scanner=RaisingScanner(policy), audit_logger=audit_logger)
290+
_enable_caplog_logger("trpc_agent_sdk.tools.safety._filter")
275291
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._filter")
276292

277293
with patch("trpc_agent_sdk.tools.safety._filter.set_safety_span_attributes") as mock_span:

trpc_agent_sdk/tools/safety/_audit.py

Lines changed: 10 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -23,17 +23,16 @@ def _ordered_rule_ids(report: SafetyReport) -> list[str]:
2323

2424

2525
def build_safety_audit_event(
26-
report: SafetyReport,
27-
*,
28-
tool_name: str = "",
29-
cwd: str = "",
30-
function_call_id: str = "",
31-
agent_name: str = "",
26+
report: SafetyReport,
27+
*,
28+
tool_name: str = "",
29+
cwd: str = "",
30+
function_call_id: str = "",
31+
agent_name: str = "",
3232
) -> SafetyAuditEvent:
3333
"""Build an audit-safe summary event from a safety report."""
3434

35-
resolved_tool_name = tool_name or str(
36-
report.metadata.get("target_tool", "") or "")
35+
resolved_tool_name = tool_name or str(report.metadata.get("target_tool", "") or "")
3736
return SafetyAuditEvent(
3837
tool_name=resolved_tool_name,
3938
decision=report.decision,
@@ -54,6 +53,7 @@ def build_safety_audit_event(
5453

5554
class SafetyAuditLogger:
5655
"""Optional JSONL writer for safety audit events."""
56+
5757
def __init__(self, path: str | Path | None = None, enabled: bool = True):
5858
self.path = Path(path) if path is not None else None
5959
self.enabled = enabled
@@ -73,13 +73,11 @@ def emit(self, event: SafetyAuditEvent) -> None:
7373
logger.warning("Failed to write safety audit event: %s", ex)
7474

7575

76-
def set_safety_span_attributes(report: SafetyReport, *,
77-
tool_name: str = "") -> None:
76+
def set_safety_span_attributes(report: SafetyReport, *, tool_name: str = "") -> None:
7877
"""Best-effort write of safety summary fields onto the current OTel span."""
7978

8079
rule_ids = _ordered_rule_ids(report)
81-
resolved_tool_name = tool_name or str(
82-
report.metadata.get("target_tool", "") or "")
80+
resolved_tool_name = tool_name or str(report.metadata.get("target_tool", "") or "")
8381
attributes: dict[str, str | bool | int | float] = {
8482
"tool.safety.decision": report.decision.value,
8583
"tool.safety.risk_level": report.risk_level.value,

0 commit comments

Comments
 (0)