Skip to content

Commit 0c0642f

Browse files
committed
fix: make safety logging tests independent of root capture
1 parent fd50f83 commit 0c0642f

2 files changed

Lines changed: 81 additions & 28 deletions

File tree

tests/tools/safety/test_code_executor_wrapper.py

Lines changed: 41 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import asyncio
1010
import logging
11+
from contextlib import contextmanager
1112
from unittest.mock import Mock
1213
from unittest.mock import patch
1314

@@ -52,11 +53,35 @@ async def execute_code(
5253
return self.result
5354

5455

55-
def _enable_caplog_logger(name: str) -> None:
56+
class _CapturingHandler(logging.Handler):
57+
def __init__(self) -> None:
58+
super().__init__(level=logging.WARNING)
59+
self.records: list[logging.LogRecord] = []
60+
61+
def emit(self, record: logging.LogRecord) -> None:
62+
self.records.append(record)
63+
64+
65+
@contextmanager
66+
def _capture_logger_records(name: str):
5667
logging.disable(logging.NOTSET)
5768
target_logger = logging.getLogger(name)
69+
handler = _CapturingHandler()
70+
previous_disabled = target_logger.disabled
71+
previous_level = target_logger.level
5872
target_logger.disabled = False
59-
target_logger.propagate = True
73+
target_logger.setLevel(logging.WARNING)
74+
target_logger.addHandler(handler)
75+
try:
76+
yield handler.records
77+
finally:
78+
target_logger.removeHandler(handler)
79+
target_logger.disabled = previous_disabled
80+
target_logger.setLevel(previous_level)
81+
82+
83+
def _record_text(records: list[logging.LogRecord]) -> str:
84+
return "\n".join(record.getMessage() for record in records)
6085

6186

6287
class RecordingAuditLogger:
@@ -449,7 +474,7 @@ def test_audit_and_span_do_not_contain_secret_values(self):
449474
assert secret not in dumped_audit
450475
assert secret not in dumped_span_calls
451476

452-
def test_fail_closed_blocks_and_logs_without_exception_detail(self, caplog):
477+
def test_fail_closed_blocks_and_logs_without_exception_detail(self):
453478
policy = SafetyPolicy(fail_closed=True)
454479
audit_logger = RecordingAuditLogger()
455480
delegate = RecordingDelegate()
@@ -458,10 +483,10 @@ def test_fail_closed_blocks_and_logs_without_exception_detail(self, caplog):
458483
scanner=RaisingScanner(policy),
459484
audit_logger=audit_logger,
460485
)
461-
_enable_caplog_logger("trpc_agent_sdk.tools.safety._code_executor")
462-
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._code_executor")
463486

464-
with patch("trpc_agent_sdk.tools.safety._code_executor.set_safety_span_attributes") as mock_span:
487+
with _capture_logger_records("trpc_agent_sdk.tools.safety._code_executor") as records, patch(
488+
"trpc_agent_sdk.tools.safety._code_executor.set_safety_span_attributes"
489+
) as mock_span:
465490
result = _execute(executor, CodeExecutionInput(code="print('ok')"))
466491

467492
assert delegate.calls == []
@@ -470,10 +495,11 @@ def test_fail_closed_blocks_and_logs_without_exception_detail(self, caplog):
470495
assert "secret-token-value" not in result.output
471496
assert audit_logger.events[0].blocked is True
472497
mock_span.assert_called_once()
473-
assert "RuntimeError" in caplog.text
474-
assert "secret-token-value" not in caplog.text
498+
log_text = _record_text(records)
499+
assert "RuntimeError" in log_text
500+
assert "secret-token-value" not in log_text
475501

476-
def test_fail_open_delegates_and_skips_audit_span(self, caplog):
502+
def test_fail_open_delegates_and_skips_audit_span(self):
477503
policy = SafetyPolicy(fail_closed=False)
478504
audit_logger = RecordingAuditLogger()
479505
delegate = RecordingDelegate()
@@ -482,18 +508,19 @@ def test_fail_open_delegates_and_skips_audit_span(self, caplog):
482508
scanner=RaisingScanner(policy),
483509
audit_logger=audit_logger,
484510
)
485-
_enable_caplog_logger("trpc_agent_sdk.tools.safety._code_executor")
486-
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._code_executor")
487511

488-
with patch("trpc_agent_sdk.tools.safety._code_executor.set_safety_span_attributes") as mock_span:
512+
with _capture_logger_records("trpc_agent_sdk.tools.safety._code_executor") as records, patch(
513+
"trpc_agent_sdk.tools.safety._code_executor.set_safety_span_attributes"
514+
) as mock_span:
489515
result = _execute(executor, CodeExecutionInput(code="print('ok')"))
490516

