Files
ss-tools/agent/tests/test_agent/test_middleware.py
root 632b730fff chore: migrate GRACE-Poly anchors to hierarchical dotted naming
Systematic rename of all semantic anchors (#region, [DEF], @RELATION)
across 1400+ files — backend Python, frontend Svelte/TS, specs, docs:
- Flat anchors become Namespace.Module.Entity
- @RELATION references updated to match new anchor paths
- Zero business logic changes
2026-07-22 11:48:15 +03:00

283 lines
11 KiB
Python

# #region Test.AgentChat.Middleware [C:3] [TYPE Module] [SEMANTICS test,agent,middleware,audit]
# @BRIEF Tests for agent/middleware.py — log_tool_event, emit_lifecycle_event, extract_trace_id_from_request.
# @RELATION BINDS_TO -> [AgentChat.Middleware]
from pathlib import Path
import sys
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "src"))
import uuid
from unittest.mock import MagicMock, patch
import pytest
# #region Test.AgentChat.TestEmitLifecycleEvent [C:2] [TYPE Function]
# @BRIEF Test emit_lifecycle_event for correct event type and payload.
class TestEmitLifecycleEvent:
def test_emits_event_with_correct_type_and_payload(self):
from ss_tools.agent.middleware import emit_lifecycle_event
with patch("ss_tools.agent.middleware.logger.reason") as mock_reason:
emit_lifecycle_event(
"AGENT_REQUEST_STARTED",
conversation_id="conv-1",
user_id="user-1",
environment_id="prod",
action="new",
)
mock_reason.assert_called_once()
args, kwargs = mock_reason.call_args
assert args[0] == "AGENT_REQUEST_STARTED"
payload = kwargs["payload"]
assert payload["conversation_id"] == "conv-1"
assert payload["user_id"] == "user-1"
assert payload["environment_id"] == "prod"
assert payload["action"] == "new"
assert kwargs["extra"]["src"] == "AgentChat.Lifecycle"
def test_filters_none_payload_values(self):
from ss_tools.agent.middleware import emit_lifecycle_event
with patch("ss_tools.agent.middleware.logger.reason") as mock_reason:
emit_lifecycle_event(
"AGENT_REQUEST_COMPLETED",
conversation_id="conv-1",
is_resume=None,
tool_names=None,
)
payload = mock_reason.call_args[1]["payload"]
assert "conversation_id" in payload
assert "is_resume" not in payload
assert "tool_names" not in payload
def test_never_includes_sensitive_fields(self):
"""Verify that sensitive field names are never in payload schema."""
from ss_tools.agent.middleware import emit_lifecycle_event
with patch("ss_tools.agent.middleware.logger.reason") as mock_reason:
emit_lifecycle_event(
"AGENT_REQUEST_STARTED",
conversation_id="conv-1",
user_id="user-1",
)
payload = mock_reason.call_args[1]["payload"]
forbidden = {"jwt", "token", "password", "secret", "message", "prompt", "file", "user_message"}
payload_keys = set(k.lower() for k in payload)
assert not (payload_keys & forbidden), f"Found forbidden key in payload: {payload_keys & forbidden}"
def test_filters_forbidden_lifecycle_fields_even_if_caller_passes_them(self):
"""Lifecycle helper enforces its no-sensitive-data invariant at runtime."""
from ss_tools.agent.middleware import emit_lifecycle_event
with patch("ss_tools.agent.middleware.logger.reason") as mock_reason:
emit_lifecycle_event(
"AGENT_REQUEST_COMPLETED",
conversation_id="conv-1",
jwt="secret",
message="private request",
files=["private.pdf"],
raw_output="private result",
)
payload = mock_reason.call_args.kwargs["payload"]
assert payload == {"conversation_id": "conv-1"}
# #endregion Test.AgentChat.TestEmitLifecycleEvent
# #region Test.AgentChat.TestExtractTraceIdFromRequest [C:2] [TYPE Function]
# @BRIEF Test extract_trace_id_from_request with valid/invalid/missing X-Trace-ID headers.
class TestExtractTraceIdFromRequest:
def make_request(self, headers: dict | None = None) -> MagicMock:
req = MagicMock()
req.headers = headers or {}
return req
def test_extracts_valid_uuid4_from_header(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
valid_id = uuid.uuid4().hex
req = self.make_request({"X-Trace-ID": valid_id})
with patch("ss_tools.agent.middleware.set_trace_id") as mock_set:
result = extract_trace_id_from_request(req)
assert result == valid_id
mock_set.assert_called_once_with(valid_id)
def test_case_insensitive_header(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
valid_id = uuid.uuid4().hex
req = self.make_request({"x-trace-id": valid_id})
with patch("ss_tools.agent.middleware.set_trace_id") as mock_set:
result = extract_trace_id_from_request(req)
assert result == valid_id
mock_set.assert_called_once_with(valid_id)
def test_seeds_when_header_missing(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
req = self.make_request({"authorization": "Bearer xyz"})
with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed:
result = extract_trace_id_from_request(req)
assert result == "new-trace"
mock_seed.assert_called_once()
def test_seeds_when_header_empty(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
req = self.make_request({"X-Trace-ID": ""})
with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed:
result = extract_trace_id_from_request(req)
assert result == "new-trace"
mock_seed.assert_called_once()
def test_seeds_on_invalid_uuid_format(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
req = self.make_request({"X-Trace-ID": "not-a-uuid-at-all"})
with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed:
result = extract_trace_id_from_request(req)
assert result == "new-trace"
mock_seed.assert_called_once()
def test_seeds_on_non_v4_uuid(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
# UUID v1
v1_id = "550e8400-e29b-11d1-a716-446655440000"
req = self.make_request({"X-Trace-ID": v1_id})
with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed:
result = extract_trace_id_from_request(req)
assert result == "new-trace"
mock_seed.assert_called_once()
def test_handles_request_without_headers(self):
from ss_tools.agent.middleware import extract_trace_id_from_request
req = MagicMock(spec=[]) # no headers attr
del req.headers
with patch("ss_tools.agent.middleware.seed_trace_id", return_value="new-trace") as mock_seed:
result = extract_trace_id_from_request(req)
assert result == "new-trace"
mock_seed.assert_called_once()
# #endregion Test.AgentChat.TestExtractTraceIdFromRequest
# #region Test.AgentChat.TestLogToolEvent [C:2] [TYPE Function]
# @BRIEF Test log_tool_event for various event types.
class TestLogToolEvent:
@pytest.mark.asyncio
async def test_logs_tool_start(self):
from ss_tools.agent.middleware import log_tool_event
event = {
"event": "on_tool_start",
"name": "migrate",
"data": {"input": {"dashboard_id": "42"}},
}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value="user-token"):
await log_tool_event(event, "conv-1")
# No exception = success
@pytest.mark.asyncio
async def test_logs_tool_end(self):
from ss_tools.agent.middleware import log_tool_event
event = {
"event": "on_tool_end",
"name": "migrate",
"data": {"output": "success"},
}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value="user-token"):
await log_tool_event(event, "conv-1")
@pytest.mark.asyncio
async def test_logs_tool_error(self):
from ss_tools.agent.middleware import log_tool_event
event = {
"event": "on_tool_error",
"name": "migrate",
"data": {"error": "Connection failed"},
}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value="user-token"):
await log_tool_event(event, "conv-1")
@pytest.mark.asyncio
async def test_logs_without_user_jwt(self):
from ss_tools.agent.middleware import log_tool_event
event = {
"event": "on_tool_start",
"name": "test_tool",
"data": {"input": {}},
}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value=None):
await log_tool_event(event, "conv-1")
@pytest.mark.asyncio
async def test_handles_missing_data_key(self):
from ss_tools.agent.middleware import log_tool_event
event = {"event": "on_tool_start", "name": "test_tool"}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value=None):
await log_tool_event(event, "conv-1")
@pytest.mark.asyncio
async def test_handles_unknown_event_kind(self):
from ss_tools.agent.middleware import log_tool_event
event = {"event": "on_custom_event", "name": "custom"}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value=None):
await log_tool_event(event, "conv-1")
@pytest.mark.asyncio
async def test_truncates_long_input(self):
from ss_tools.agent.middleware import log_tool_event
long_input = "x" * 1000
event = {
"event": "on_tool_start",
"name": "big_tool",
"data": {"input": long_input},
}
with patch("ss_tools.agent.middleware.get_user_jwt", return_value="token"):
await log_tool_event(event, "conv-1")
@pytest.mark.asyncio
async def test_includes_trace_id(self):
from ss_tools.agent.middleware import log_tool_event
test_trace_id = "abc123"
event = {
"event": "on_tool_start",
"name": "trace_test",
"data": {"input": {"key": "val"}},
}
with (
patch("ss_tools.agent.middleware.get_user_jwt", return_value="token"),
patch("ss_tools.agent.middleware.get_trace_id", return_value=test_trace_id),
patch("ss_tools.agent.middleware.logger.reason") as mock_reason,
):
await log_tool_event(event, "conv-1")
payload = mock_reason.call_args[1]["payload"]
assert payload["trace_id"] == test_trace_id
@pytest.mark.asyncio
async def test_handles_empty_trace_id(self):
from ss_tools.agent.middleware import log_tool_event
event = {
"event": "on_tool_start",
"name": "no_trace",
"data": {"input": {}},
}
with (
patch("ss_tools.agent.middleware.get_user_jwt", return_value="token"),
patch("ss_tools.agent.middleware.get_trace_id", return_value=""),
patch("ss_tools.agent.middleware.logger.reason"),
):
await log_tool_event(event, "conv-1")
# No exception = success
# #endregion Test.AgentChat.TestLogToolEvent
# #endregion Test.AgentChat.Middleware