"""Unit tests for request context helpers.""" import asyncio import base64 import json from collections.abc import Iterator from typing import Any, cast from uuid import UUID import pytest from fastapi import FastAPI from fastapi.testclient import TestClient from starlette.datastructures import Headers from starlette.types import ASGIApp, Message, Receive, Scope, Send from contributor.api.context import ( RequestContext, exclude, get_correlation_id, get_request_context, request_context_from_headers, reset_correlation_id, reset_request_context, set_correlation_id, set_request_context, ) from contributor.api.context.asgi.middleware import ( CorrelationIdMiddleware, RequestContextMiddleware, ) @pytest.fixture(autouse=True) def clear_context_state() -> Iterator[None]: request_token = set_request_context(None) correlation_token = set_correlation_id(None) yield reset_request_context(request_token) reset_correlation_id(correlation_token) def test_exclude_removes_requested_keys() -> None: data = {"a": 1, "b": 2, "c": 3} result = exclude(data, keys=("b", "c")) assert result == {"a": 1} def test_request_context_dict_returns_all_fields_by_default() -> None: context = RequestContext( requestor_service_name="ows-contributor", identity_id="id-1", profile_id=123, ) result = context.dict() assert result == { "requestor_service_name": "ows-contributor", "identity_id": "id-1", "identity_uuid": None, "profile_type": None, "profile_id": 123, "profile_uuid": None, "authorization": None, } def test_request_context_dict_excludes_empty_values_when_requested() -> None: context = RequestContext( requestor_service_name="ows-contributor", profile_id=0, authorization="", ) result = context.dict(exclude_empty=True) assert result == {"requestor_service_name": "ows-contributor"} def test_request_context_dict_excludes_specified_keys() -> None: context = RequestContext( requestor_service_name="ows-contributor", identity_id="id-1", profile_uuid=UUID("3f059f11-8454-4a27-a2b5-9804c9339286"), ) result = context.dict(exclude_keys=("identity_id", "profile_uuid")) assert "identity_id" not in result assert "profile_uuid" not in result assert result["requestor_service_name"] == "ows-contributor" def test_request_context_dict_prefers_exclude_keys_when_both_flags_set() -> None: context = RequestContext( requestor_service_name="ows-contributor", authorization="", ) result = context.dict(exclude_empty=True, exclude_keys=("identity_id",)) assert "identity_id" not in result assert "authorization" in result def test_set_get_and_reset_request_context() -> None: context = RequestContext(identity_id="abc") token = set_request_context(context) try: assert get_request_context() == context finally: reset_request_context(token) assert get_request_context() is None def test_set_get_and_reset_correlation_id() -> None: token = set_correlation_id("corr-id") try: assert get_correlation_id() == "corr-id" finally: reset_correlation_id(token) assert get_correlation_id() is None def test_request_context_from_headers_parses_values() -> None: result = request_context_from_headers( Headers( { "orchard-requestor-service": "test-service", "orchard-identity-id": "identity-id", "orchard-identity-uuid": "identity-uuid", "orchard-profile-type": "LabelProfile", "orchard-profile-id": "123", "orchard-profile-uuid": "3f059f11-8454-4a27-a2b5-9804c9339286", "authorization": "Bearer token", } ) ) assert result == RequestContext( requestor_service_name="ows_contributor", identity_id="identity-id", identity_uuid="identity-uuid", profile_type="LabelProfile", profile_id=123, profile_uuid=UUID("3f059f11-8454-4a27-a2b5-9804c9339286"), authorization="Bearer token", ) def test_request_context_from_headers_decodes_jwt_claims() -> None: def _b64(data: dict[str, Any]) -> str: encoded = base64.urlsafe_b64encode(json.dumps(data).encode("utf-8")).decode( "ascii" ) return encoded.rstrip("=") header = _b64({"alg": "none", "typ": "JWT"}) payload = _b64( { "https://grass.theorchard.com/profiles": [ {"profile_type": "ContentProfile", "profile_id": 318696} ], "https://grass.theorchard.com/user_metadata": { "orchardIdentityId": "10436b38-5e11-472d-b6a4-bf1ee2b1b438" }, } ) token = f"{header}.{payload}.sig" result = request_context_from_headers(Headers({"authorization": f"Bearer {token}"})) assert result.requestor_service_name == "ows_contributor" assert result.profile_type == "ContentProfile" assert result.profile_id == 318696 assert result.identity_id == "10436b38-5e11-472d-b6a4-bf1ee2b1b438" assert result.identity_uuid == "10436b38-5e11-472d-b6a4-bf1ee2b1b438" def test_request_context_from_headers_uses_none_for_invalid_profile_id() -> None: result = request_context_from_headers(Headers({"orchard-profile-id": "12x"})) assert result.profile_id is None def _build_test_app() -> FastAPI: app = FastAPI() app.add_middleware(CorrelationIdMiddleware) app.add_middleware(RequestContextMiddleware) @app.get("/context") def context_endpoint() -> dict[str, Any]: request_context = get_request_context() return { "correlation_id": get_correlation_id(), "request_context": request_context.dict() if request_context else None, } return app def test_request_context_middleware_sets_context_for_http_requests() -> None: app = _build_test_app() with TestClient(app) as client: response = client.get( "/context", headers={ "Correlation-Id": "req-corr-id", "orchard-requestor-service": "test-service", "orchard-profile-type": "LabelProfile", "orchard-profile-id": "456", }, ) assert response.status_code == 200 payload = response.json() assert payload["correlation_id"] == "req-corr-id" assert payload["request_context"]["requestor_service_name"] == "ows_contributor" assert payload["request_context"]["profile_type"] == "LabelProfile" assert payload["request_context"]["profile_id"] == 456 assert get_request_context() is None assert get_correlation_id() is None def test_correlation_id_middleware_generates_response_header_when_missing() -> None: app = _build_test_app() with TestClient(app) as client: response = client.get("/context") assert response.status_code == 200 correlation_id = response.headers["Correlation-Id"] assert UUID(correlation_id) def test_request_context_middleware_passthrough_for_non_http() -> None: called = False async def app(scope: Scope, receive: Receive, send: Send) -> None: nonlocal called called = True middleware = RequestContextMiddleware(cast(ASGIApp, app)) async def receive() -> Message: return {"type": "websocket.connect"} async def send(message: Message) -> None: return None asyncio.run(middleware({"type": "websocket", "headers": []}, receive, send)) assert called is True def test_correlation_id_middleware_passthrough_for_non_http() -> None: called = False async def app(scope: Scope, receive: Receive, send: Send) -> None: nonlocal called called = True middleware = CorrelationIdMiddleware(cast(ASGIApp, app)) async def receive() -> Message: return {"type": "websocket.connect"} async def send(message: Message) -> None: return None asyncio.run(middleware({"type": "websocket", "headers": []}, receive, send)) assert called is True