Files
han-app/codebase/backend/api-backend/tests/unit/test_audit_context.py
T

155 lines
4.8 KiB
Python

import hashlib
import uuid
from types import SimpleNamespace
import pytest
from starlette.requests import Request
from app.main import (
app,
client_ip,
request_trace_id,
required_user_audit_context,
user_agent_hash,
)
from app.schemas import Device
from app.services import AuditContext, DomainError, audit, device_snapshot
def request(
*,
peer: str = "172.18.0.5",
forwarded: str | None = None,
traceparent: str | None = None,
user_agent: str | None = None,
ux_session_id: str | None = None,
trusted: str = "172.16.0.0/12",
) -> Request:
headers: list[tuple[bytes, bytes]] = []
for name, value in (
("x-forwarded-for", forwarded),
("traceparent", traceparent),
("user-agent", user_agent),
("x-ux-session-id", ux_session_id),
):
if value is not None:
headers.append((name.encode(), value.encode()))
settings = SimpleNamespace(trusted_proxy_cidrs=trusted)
app = SimpleNamespace(state=SimpleNamespace(settings=settings))
return Request(
{
"type": "http",
"method": "GET",
"path": "/",
"headers": headers,
"client": (peer, 12345),
"server": ("test", 443),
"scheme": "https",
"query_string": b"",
"app": app,
}
)
def test_client_ip_only_trusts_forwarded_header_from_configured_proxy() -> None:
assert client_ip(request(forwarded="203.0.113.10")) == "203.0.113.10"
assert (
client_ip(request(peer="198.51.100.7", forwarded="203.0.113.10"))
== "198.51.100.7"
)
assert client_ip(request(peer="not-an-ip")) is None
def test_trace_id_uses_valid_w3c_header_and_generates_fallback() -> None:
trace_id = "1" * 32
assert request_trace_id(request(traceparent=f"00-{trace_id}-{'2' * 16}-01")) == trace_id
fallback = request_trace_id(request(traceparent="invalid"))
assert len(fallback) == 32
int(fallback, 16)
def test_user_agent_and_device_are_hashed_without_raw_identifiers() -> None:
agent = "Example Browser/1.0"
assert user_agent_hash(request(user_agent=agent)) == hashlib.sha256(agent.encode()).hexdigest()
snapshot = device_snapshot(
Device(platform="web", app_version="1.2.3", device_id="raw-device-id")
)
assert snapshot == {
"platform": "web",
"app_version": "1.2.3",
"device_id_hash": hashlib.sha256(b"raw-device-id").hexdigest(),
}
assert "raw-device-id" not in str(snapshot)
def test_audit_copies_request_context_and_bounded_metadata() -> None:
session_id = uuid.uuid4()
user_id = uuid.uuid4()
context = AuditContext(
request_id=str(uuid.uuid4()),
trace_id="a" * 32,
ux_session_id=session_id,
user_agent_hash="b" * 64,
client_ip="203.0.113.10",
)
event = audit(
"dialog.created",
context,
user_id,
"dialog",
uuid.uuid4(),
metadata={"status": "open"},
)
assert event.user_id == user_id
assert event.ux_session_id == session_id
assert event.trace_id == "a" * 32
assert event.user_agent_hash == "b" * 64
assert event.metadata_json == {"status": "open"}
assert event.outcome == "success"
assert not hasattr(event, "ip")
failed = audit("message.failed", context, user_id, outcome="failed")
assert failed.outcome == "failed"
def test_post_session_routes_publish_required_ux_header_in_openapi() -> None:
operation = app.openapi()["paths"]["/api/v1/dialogs"]["post"]
ux_header = next(
parameter
for parameter in operation["parameters"]
if parameter["name"] == "X-Ux-Session-Id"
)
assert ux_header["in"] == "header"
assert ux_header["required"] is True
@pytest.mark.asyncio
async def test_user_audit_context_requires_existing_owned_session() -> None:
session_id = uuid.uuid4()
user = SimpleNamespace(id=uuid.uuid4())
incoming = request(
forwarded="203.0.113.10",
ux_session_id=str(session_id),
user_agent="Example Browser/1.0",
)
incoming.state.request_id = str(uuid.uuid4())
incoming.state.trace_id = "c" * 32
incoming.state.user_agent_hash = user_agent_hash(incoming)
class Db:
async def scalar(self, _query):
return session_id
context = await required_user_audit_context(incoming, Db(), user, str(session_id))
assert context.ux_session_id == session_id
assert context.client_ip == "203.0.113.10"
missing = request(forwarded="203.0.113.10")
missing.state.request_id = str(uuid.uuid4())
missing.state.trace_id = "d" * 32
missing.state.user_agent_hash = None
with pytest.raises(DomainError, match="X-Ux-Session-Id is required"):
await required_user_audit_context(missing, Db(), user, "")