491517
assert result.output == "delegate output"
492518
assert len(delegate.calls) == 1
493519
assert audit_logger.events == []
494520
mock_span.assert_not_called()
495-
assert "RuntimeError" in caplog.text
496-
assert "secret-token-value" not in caplog.text
521+
log_text = _record_text(records)
522+
assert "RuntimeError" in log_text
523+
assert "secret-token-value" not in log_text
497524

498525
def test_package_export_is_safety_only(self):
499526
import trpc_agent_sdk.code_executors as code_executors

tests/tools/safety/test_filter_and_audit.py

Lines changed: 40 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import asyncio
1010
import logging
11+
from contextlib import contextmanager
1112
from typing import Any
1213
from unittest.mock import patch
1314

@@ -48,11 +49,35 @@ def _ensure_safety_filter_registered() -> None:
4849
register_tool_filter("tool_safety_guard")(ToolSafetyFilter)
4950

5051

51-
def _enable_caplog_logger(name: str) -> None:
52+
class _CapturingHandler(logging.Handler):
53+
def __init__(self) -> None:
54+
super().__init__(level=logging.WARNING)
55+
self.records: list[logging.LogRecord] = []
56+
57+
def emit(self, record: logging.LogRecord) -> None:
58+
self.records.append(record)
59+
60+
61+
@contextmanager
62+
def _capture_logger_records(name: str):
5263
logging.disable(logging.NOTSET)
5364
target_logger = logging.getLogger(name)
65+
handler = _CapturingHandler()
66+
previous_disabled = target_logger.disabled
67+
previous_level = target_logger.level
5468
target_logger.disabled = False
55-
target_logger.propagate = True
69+
target_logger.setLevel(logging.WARNING)
70+
target_logger.addHandler(handler)
71+
try:
72+
yield handler.records
73+
finally:
74+
target_logger.removeHandler(handler)
75+
target_logger.disabled = previous_disabled
76+
target_logger.setLevel(previous_level)
77+
78+
79+
def _record_text(records: list[logging.LogRecord]) -> str:
80+
return "\n".join(record.getMessage() for record in records)
5681

5782

5883
class StaticScanner:
@@ -269,37 +294,38 @@ def test_shell_tools_default_to_shell_language(self):
269294
assert result == {"ok": True}
270295
assert scanner.targets[0].language == ScriptLanguage.SHELL
271296

272-
def test_fail_closed_blocks_and_logs_without_exception_detail(self, caplog):
297+
def test_fail_closed_blocks_and_logs_without_exception_detail(self):
273298
policy = SafetyPolicy(fail_closed=True)
274299
filter_ = ToolSafetyFilter(scanner=RaisingScanner(policy), audit_logger=RecordingAuditLogger())
275-
_enable_caplog_logger("trpc_agent_sdk.tools.safety._filter")
276-
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._filter")
277300

278-
result, _, calls = _run_filter(filter_, {"command": "echo ok"})
301+
with _capture_logger_records("trpc_agent_sdk.tools.safety._filter") as records:
302+
result, _, calls = _run_filter(filter_, {"command": "echo ok"})
279303

280304
assert calls == []
281305
assert result["blocked"] is True
282306
assert result["error"] == "Tool safety scan failed closed"
283307
assert result["safety_report"]["decision"] == "deny"
284-
assert "RuntimeError" in caplog.text
285-
assert "secret-token-value" not in caplog.text
308+
log_text = _record_text(records)
309+
assert "RuntimeError" in log_text
310+
assert "secret-token-value" not in log_text
286311

287-
def test_fail_open_allows_and_logs_without_exception_detail(self, caplog):
312+
def test_fail_open_allows_and_logs_without_exception_detail(self):
288313
policy = SafetyPolicy(fail_closed=False)
289314
audit_logger = RecordingAuditLogger()
290315
filter_ = ToolSafetyFilter(scanner=RaisingScanner(policy), audit_logger=audit_logger)
291-
_enable_caplog_logger("trpc_agent_sdk.tools.safety._filter")
292-
caplog.set_level(logging.WARNING, logger="trpc_agent_sdk.tools.safety._filter")
293316

294-
with patch("trpc_agent_sdk.tools.safety._filter.set_safety_span_attributes") as mock_span:
317+
with _capture_logger_records("trpc_agent_sdk.tools.safety._filter") as records, patch(
318+
"trpc_agent_sdk.tools.safety._filter.set_safety_span_attributes"
319+
) as mock_span:
295320
result, _, calls = _run_filter(filter_, {"command": "echo ok"})
296321

297322
assert result == {"ok": True}
298323
assert calls == ["handler"]
299324
assert audit_logger.events == []
300325
mock_span.assert_not_called()
301-
assert "RuntimeError" in caplog.text
302-
assert "secret-token-value" not in caplog.text
326+
log_text = _record_text(records)
327+
assert "RuntimeError" in log_text
328+
assert "secret-token-value" not in log_text
303329

304330
def test_blocked_safety_filter_prevents_later_filters(self):
305331
policy = SafetyPolicy()

0 commit comments

Comments
 (0)