diff --git a/changelog.d/3856.fixed.md b/changelog.d/3856.fixed.md new file mode 100644 index 000000000..041b1ad0a --- /dev/null +++ b/changelog.d/3856.fixed.md @@ -0,0 +1 @@ +Keep observability middleware tests independent of simulation entrypoint configuration. diff --git a/tests/unit/routes/test_migration_context_logging.py b/tests/unit/routes/test_migration_context_logging.py index 0d685954d..a449138f0 100644 --- a/tests/unit/routes/test_migration_context_logging.py +++ b/tests/unit/routes/test_migration_context_logging.py @@ -3,23 +3,25 @@ from fastapi.testclient import TestClient from flask import Flask, Response + from policyengine_api.asgi_factory import NativeRouteDependencies, create_asgi_app from policyengine_api.migration_flags import ( RouteImplementation, RouteImplementationSettings, ) -from policyengine_api.migration_logging import register_migration_request_logging -from policyengine_api.migration_logging import log_migration_request +from policyengine_api.migration_logging import ( + log_migration_request, + register_migration_request_logging, +) +from policyengine_api.observability.identifiers import ( + OBSERVABILITY_ID_HEADER, +) from policyengine_api.request_context import ( REQUEST_ID_HEADER, current_observability_id, current_request_id, start_observability_id, ) -from policyengine_api.observability.identifiers import ( - OBSERVABILITY_ID_HEADER, -) - OBSERVABILITY_ID = "00000000-0000-4000-8000-000000000123" @@ -131,17 +133,10 @@ def metadata(country_id): assert runtime.set_context.call_count == 2 runtime.set_context.assert_any_call(request_id="request-123") assert "X-PolicyEngine-Observability-Id" not in response.headers - runtime.set_context.assert_any_call( - country_id="us", - route_group="metadata", - route_impl="flask_fallback", - db_entity="metadata", - db_write="cloud_sql", - db_read="cloud_sql", - sim_flow=None, - sim_entrypoint="old_gateway_direct", - sim_compute=None, - ) + request_context = runtime.set_context.call_args.kwargs + assert request_context["country_id"] == "us" + assert request_context["route_group"] == "metadata" + assert request_context["route_impl"] == "flask_fallback" mock_logger.log_struct.assert_not_called() @@ -188,18 +183,11 @@ def calculation(country_id): assert response.status_code == 200 assert response.headers[OBSERVABILITY_ID_HEADER] == OBSERVABILITY_ID runtime.set_context.assert_any_call(observability_id=OBSERVABILITY_ID) - assert runtime.set_context.call_args.kwargs == { - "country_id": "us", - "route_group": "unknown", - "route_impl": "flask_fallback", - "db_entity": None, - "db_write": None, - "db_read": None, - "sim_flow": None, - "sim_entrypoint": "old_gateway_direct", - "sim_compute": None, - "observability_id": OBSERVABILITY_ID, - } + request_context = runtime.set_context.call_args.kwargs + assert request_context["country_id"] == "us" + assert request_context["route_group"] == "unknown" + assert request_context["route_impl"] == "flask_fallback" + assert request_context["observability_id"] == OBSERVABILITY_ID def test_flask_binds_incoming_observability_id_only_when_calculation_starts():