154 lines
4.7 KiB
Python
154 lines
4.7 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_is_hashed_and_device_parameters_are_preserved() -> 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": "raw-device-id",
|
|
}
|
|
|
|
|
|
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, "")
|