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, "")