diff --git a/.github/scripts/deploy_cloud_run_candidate.sh b/.github/scripts/deploy_cloud_run_candidate.sh index cb19b9bf3..2e5d80773 100755 --- a/.github/scripts/deploy_cloud_run_candidate.sh +++ b/.github/scripts/deploy_cloud_run_candidate.sh @@ -8,6 +8,7 @@ cloud_run_set_defaults bash .github/scripts/validate_cloud_run_deploy_env.sh env_vars=( + "APP_ENVIRONMENT=${DEPLOYMENT_ENVIRONMENT}" "POLICYENGINE_DB_INSTANCE_CONNECTION_NAME=${POLICYENGINE_DB_INSTANCE_CONNECTION_NAME}" "POLICYENGINE_DB_USER=${POLICYENGINE_DB_USER:-policyengine}" "POLICYENGINE_DB_NAME=${POLICYENGINE_DB_NAME:-policyengine}" @@ -32,6 +33,14 @@ env_vars=( "RUNTIME_CACHE_MODE=deployed" "RUNTIME_CACHE_ENVIRONMENT=${CLOUD_RUN_RUNTIME_CACHE_ENVIRONMENT}" "RUNTIME_CACHE_SERVICE=api" + "OBSERVABILITY_SERVICE_NAMESPACE=${OBSERVABILITY_SERVICE_NAMESPACE}" + "OBSERVABILITY_TRACE_PROJECT_ID=${OBSERVABILITY_TRACE_PROJECT_ID}" + "OTEL_EXPORTER_OTLP_ENDPOINT=${OTEL_EXPORTER_OTLP_ENDPOINT}" + "OTEL_EXPORTER_OTLP_PROTOCOL=grpc" + "OTEL_TRACES_EXPORTER=otlp" + "OTEL_METRICS_EXPORTER=otlp" + "OTEL_TRACES_SAMPLER_ARG=1.0" + "POLICYENGINE_OTEL_GOOGLE_AUDIENCE=${POLICYENGINE_OTEL_GOOGLE_AUDIENCE}" "V2_SUPABASE_PROJECT_REF=${V2_SUPABASE_PROJECT_REF}" "V2_SUPABASE_ENVIRONMENT=${V2_SUPABASE_ENVIRONMENT}" "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE=${V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE}" diff --git a/.github/scripts/validate_cloud_run_deploy_env.sh b/.github/scripts/validate_cloud_run_deploy_env.sh index 732b607d9..8647f1235 100755 --- a/.github/scripts/validate_cloud_run_deploy_env.sh +++ b/.github/scripts/validate_cloud_run_deploy_env.sh @@ -36,6 +36,10 @@ cloud_run_require_env \ CLOUD_RUN_VPC_NETWORK \ CLOUD_RUN_VPC_SUBNET \ CLOUD_RUN_VPC_EGRESS \ + OBSERVABILITY_SERVICE_NAMESPACE \ + OBSERVABILITY_TRACE_PROJECT_ID \ + OTEL_EXPORTER_OTLP_ENDPOINT \ + POLICYENGINE_OTEL_GOOGLE_AUDIENCE \ V2_SUPABASE_PROJECT_REF \ V2_SUPABASE_ENVIRONMENT \ V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE \ diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index 1390bb8cb..6dffa5e2b 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -273,6 +273,10 @@ jobs: GATEWAY_AUTH_AUDIENCE: ${{ secrets.GATEWAY_AUTH_AUDIENCE }} GATEWAY_AUTH_CLIENT_ID: ${{ secrets.GATEWAY_AUTH_CLIENT_ID }} GATEWAY_AUTH_CLIENT_SECRET_RESOURCE: ${{ secrets.GATEWAY_AUTH_CLIENT_SECRET_RESOURCE }} + OBSERVABILITY_SERVICE_NAMESPACE: ${{ vars.OBSERVABILITY_SERVICE_NAMESPACE }} + OBSERVABILITY_TRACE_PROJECT_ID: ${{ vars.OBSERVABILITY_TRACE_PROJECT_ID }} + OTEL_EXPORTER_OTLP_ENDPOINT: ${{ vars.OBSERVABILITY_OTLP_ENDPOINT }} + POLICYENGINE_OTEL_GOOGLE_AUDIENCE: ${{ vars.OBSERVABILITY_OTLP_GOOGLE_AUDIENCE }} - name: Resolve exact Cloud Run staging candidate id: candidate run: bash .github/scripts/resolve_cloud_run_candidate_state.sh >> "$GITHUB_OUTPUT" @@ -497,6 +501,10 @@ jobs: GATEWAY_AUTH_AUDIENCE: ${{ secrets.GATEWAY_AUTH_AUDIENCE }} GATEWAY_AUTH_CLIENT_ID: ${{ secrets.GATEWAY_AUTH_CLIENT_ID }} GATEWAY_AUTH_CLIENT_SECRET_RESOURCE: ${{ secrets.GATEWAY_AUTH_CLIENT_SECRET_RESOURCE }} + OBSERVABILITY_SERVICE_NAMESPACE: ${{ vars.OBSERVABILITY_SERVICE_NAMESPACE }} + OBSERVABILITY_TRACE_PROJECT_ID: ${{ vars.OBSERVABILITY_TRACE_PROJECT_ID }} + OTEL_EXPORTER_OTLP_ENDPOINT: ${{ vars.OBSERVABILITY_OTLP_ENDPOINT }} + POLICYENGINE_OTEL_GOOGLE_AUDIENCE: ${{ vars.OBSERVABILITY_OTLP_GOOGLE_AUDIENCE }} - name: Resolve exact Cloud Run production candidate id: candidate run: bash .github/scripts/resolve_cloud_run_candidate_state.sh >> "$GITHUB_OUTPUT" diff --git a/AGENTS.md b/AGENTS.md index b29042324..f3034a66d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -25,6 +25,10 @@ migration revisions, read When adding or moving API v2 route, service, or database-access modules, read `docs/engineering/skills/v2-code-organization.md`. +When changing correlation identifiers, telemetry transport, spans, stage +names, logging, or observability failure handling, read +`docs/engineering/skills/observability.md`. + When modifying the `Makefile`, read `docs/engineering/skills/repository-maintenance.md`. diff --git a/changelog.d/3847.changed.md b/changelog.d/3847.changed.md new file mode 100644 index 000000000..1cd7dafbd --- /dev/null +++ b/changelog.d/3847.changed.md @@ -0,0 +1,3 @@ +Route API v1 structured logs, traces, and metrics through the explicit +policyengine-observability 3.x runtime and propagate request context to +the simulation entry service. diff --git a/docs/engineering/skills/README.md b/docs/engineering/skills/README.md index 1d6f31070..768847909 100644 --- a/docs/engineering/skills/README.md +++ b/docs/engineering/skills/README.md @@ -16,6 +16,9 @@ Current skills: - `github-prs.md`: PR workflow and migration PR handoff expectations. - `migration_contracts.md`: API v2 migration route contracts, route-group metadata, generated migration artifacts, and quality guards. +- `observability.md`: identifier ownership, HTTP and simulation transport, + persistence, registered runtime stages, trace boundaries, and failure + isolation. - `repository-maintenance.md`: mandatory Makefile target and `.PHONY` maintenance rules. - `testing.md`: focused test commands and dependency boundaries for migration diff --git a/docs/engineering/skills/observability.md b/docs/engineering/skills/observability.md new file mode 100644 index 000000000..b44ba8242 --- /dev/null +++ b/docs/engineering/skills/observability.md @@ -0,0 +1,122 @@ +# Observability engineering rules + +Read this file before changing request correlation, calculation correlation, +logging, traces, metrics, simulation requests, or runtime stage names. + +## Identifier registry + +Use each identifier for its defined scope: + +| Identifier | Scope | Created by | Durable | +| --- | --- | --- | --- | +| `request_id` | One HTTP request | HTTP request instrumentation | No | +| `observability_id` | One complete household calculation or society report | The first service that accepts the calculation | Yes for asynchronous reports | +| `submission_claim_id` | One attempt to acquire ownership of a simulation submission | API v1 economy service | Only as submission metadata | +| `job_id` | One annual simulation job | Simulation API | Yes, as functional job state | +| `batch_job_id` | One budget window simulation job | Simulation API | Yes, as functional batch state | +| `evaluation_id` | One Stage 12 comparison report | Stage 12 report construction | Yes, as functional report state | +| `simulation_execution_id` | One Stage 12 baseline or reform simulation | Stage 12 coordinator | Yes, as functional simulation state | + +`observability_id` is a canonical UUID string used only to query diagnostic +records. Do not use it for idempotency, cache ownership, authorization, +database identity, routing, or calculation behavior. Do not create one for +health checks, metadata reads, invalid requests, missing reports, or ordinary +status requests. + +## HTTP lifecycle + +Transport `observability_id` only in +`X-PolicyEngine-Observability-Id`. HTTP middleware validates an incoming value +and stores it as a candidate. Middleware must not bind it or generate a new +value for every request. + +A calculation boundary calls `start_observability_id`. This uses an already +bound value, then a valid incoming candidate, and otherwise creates a UUID. It +binds the selected value to the request and observability runtime. Household +calculation routes call it only after request validation. On a cache miss, the +route constructs and validates the PolicyEngine situation, binds the identifier +before the successful normalization span ends, and reuses that prepared +simulation for the calculation. Pre-acceptance validation must not emit a run +stage span that cannot carry the selected identifier. A successful cache hit +binds the identifier after the cache returns. A request that fails situation +parsing must never call `start_observability_id`. + +For a new economy report, acquire submission ownership before calling +`start_observability_id`. Persist the selected identifier in the report state +written after submission. A request that reads existing report state calls +`restore_observability_id`; a persisted value is authoritative even when the +polling request supplies a different header. Older records with a null value +remain null. Never create a replacement identifier while polling. + +The HTTP response contains the header only when the request started a +calculation, continued one synchronously, or restored an existing report +identifier. + +## Simulation client transport + +The simulation HTTP client reads identifiers from request context. Its request +hook sends `request_id` and any bound `observability_id` in their respective +headers. API v1 never replaces its selected identifier with a value from an +HTTP response. + +Do not add `observability_id` to a simulation JSON body, `_telemetry`, or +execution result data class. The simulation API carries it through HTTP headers, +persists it beside functional job state, and transports captured trace context +to Modal as a separate function argument. + +## Runtime stages + +All API v1 span names for supported calculation configurations are defined in +`policyengine_api.observability.stages`. Runtime code imports the applicable +`StagePlan` and calls `plan.name(Stage.VALUE)`. Do not add span name string +literals in route or service code. + +When adding a calculation configuration or stage: + +1. Add it to `RunConfiguration` or `Stage`. +2. Add it to the applicable plan in `RUN_STAGE_REGISTRY`. +3. Import that plan in runtime code. +4. Add focused tests for the stage and identifier lifecycle. + +## Trace boundaries + +HTTP instrumentation carries W3C trace context across synchronous calls. The +simulation API carries that trace context through asynchronous Modal dispatch. +The submission and worker work can therefore form one distributed trace. +Native ASGI instrumentation must start request spans with the matched route +template. Never use the unresolved request path as a span name. + +A later polling request starts a new trace. Its logs and spans share the +persisted `observability_id` with the submission trace. Measure a complete +report by querying all diagnostic records with that identifier, then use the +registered stage names to break down elapsed time. + +When an economy report selects or restores its identifier inside a nested +stage, reapply the identifier after that stage exits so each containing economy +span receives it. HTTP completion handling must reapply a bound identifier +before ending the server request span. Keep these calls behind local exception +boundaries so a runtime failure cannot change the response. + +Stage 12 authoritative and comparison executions use the same +`observability_id` as the report that dispatched them. Their `evaluation_id` +and simulation execution identifiers remain separate functional identifiers. + +## Failure behavior + +Invalid observability configuration must fail during build or deployment +validation. After a service starts serving application traffic, logging, +tracing, metrics, context binding, and export failures must not alter an +application result or HTTP status. + +Normalize all identifiers received from callers or downstream services. Keep a +local exception boundary around runtime context binding because observability +package failures must not escape into request processing. + +Focused tests must cover: + +- identifier creation at calculation submission boundaries; +- HTTP propagation to the simulation API; +- persistence and restoration during polling; +- persisted values taking precedence over polling headers; +- older records with null identifiers; +- observability failures leaving application results unchanged. diff --git a/docs/operations/api-v1-observability.md b/docs/operations/api-v1-observability.md new file mode 100644 index 000000000..1b427f8d5 --- /dev/null +++ b/docs/operations/api-v1-observability.md @@ -0,0 +1,196 @@ +# API v1 observability operating policy + +## Scope + +Only these workloads participate: + +- The `policyengine-api` and `policyengine-api-staging` Cloud Run services. +- The `policyengine-simulation-entry` and + `policyengine-simulation-entry-staging` Cloud Run services. +- The `policyengine-simulation-gateway` Modal application. +- Versioned Modal applications whose names match + `policyengine-simulation-py--` or + `policyengine-simulation-v2-py--`. + +Cloud Run candidate, canary, and tagged revisions use the identity of their +containing service and are included. Modal smoke, precompute, and ephemeral +applications are excluded. `policyengine-household-api` and +`policyengine-uk-chat` remain unchanged and receive no migration or +service-specific validation in this work. + +The live log sink filters, collector invocation permissions, and Modal +Workload Identity Federation condition enforce this scope. Telemetry attributes +such as `service.namespace` do not grant access. + +## Calculation correlation + +`request_id` identifies one HTTP request. `observability_id` identifies one +complete household calculation or society report across API v1, the simulation +entry service, the Modal gateway, and workers. Functional identifiers such as +`job_id`, `batch_job_id`, `evaluation_id`, and `simulation_execution_id` +continue to identify stored application state. + +API v1 sends `observability_id` only in the +`X-PolicyEngine-Observability-Id` header. Household calculation routes bind it +after request validation. Economy report routes bind it after acquiring +submission ownership and persist it with report state. Polling restores the +persisted value and does not create a new one for an older record whose value +is null. + +The service binds or reapplies the selected value before each accepted +calculation stage, containing economy stage, and HTTP server span ends. Input +validation that rejects a request occurs before this boundary and does not +create an identifier. + +The simulation API preserves the same header through synchronous HTTP calls, +then passes captured observability context to Modal functions in a separate +keyword argument. Calculation payloads and API v1 execution data classes do +not contain the identifier. + +Submission work may appear in one distributed trace. Later polling requests +create additional traces with the same `observability_id`. Query all logs and +spans carrying that value to measure the complete report, and use the stage +names registered in `policyengine_api.observability.stages` and the simulation +API stage registry to attribute runtime to individual operations. + +## Storage decision + +Production and nonproduction application logs use the existing +`policyengine-observability` bucket in the `policyengine-observability` Google +Cloud project. The bucket is global, analytics enabled, and retains records for +30 days. `deployment.environment.name` distinguishes production and staging. + +A single bucket is appropriate for the initial rollout because the same small +operator group requires both environments, the current bucket already exists, +and a shared analytics surface simplifies request investigations. Access is +controlled at the bucket and project level rather than by environment. This +decision must be revisited before an environment requires different readers, +retention, residency, or deletion policy. + +Cloud Run writes structured JSON to standard output. Exact-service sinks in +the source projects route every application, request, platform, and internal +diagnostic record from the listed services into this bucket. This includes +records that do not use the application schema, while the exact Cloud Run +service names keep unrelated workloads out. Modal writes application records +to standard output and uses the package's bounded asynchronous Cloud Logging +destination under the `policyengine-api-v1-modal` log ID. A central exclusion +prevents a directly ingested record from also being retained in `_Default`. + +Traces and metrics use Cloud Trace and Cloud Monitoring in the same project. +They are correlated with logs by resource identity, trace ID, request ID, and +job ID; they are not stored in the log bucket. + +## Initial trace sampling + +The initial rollout uses these head-sampling settings: + +| Environment | Services | Parent-based ratio | +| --- | --- | ---: | +| Production | API v1, simulation entry, gateway, executors | 1.0 | +| Staging | API v1, simulation entry, gateway, executors | 1.0 | +| Local development | All | No remote exporter | + +Sampling every initial production trace provides a complete baseline for +volume, cost, errors, and slow operations. After at least one representative +week, operators may lower the normal-request ratio only after recording a cost +and coverage review. A lower head-sampling ratio cannot retroactively retain a +request after its outcome or duration becomes known. Retaining all errors or +slow requests with a lower normal ratio therefore requires a reviewed, +bounded collector tail-sampling policy. + +The checked-in collector configuration implements error and 30-second latency +policies and initially retains 100% of all remaining traces. Any later +reduction applies to the collector's general tail policy while SDK head +sampling remains at 100%, allowing the collector to evaluate completed traces. + +Parent sampling decisions are preserved. Society-wide simulation dispatches +must not override a sampled parent with an unsampled child. + +## Asynchronous trace relationships + +Captured dispatch context contains its UTC capture time. A worker uses the +dispatch span as its parent only when all of these conditions hold: + +- The work is a direct continuation of one dispatch. +- The worker starts no more than five minutes after capture. +- The invocation is not an independent retry. +- The invocation does not aggregate multiple dispatches. + +Otherwise the worker starts a new trace and links the dispatch span. Request +and job identifiers remain the same in either representation. Malformed or +expired remote context is ignored without rejecting the job. + +## Data policy + +### Required log fields + +- `schema_version` +- `timestamp` +- `severity` +- `message` or `event.name` +- `service.name`, `service.namespace`, `service.version`, `service.role` +- `deployment.environment.name` and `cloud.platform` +- Request, operation, trace, span, duration, outcome, and bounded error fields + when applicable + +Application attributes are stored below `attributes`. Local logs and spans +accept explicitly supplied strings, integers, finite floating-point values, +Booleans, and enum values after the package rejects prohibited names and +redacts configured sensitive values. Attribute strings are truncated at 1,024 +characters. There is no numeric attribute-count limit. The runtime does not +automatically capture function arguments, request bodies, or response bodies. + +Only `observability_id` is transported across an asynchronous process boundary. +Metric labels use the separate bounded list below. + +### Trace attributes + +Traces may contain the standard service resource fields, HTTP route templates, +HTTP methods, status codes, operation names, bounded deployment identifiers, +request IDs, job IDs, and explicitly supplied operational attributes that pass +the package's name, type, truncation, and redaction checks. Raw URLs, query +values, request bodies, response bodies, and arbitrary context are not +recorded. + +### Metric labels + +Metrics use only these bounded labels: + +- `service.name` +- `service.role` +- `deployment.environment.name` +- `cloud.platform` +- `http.route` +- `http.request.method` +- `http.response.status_code_class` +- `operation.name` +- `operation.kind` +- `outcome` + +Request IDs, trace IDs, job IDs, simulation IDs, raw paths, error messages, +user-provided values, unrestricted geography values, and unrestricted version +values are prohibited metric labels. + +### Prohibited telemetry data + +Logs, traces, metrics, and dispatch context must not contain: + +- Authorization headers, cookies, credentials, tokens, or secret values +- Request or response bodies +- Household situations, entity records, or person-level values +- Reform definitions or parameter payloads +- Raw IP addresses +- Prompts, model inputs, or model responses +- Exception local variables +- Function arguments or return values captured automatically + +Exception messages and stacks are truncated and passed through configured +secret-value redaction before remote delivery. + +## Operational limits + +Remote application delivery is best effort. Every application queue, exporter, +retry, network request, flush, and shutdown action has a finite bound. Queue +overflow drops the new record and increments a local counter. Internal +diagnostics are rate limited and written directly to standard error so they do +not recurse through a failing exporter. diff --git a/gcp/observability/README.md b/gcp/observability/README.md new file mode 100644 index 000000000..e40a3acda --- /dev/null +++ b/gcp/observability/README.md @@ -0,0 +1,106 @@ +# Google Cloud observability runtime + +This directory contains the source used to build the API v1 OpenTelemetry +Collector: + +| File | Purpose | +| --- | --- | +| `collector/config.yaml` | OTLP gRPC receiver, bounded processing, trace sampling, and Google Telemetry API export configuration | +| `collector/Dockerfile` | Collector container image built with that configuration | + +The Google Cloud IAM, logging sinks, collector service, dashboard, and alert +policies were provisioned separately. This repository does not manage or apply +those resources. + +## Collector behavior + +The collector accepts OTLP traces and metrics over gRPC. It has no application +log pipeline. Cloud Run services write structured JSON to standard output, and +authorized Modal applications use the observability package's bounded Cloud +Logging destination. + +The collector applies a memory limit, batches exports, and sends signals to +`telemetry.googleapis.com` using its Google service account. Its trace policy +retains errors, operations lasting at least 30 seconds, and currently 100% of +all remaining traces. Participating SDKs therefore use 100% head sampling so +the collector can evaluate complete traces. + +Changing `collector/config.yaml` does not update the live service. The image +must be rebuilt and the existing `policyengine-api-v1-otel-collector` Cloud Run +service must be updated through a separately managed deployment process. No +collector deployment workflow exists in this repository. + +## Participating workloads + +The live GCP permissions and routing configuration cover only: + +- `policyengine-api` and `policyengine-api-staging` in the API project. +- `policyengine-simulation-entry` and + `policyengine-simulation-entry-staging` in the simulation entry project. +- The `policyengine-simulation-gateway` Modal application. +- Versioned Modal applications matching + `policyengine-simulation-py--` or + `policyengine-simulation-v2-py--`. + +Modal smoke, precompute, ephemeral, Household API, and UK Chat applications are +excluded. + +## Consumer configuration + +The API and simulation repositories configure these GitHub Actions variables: + +- `OBSERVABILITY_SERVICE_NAMESPACE` +- `OBSERVABILITY_TRACE_PROJECT_ID` +- `OBSERVABILITY_OTLP_ENDPOINT` +- `OBSERVABILITY_OTLP_GOOGLE_AUDIENCE` + +The simulation repository additionally configures: + +- `OBSERVABILITY_LOGGING_PROJECT_ID` +- `OBSERVABILITY_LOG_NAME` +- `OBSERVABILITY_GOOGLE_WORKLOAD_IDENTITY_PROVIDER` +- `OBSERVABILITY_GOOGLE_SERVICE_ACCOUNT_EMAIL` + +API Cloud Run services use source-project logging sinks and therefore require +no direct Cloud Logging credentials. + +## Live infrastructure record + +The infrastructure was applied and verified on 2026-09-22: + +- The global central log bucket in `policyengine-observability` has log + analytics enabled and 30-day retention. +- Exact Cloud Run service filters route the two API services and two simulation + entry services to the central bucket. +- Modal application logs use the `policyengine-api-v1-modal` log ID, and an + exclusion prevents duplicate retention in `_Default`. +- The authenticated `policyengine-api-v1-otel-collector` service runs in + `us-central1`. +- The collector uses the `policyengine-otel-collector` service account. +- A dedicated `modal-api-v1` Workload Identity Federation provider restricts + access by workspace, environment, and application name. +- The only project-level `roles/logging.logWriter` identity is the API v1 Modal + service account. +- The Cloud Monitoring dashboard and six alert policies are enabled. +- The alert policies have no notification channels, so they record incidents + without sending email, Slack, or paging notifications. + +The consumer services require `policyengine-observability` 3.0.1. Record the +deployed consumer revisions and a representative cost and volume observation +interval after the API v1 rollout. + +## Rollback + +1. Remove the OTel endpoint from participating service configuration. +2. Remove Modal remote logging configuration. +3. Revert participating services to their previous package versions and + deployment revisions. +4. Remove the API v1 source sinks and restore the previous central direct-log + sink and `_Default` exclusion configuration. +5. Remove collector invocation permissions, disable the `modal-api-v1` + provider, and disable or delete the collector service. +6. Retain the central bucket for its configured retention period unless stored + data caused the incident. + +Rollback does not modify the existing `modal/modal` provider or excluded +applications. diff --git a/gcp/observability/collector/Dockerfile b/gcp/observability/collector/Dockerfile new file mode 100644 index 000000000..ae3f63207 --- /dev/null +++ b/gcp/observability/collector/Dockerfile @@ -0,0 +1,6 @@ +FROM us-docker.pkg.dev/cloud-ops-agents-artifacts/google-cloud-opentelemetry-collector/otelcol-google:0.160.0 + +COPY config.yaml /etc/otelcol-google/config.yaml + +ENTRYPOINT ["/otelcol-google"] +CMD ["--config=/etc/otelcol-google/config.yaml"] diff --git a/gcp/observability/collector/config.yaml b/gcp/observability/collector/config.yaml new file mode 100644 index 000000000..2fa3ced55 --- /dev/null +++ b/gcp/observability/collector/config.yaml @@ -0,0 +1,71 @@ +receivers: + otlp: + protocols: + grpc: + endpoint: 0.0.0.0:8080 + +processors: + memory_limiter: + check_interval: 1s + limit_mib: 384 + spike_limit_mib: 64 + resource/destination: + attributes: + - key: gcp.project_id + value: ${env:OBSERVABILITY_PROJECT_ID} + action: upsert + tail_sampling: + decision_wait: 30s + num_traces: 50000 + expected_new_traces_per_sec: 500 + policies: + - name: retain-errors + type: status_code + status_code: + status_codes: [ERROR] + - name: retain-slow-operations + type: latency + latency: + threshold_ms: 30000 + - name: initial-full-sample + type: probabilistic + probabilistic: + sampling_percentage: 100 + batch: + send_batch_size: 200 + send_batch_max_size: 1000 + timeout: 5s + +exporters: + otlp_grpc: + endpoint: telemetry.googleapis.com:443 + auth: + authenticator: googleclientauth + sending_queue: + enabled: true + queue_size: 2000 + retry_on_failure: + enabled: true + initial_interval: 1s + max_interval: 5s + max_elapsed_time: 30s + +extensions: + googleclientauth: + health_check: + endpoint: 0.0.0.0:13133 + +service: + extensions: [googleclientauth, health_check] + pipelines: + traces: + receivers: [otlp] + processors: [memory_limiter, resource/destination, tail_sampling, batch] + exporters: [otlp_grpc] + metrics: + receivers: [otlp] + processors: [memory_limiter, resource/destination, batch] + exporters: [otlp_grpc] + telemetry: + logs: + level: info diff --git a/policyengine_api/api.py b/policyengine_api/api.py index 8c470596b..2c0d23a35 100644 --- a/policyengine_api/api.py +++ b/policyengine_api/api.py @@ -23,7 +23,11 @@ def log_timing(message): from policyengine_api.extensions import cache from policyengine_api.migration_logging import register_migration_request_logging +from policyengine_api.observability.identifiers import OBSERVABILITY_ID_HEADER +from policyengine_api.request_context import REQUEST_ID_HEADER +from policyengine_api.observability import runtime as observability_runtime from policyengine_api.runtime_cache.settings import load_runtime_cache_settings +from policyengine_observability import instrument_flask log_timing("Caching utilities import completed") @@ -60,6 +64,8 @@ def log_timing(message): app = application = flask.Flask(__name__) log_timing("Flask app created") +instrument_flask(app, observability_runtime) +log_timing("Observability initialised") runtime_cache_settings = load_runtime_cache_settings() if runtime_cache_settings.enabled: @@ -99,10 +105,10 @@ def log_timing(message): cache.init_app(app) log_timing("Caching initialised") -CORS(app) +CORS(app, expose_headers=[REQUEST_ID_HEADER, OBSERVABILITY_ID_HEADER]) log_timing("CORS initialised") -register_migration_request_logging(app) +register_migration_request_logging(app, runtime=observability_runtime) log_timing("Migration request logging initialised") app.register_blueprint(error_bp) diff --git a/policyengine_api/asgi.py b/policyengine_api/asgi.py index d49ec1405..1f981c56a 100644 --- a/policyengine_api/asgi.py +++ b/policyengine_api/asgi.py @@ -7,6 +7,7 @@ from policyengine_api.api import app as flask_app from policyengine_api.asgi_factory import create_asgi_app from policyengine_api.data.orm import close_v1_engines +from policyengine_api.observability import get_runtime from policyengine_api.readiness import mark_not_ready, mark_ready from policyengine_api.runtime_cache.client import close_runtime_cache_clients from policyengine_api.warmup import run_startup_warmup @@ -15,6 +16,7 @@ def _close_runtime_resources() -> None: close_v1_engines() close_runtime_cache_clients() + get_runtime().shutdown() app = application = create_asgi_app( diff --git a/policyengine_api/asgi_factory.py b/policyengine_api/asgi_factory.py index bcc9802ed..408cd8aa5 100644 --- a/policyengine_api/asgi_factory.py +++ b/policyengine_api/asgi_factory.py @@ -10,7 +10,6 @@ from fastapi import FastAPI, Request from fastapi.exception_handlers import request_validation_exception_handler from fastapi.exceptions import RequestValidationError -from fastapi.routing import APIRoute from policyengine_api.constants import VERSION from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies from policyengine_api.fastapi_routes.health import build_core_health_router @@ -29,15 +28,25 @@ RouteImplementationSettings, ) from policyengine_api.migration_logging import log_migration_request +from policyengine_api.observability import get_runtime from policyengine_api.request_context import ( REQUEST_ID_HEADER, + _asgi_incoming_observability_id, + _asgi_observability_id, _asgi_request_id, + current_observability_id, generate_request_id, ) +from policyengine_observability import ObservabilityRuntime +from policyengine_api.observability.identifiers import ( + OBSERVABILITY_ID_HEADER, + normalize_observability_id, +) from starlette.datastructures import MutableHeaders from starlette.middleware.cors import CORSMiddleware from starlette.middleware.gzip import GZipMiddleware from starlette.responses import PlainTextResponse, Response +from starlette.routing import Match, Mount from starlette.types import ASGIApp @@ -48,12 +57,48 @@ def _apply_request_id_header( response.headers[REQUEST_ID_HEADER] = request_id +def _apply_observability_id_header( + response: Response, + observability_id: str, +) -> None: + response.headers[OBSERVABILITY_ID_HEADER] = observability_id + + +def _matched_route_template(routes, scope: dict) -> str | None: + for route in routes: + match, _ = route.matches(scope) + if match is Match.FULL: + if isinstance(route, Mount): + return None + route_template = getattr(route, "path_format", None) or getattr( + route, "path", None + ) + if isinstance(route_template, str): + return route_template + included_router = getattr(route, "original_router", None) + if included_router is not None: + nested_template = _matched_route_template( + included_router.routes, + scope, + ) + if nested_template is not None: + return nested_template + return None + + +def _native_route_template(app: FastAPI, scope: dict) -> str | None: + """Return the matched FastAPI route template, excluding the Flask mount.""" + + return _matched_route_template(app.router.routes, scope) + + def create_asgi_app( wsgi_app, *, route_settings: RouteImplementationSettings | None = None, dependencies: NativeRouteDependencies | None = None, shutdown_callback: Callable[[], None] | None = None, + observability_runtime: ObservabilityRuntime | None = None, ) -> ASGIApp: """Create the Stage 2 FastAPI shell around the existing Flask app.""" @@ -61,6 +106,7 @@ def create_asgi_app( route_settings = RouteImplementationSettings.from_environment() if dependencies is None: dependencies = NativeRouteDependencies.defaults() + request_runtime = observability_runtime or get_runtime() @asynccontextmanager async def lifespan(_app: FastAPI): @@ -97,6 +143,9 @@ async def add_headers_to_unhandled_errors( request.headers.get(REQUEST_ID_HEADER) or generate_request_id(), ) _apply_request_id_header(response, request_id) + observability_id = current_observability_id() + if observability_id is not None: + _apply_observability_id_header(response, observability_id) return response @app.exception_handler(RequestValidationError) @@ -119,12 +168,48 @@ async def oversized_v2_request( async def add_request_context_and_migration_logging(request, call_next): started_at = time.time() request_id = request.headers.get(REQUEST_ID_HEADER) or generate_request_id() + incoming_observability_id = normalize_observability_id( + request.headers.get(OBSERVABILITY_ID_HEADER) + ) MutableHeaders(scope=request.scope)[REQUEST_ID_HEADER] = request_id request.state.policyengine_request_id = request_id + request.state.policyengine_observability_id = None context_token = _asgi_request_id.set(request_id) + observability_context_token = _asgi_observability_id.set(None) + incoming_observability_context_token = _asgi_incoming_observability_id.set( + incoming_observability_id + ) + native_route_template = _native_route_template(app, request.scope) + native_request = native_route_template is not None + initial_route = native_route_template or request.url.path + + if native_request: + try: + runtime_request_id = request_runtime.begin_request( + headers={ + **dict(request.headers), + REQUEST_ID_HEADER: request_id, + }, + method=request.method, + route=initial_route, + ) + if isinstance(runtime_request_id, str) and runtime_request_id: + request_id = runtime_request_id + MutableHeaders(scope=request.scope)[REQUEST_ID_HEADER] = request_id + request.state.policyengine_request_id = request_id + _asgi_request_id.reset(context_token) + context_token = _asgi_request_id.set(request_id) + except Exception: + pass + try: + request_runtime.set_context( + request_id=request_id, + ) + except Exception: + pass def log_native_route(status_code: int) -> None: - if not isinstance(request.scope.get("route"), APIRoute): + if not native_request: return try: log_migration_request( @@ -142,17 +227,70 @@ def log_native_route(status_code: int) -> None: except Exception: pass + def finish_native_route( + status_code: int, + error: BaseException | None = None, + ) -> None: + if not native_request: + return + resolved_route = getattr(request.scope.get("route"), "path", initial_route) + try: + request_runtime.update_request_route(resolved_route) + except Exception: + pass + try: + request_runtime.update_request_status(status_code) + except Exception: + pass + try: + request_runtime.end_request( + status_code=status_code, + error=error, + ) + except Exception: + pass + try: try: response = await call_next(request) - except Exception: + except Exception as error: log_native_route(500) + finish_native_route(500, error) raise + response_observability_id = ( + current_observability_id() + or normalize_observability_id( + response.headers.get(OBSERVABILITY_ID_HEADER) + ) + ) + if native_request and response_observability_id is not None: + try: + # Route-level spans have completed at this point, so this + # call applies the selected calculation identifier to the + # still-current server request span. + request_runtime.set_context( + observability_id=response_observability_id + ) + except Exception: + pass + if native_request: + try: + for name, value in request_runtime.response_headers().items(): + response.headers[name] = value + except Exception: + pass _apply_request_id_header(response, request_id) + if response_observability_id is not None: + _apply_observability_id_header(response, response_observability_id) + elif OBSERVABILITY_ID_HEADER in response.headers: + del response.headers[OBSERVABILITY_ID_HEADER] log_native_route(response.status_code) + finish_native_route(response.status_code) return response finally: _asgi_request_id.reset(context_token) + _asgi_observability_id.reset(observability_context_token) + _asgi_incoming_observability_id.reset(incoming_observability_context_token) app.include_router(build_core_health_router(dependencies)) app.include_router(build_v2_router(dependencies)) @@ -171,7 +309,7 @@ def log_native_route(status_code: int) -> None: allow_origin_regex=".*", allow_methods=["DELETE", "GET", "HEAD", "OPTIONS", "PATCH", "POST", "PUT"], allow_headers=["*"], - expose_headers=[REQUEST_ID_HEADER], + expose_headers=[REQUEST_ID_HEADER, OBSERVABILITY_ID_HEADER], allow_credentials=False, max_age=600, ) diff --git a/policyengine_api/country.py b/policyengine_api/country.py index 1caa0f50f..7af395ffc 100644 --- a/policyengine_api/country.py +++ b/policyengine_api/country.py @@ -2,8 +2,9 @@ import inspect import json import logging +from dataclasses import dataclass from policyengine_core.taxbenefitsystems import TaxBenefitSystem -from typing import Union +from typing import Any, Union from policyengine_api.utils import get_safe_json from policyengine_core.parameters import ( ParameterNode, @@ -39,6 +40,18 @@ logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class PreparedCountryCalculation: + """A parsed PolicyEngine situation ready for requested calculations.""" + + simulation: Any + system: TaxBenefitSystem + household: dict + requested_computations: list[tuple[str, str, str, str]] + has_axes: bool + spm_requested: bool + + def _serialize_float(value): serialized = float(str(value)) if serialized == float("inf"): @@ -431,6 +444,7 @@ def calculate( reform: Union[dict, None], spm: dict | None = None, spm_requested: bool = False, + prepared: PreparedCountryCalculation | None = None, ) -> CalculationResult: """Calculate requested variables, optionally under a chosen measurement. @@ -439,16 +453,18 @@ def calculate( inherited bundle default was not chosen, and a variable that depends on it stays unavailable the way every other uncomputable variable does. """ - simulation, system = self._create_simulation(household, reform, spm=spm) - - household = json.loads(json.dumps(household)) - - has_axes = "axes" in household - requested_computations = get_requested_computations( + prepared = prepared or self.prepare_calculation( household, - include_provided_values=has_axes, - variable_names=set(system.variables) if has_axes else None, + reform, + spm=spm, + spm_requested=spm_requested, ) + simulation = prepared.simulation + system = prepared.system + household = prepared.household + has_axes = prepared.has_axes + requested_computations = prepared.requested_computations + spm_requested = prepared.spm_requested calculation_warnings: list[str] = [] for ( @@ -520,6 +536,32 @@ def calculate( **calculation_spm_receipt(simulation), ) + def prepare_calculation( + self, + household: dict, + reform: Union[dict, None], + spm: dict | None = None, + spm_requested: bool = False, + ) -> PreparedCountryCalculation: + """Parse a situation before the API accepts it as a calculation.""" + + simulation, system = self._create_simulation(household, reform, spm=spm) + household = json.loads(json.dumps(household)) + has_axes = "axes" in household + requested_computations = get_requested_computations( + household, + include_provided_values=has_axes, + variable_names=set(system.variables) if has_axes else None, + ) + return PreparedCountryCalculation( + simulation=simulation, + system=system, + household=household, + requested_computations=requested_computations, + has_axes=has_axes, + spm_requested=spm_requested, + ) + def _create_simulation( self, household: dict, diff --git a/policyengine_api/gcp_logging.py b/policyengine_api/gcp_logging.py index be3c96e1b..6c3b55b7e 100644 --- a/policyengine_api/gcp_logging.py +++ b/policyengine_api/gcp_logging.py @@ -1,62 +1,50 @@ -import logging -import os -from typing import Optional - - -class _LazyGoogleLogger: - """Lazily initialize Google Cloud Logging and fall back to stderr.""" - - def __init__(self, logger_name: str): - self._logger_name = logger_name - self._google_logger = None - self._initialization_failed = False - self._fallback_logger = logging.getLogger(logger_name) - - def _get_google_logger(self): - if not os.environ.get("K_SERVICE"): - self._initialization_failed = True - return None - if self._google_logger is not None: - return self._google_logger - if self._initialization_failed: - return None - try: - from google.cloud.logging import Client +"""Compatibility facade for application-owned structured logging. - self._google_logger = Client().logger(self._logger_name) - return self._google_logger - except Exception: - self._initialization_failed = True - return None +Existing API modules call ``logger.log_struct``. The facade keeps that small +surface while sending records through the explicitly owned v2 runtime. Cloud +Run captures the resulting JSON from standard output, so request threads never +call the Cloud Logging API. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any + +from policyengine_api.observability import get_runtime + + +class _RuntimeLogger: + """Adapt the former ``log_struct`` call shape to the v2 runtime.""" def log_struct( self, - info: dict, + info: Mapping[str, Any], severity: str = "INFO", *, - labels: Optional[dict] = None, + labels: Mapping[str, Any] | None = None, ) -> None: - """Record structured diagnostics without changing caller behavior.""" - - google_logger = self._get_google_logger() - if google_logger is not None: - try: - google_logger.log_struct(info, severity=severity, labels=labels) - return - except Exception: - # Observability must never invalidate a successful request or - # cache operation. Cloud Run collects stderr as a fallback - # when the structured logging API is unavailable. - self._google_logger = None - self._initialization_failed = True - - level = getattr(logging, severity.upper(), logging.INFO) + """Record an allowlisted structured message without affecting callers.""" + try: - self._fallback_logger.log(level, "%s", info) + message = str(info.get("message") or "API event") + attributes = { + key: value + for key, value in info.items() + if key not in {"message", "migration", "response_text"} + } + migration = info.get("migration") + if isinstance(migration, Mapping): + attributes.update(migration) + if labels: + attributes.update(labels) + get_runtime().log( + message, + severity=severity, + attributes=attributes, + ) except Exception: - # Logging is diagnostic only. A broken local handler must not - # change the result of the operation that attempted to log. pass -logger = _LazyGoogleLogger("policyengine-api") +logger = _RuntimeLogger() diff --git a/policyengine_api/libs/simulation_entrypoint.py b/policyengine_api/libs/simulation_entrypoint.py index 7d4b2f19d..aa2a2bd17 100644 --- a/policyengine_api/libs/simulation_entrypoint.py +++ b/policyengine_api/libs/simulation_entrypoint.py @@ -7,6 +7,7 @@ from urllib.parse import urlparse import httpx +from policyengine_observability import instrument_httpx from policyengine_api.gcp_logging import logger from policyengine_api.libs.gateway_auth import ( GatewayAuthError, @@ -16,10 +17,13 @@ gateway_auth_required, ) from policyengine_api.migration_flags import get_sim_entrypoint +from policyengine_api.observability import get_runtime from policyengine_api.request_context import ( REQUEST_ID_HEADER, + current_observability_id, current_request_id, ) +from policyengine_api.observability.identifiers import OBSERVABILITY_ID_HEADER from policyengine_api.worker_spm import validate_worker_spm, raise_worker_spm_error @@ -60,6 +64,9 @@ def _attach_current_request_id(request: httpx.Request) -> None: request_id = current_request_id() if request_id is not None: request.headers[REQUEST_ID_HEADER] = request_id + observability_id = current_observability_id() + if observability_id is not None: + request.headers[OBSERVABILITY_ID_HEADER] = observability_id @dataclass @@ -70,7 +77,6 @@ class ModalSimulationExecution: job_id: str status: str - run_id: Optional[str] = None result: Optional[dict] = None error: Optional[str] = None policyengine_bundle: Optional[dict] = None @@ -142,6 +148,7 @@ def __init__(self, entrypoint: str | None = None): auth=auth, event_hooks={"request": [_attach_current_request_id]}, ) + instrument_httpx(self.client, get_runtime()) def _normalize_submission_payload(self, payload: dict) -> dict: if "data" in payload or "data_version" in payload: @@ -218,7 +225,7 @@ def run(self, payload: dict) -> ModalSimulationExecution: { "message": "Simulation entrypoint job submitted", "job_id": data.get("job_id"), - "run_id": data.get("run_id"), + "observability_id": current_observability_id(), "status": data.get("status"), }, severity="INFO", @@ -229,14 +236,13 @@ def run(self, payload: dict) -> ModalSimulationExecution: status=data["status"], policyengine_bundle=data.get("policyengine_bundle"), resolved_app_name=data.get("resolved_app_name"), - run_id=data.get("run_id"), ) except httpx.HTTPStatusError as e: logger.log_struct( { "message": f"Simulation entrypoint HTTP error: {e.response.status_code}", - "run_id": (payload.get("_telemetry") or {}).get("run_id"), + "observability_id": current_observability_id(), "response_text": e.response.text[:500], }, severity="ERROR", @@ -247,7 +253,7 @@ def run(self, payload: dict) -> ModalSimulationExecution: logger.log_struct( { "message": f"Simulation entrypoint request error: {str(e)}", - "run_id": (payload.get("_telemetry") or {}).get("run_id"), + "observability_id": current_observability_id(), }, severity="ERROR", ) @@ -296,7 +302,7 @@ def run_budget_window_batch(self, payload: dict) -> ModalBudgetWindowBatchExecut logger.log_struct( { "message": f"Simulation batch API request error: {str(e)}", - "run_id": (payload.get("_telemetry") or {}).get("run_id"), + "observability_id": current_observability_id(), }, severity="ERROR", ) @@ -431,7 +437,6 @@ def get_execution_by_id(self, job_id: str) -> ModalSimulationExecution: return ModalSimulationExecution( job_id=job_id, status=data["status"], - run_id=data.get("run_id"), result=data.get("result"), error=data.get("error"), policyengine_bundle=data.get("policyengine_bundle"), diff --git a/policyengine_api/migration_logging.py b/policyengine_api/migration_logging.py index 75b1ed447..e4432ace1 100644 --- a/policyengine_api/migration_logging.py +++ b/policyengine_api/migration_logging.py @@ -5,6 +5,7 @@ import time import flask +from policyengine_observability import ObservabilityRuntime from policyengine_api.gcp_logging import logger from policyengine_api.migration_flags import ( RouteImplementation, @@ -15,6 +16,10 @@ REQUEST_ID_HEADER, generate_request_id, ) +from policyengine_api.observability.identifiers import ( + OBSERVABILITY_ID_HEADER, + normalize_observability_id, +) V2_METADATA_RESOURCE_SEGMENTS = frozenset( @@ -66,33 +71,69 @@ def _is_v2_household_resource(method: str, path: str) -> bool: ) -def register_migration_request_logging(app: flask.Flask) -> None: +def register_migration_request_logging( + app: flask.Flask, + *, + runtime: ObservabilityRuntime | None = None, +) -> None: """Register request IDs and migration logging for Flask.""" @app.before_request def set_request_migration_context(): flask.g.request_started_at = time.time() - flask.g.request_id = ( + try: + captured = runtime.capture_context() if runtime is not None else {} + except Exception: + captured = {} + flask.g.request_id = captured.get("request_id") or ( flask.request.headers.get(REQUEST_ID_HEADER) or generate_request_id() ) + flask.g.incoming_observability_id = normalize_observability_id( + flask.request.headers.get(OBSERVABILITY_ID_HEADER) + ) + flask.g.observability_id = None + if runtime is not None: + try: + runtime.set_context(request_id=flask.g.request_id) + except Exception: + pass @app.after_request def log_request_migration_context(response): request_id = getattr(flask.g, "request_id", None) if request_id is not None: response.headers[REQUEST_ID_HEADER] = request_id + observability_id = getattr(flask.g, "observability_id", None) + if observability_id is not None: + response.headers[OBSERVABILITY_ID_HEADER] = observability_id try: - log_migration_request( - request_id=request_id, - method=flask.request.method, - path=flask.request.path, - status_code=response.status_code, - started_at=getattr(flask.g, "request_started_at", None), - country_id=flask.request.view_args.get("country_id") + country_id = ( + flask.request.view_args.get("country_id") if flask.request.view_args - else None, - route_impl=RouteImplementation.FLASK_FALLBACK, + else None ) + if runtime is not None: + runtime_context = { + "country_id": country_id, + **_migration_context( + method=flask.request.method, + path=flask.request.path, + route_impl=RouteImplementation.FLASK_FALLBACK, + ), + } + if observability_id is not None: + runtime_context["observability_id"] = observability_id + runtime.set_context(**runtime_context) + else: + log_migration_request( + request_id=request_id, + method=flask.request.method, + path=flask.request.path, + status_code=response.status_code, + started_at=getattr(flask.g, "request_started_at", None), + country_id=country_id, + route_impl=RouteImplementation.FLASK_FALLBACK, + ) except Exception: try: app.logger.exception("Failed to log migration request context") @@ -117,6 +158,33 @@ def log_migration_request( if started_at is not None: elapsed_ms = round((time.time() - started_at) * 1000, 2) + migration_context = _migration_context( + method=method, + path=path, + route_impl=route_impl, + ) + + logger.log_struct( + { + "message": "API request served", + "request_id": request_id, + "method": method, + "path": path, + "status_code": status_code, + "latency_ms": elapsed_ms, + "country_id": country_id, + "migration": migration_context, + }, + severity="INFO" if status_code < 500 else "ERROR", + ) + + +def _migration_context( + *, + method: str, + path: str, + route_impl: RouteImplementation | None, +) -> dict[str, str | None]: route_group = infer_route_group(path) is_v2_metadata_read = _is_v2_metadata_resource_read(method, path) is_v2_policy_resource = _is_v2_policy_resource(method, path) @@ -124,7 +192,7 @@ def log_migration_request( uses_explicit_v2_source = ( is_v2_metadata_read or is_v2_policy_resource or is_v2_household_resource ) - migration_context = get_migration_log_context( + return get_migration_log_context( route_group, route_impl=route_impl, use_configured_db_sources=not uses_explicit_v2_source, @@ -142,17 +210,3 @@ def log_migration_request( else None ), ) - - logger.log_struct( - { - "message": "API request served", - "request_id": request_id, - "method": method, - "path": path, - "status_code": status_code, - "latency_ms": elapsed_ms, - "country_id": country_id, - "migration": migration_context, - }, - severity="INFO" if status_code < 500 else "ERROR", - ) diff --git a/policyengine_api/observability/__init__.py b/policyengine_api/observability/__init__.py new file mode 100644 index 000000000..c7c1121f3 --- /dev/null +++ b/policyengine_api/observability/__init__.py @@ -0,0 +1,5 @@ +"""API observability runtime, identifiers, and registered stage plans.""" + +from .runtime import _build_runtime, get_runtime, runtime, set_runtime_context + +__all__ = ["_build_runtime", "get_runtime", "runtime", "set_runtime_context"] diff --git a/policyengine_api/observability/identifiers.py b/policyengine_api/observability/identifiers.py new file mode 100644 index 000000000..e836fc887 --- /dev/null +++ b/policyengine_api/observability/identifiers.py @@ -0,0 +1,25 @@ +"""Diagnostic correlation identifiers for API requests and report work.""" + +from __future__ import annotations + +from typing import Any +from uuid import UUID, uuid4 + +OBSERVABILITY_ID_HEADER = "X-PolicyEngine-Observability-Id" + + +def generate_observability_id() -> str: + """Create an identifier used only to correlate observability records.""" + + return str(uuid4()) + + +def normalize_observability_id(value: Any) -> str | None: + """Return a canonical UUID string, or ``None`` for malformed input.""" + + if not isinstance(value, str): + return None + try: + return str(UUID(value)) + except (ValueError, AttributeError): + return None diff --git a/policyengine_api/observability/runtime.py b/policyengine_api/observability/runtime.py new file mode 100644 index 000000000..857cada38 --- /dev/null +++ b/policyengine_api/observability/runtime.py @@ -0,0 +1,77 @@ +"""Explicit API v1 observability runtime ownership.""" + +from __future__ import annotations + +import os +from dataclasses import replace +from importlib.metadata import PackageNotFoundError, version + +from policyengine_observability import ( + DeploymentIdentity, + GoogleCloudLogFormatter, + LoggingConfig, + ObservabilityConfig, + ObservabilityRuntime, + ServiceIdentity, + StdoutLogDestination, + configure, +) + + +DISPATCH_ATTRIBUTE_KEYS = frozenset({"observability_id"}) + + +def _package_version() -> str: + try: + return version("policyengine-api") + except PackageNotFoundError: + return "4.1.0" + + +def _build_runtime() -> ObservabilityRuntime: + environment = os.getenv("APP_ENVIRONMENT", "local").strip() or "local" + trace_project = os.getenv("OBSERVABILITY_TRACE_PROJECT_ID", "").strip() + formatter = GoogleCloudLogFormatter(trace_project) if trace_project else None + config = ObservabilityConfig.from_env( + service=ServiceIdentity( + name="policyengine-api", + namespace=os.getenv( + "OBSERVABILITY_SERVICE_NAMESPACE", + "policyengine.api-v1", + ), + version=_package_version(), + role="api", + ), + deployment=DeploymentIdentity( + environment=environment, + platform="google_cloud_run", + region=os.getenv("CLOUD_RUN_REGION") or "us-central1", + instance_id=os.getenv("K_REVISION"), + ), + logging=LoggingConfig( + destinations=(StdoutLogDestination(formatter=formatter),), + capture_standard_library=True, + ), + dispatch_attribute_keys=DISPATCH_ATTRIBUTE_KEYS, + ) + config = replace( + config, + otel=replace(config.otel, sampling_ratio=1.0), + ) + return configure(config) + + +runtime = _build_runtime() + + +def get_runtime() -> ObservabilityRuntime: + return runtime + + +def set_runtime_context(**attributes: object) -> None: + """Bind local telemetry attributes without affecting application behavior.""" + + try: + runtime.set_context(**attributes) + except Exception: + pass diff --git a/policyengine_api/observability/stages.py b/policyengine_api/observability/stages.py new file mode 100644 index 000000000..eaba1b7d9 --- /dev/null +++ b/policyengine_api/observability/stages.py @@ -0,0 +1,103 @@ +"""Canonical stage registry for API calculation configurations.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import StrEnum +from types import MappingProxyType +from typing import Mapping + + +class RunConfiguration(StrEnum): + HOUSEHOLD = "household" + ECONOMY_ANNUAL = "economy_annual" + ECONOMY_BUDGET_WINDOW = "economy_budget_window" + + +class Stage(StrEnum): + HOUSEHOLD_LOAD_INPUTS = "household.load_inputs" + HOUSEHOLD_CACHE_LOOKUP = "household.cache_lookup" + HOUSEHOLD_INPUT_NORMALIZATION = "household.input_normalization" + HOUSEHOLD_CALCULATION = "household.calculation" + HOUSEHOLD_CACHE_WRITE = "household.cache_write" + + ECONOMY_REQUEST = "economy.request" + ECONOMY_BUDGET_WINDOW_REQUEST = "economy.budget_window_request" + ECONOMY_LOAD_POLICIES = "economy.load_policies" + ECONOMY_RESOLVE_CACHED_OR_NEW = "economy.resolve_cached_or_new_impact" + ECONOMY_RESOLVE_RUNTIME_BUNDLE = "economy.resolve_runtime_bundle" + ECONOMY_SUBMIT = "economy.submit_impact" + ECONOMY_START_BUDGET_WINDOW = "economy.start_budget_window_batch" + ECONOMY_POLL_BUDGET_WINDOW = "economy.poll_budget_window_batch" + ECONOMY_HANDLE_EXECUTION_STATE = "economy.handle_execution_state" + ECONOMY_READ_COMPLETED = "economy.read_completed_impact" + ECONOMY_READ_FAILED = "economy.read_failed_impact" + ECONOMY_POLL_ACTIVE = "economy.poll_active_impact" + ECONOMY_PERSIST_COMPUTING = "economy.persist_computing_impact" + ECONOMY_PERSIST_COMPLETED = "economy.persist_completed_impact" + ECONOMY_PERSIST_FAILED = "economy.persist_failed_impact" + + +@dataclass(frozen=True) +class StagePlan: + configuration: RunConfiguration + stages: tuple[Stage, ...] + + def name(self, stage: Stage) -> str: + if stage not in self.stages: + raise ValueError( + f"{stage.value!r} is not registered for {self.configuration.value!r}" + ) + return stage.value + + +_ECONOMY_COMMON = ( + Stage.ECONOMY_LOAD_POLICIES, + Stage.ECONOMY_RESOLVE_CACHED_OR_NEW, + Stage.ECONOMY_RESOLVE_RUNTIME_BUNDLE, + Stage.ECONOMY_HANDLE_EXECUTION_STATE, + Stage.ECONOMY_READ_COMPLETED, + Stage.ECONOMY_READ_FAILED, + Stage.ECONOMY_POLL_ACTIVE, + Stage.ECONOMY_PERSIST_COMPUTING, + Stage.ECONOMY_PERSIST_COMPLETED, + Stage.ECONOMY_PERSIST_FAILED, +) + +RUN_STAGE_REGISTRY: Mapping[RunConfiguration, StagePlan] = MappingProxyType( + { + RunConfiguration.HOUSEHOLD: StagePlan( + RunConfiguration.HOUSEHOLD, + ( + Stage.HOUSEHOLD_LOAD_INPUTS, + Stage.HOUSEHOLD_CACHE_LOOKUP, + Stage.HOUSEHOLD_INPUT_NORMALIZATION, + Stage.HOUSEHOLD_CALCULATION, + Stage.HOUSEHOLD_CACHE_WRITE, + ), + ), + RunConfiguration.ECONOMY_ANNUAL: StagePlan( + RunConfiguration.ECONOMY_ANNUAL, + ( + Stage.ECONOMY_REQUEST, + *_ECONOMY_COMMON, + Stage.ECONOMY_SUBMIT, + ), + ), + RunConfiguration.ECONOMY_BUDGET_WINDOW: StagePlan( + RunConfiguration.ECONOMY_BUDGET_WINDOW, + ( + Stage.ECONOMY_BUDGET_WINDOW_REQUEST, + *_ECONOMY_COMMON, + Stage.ECONOMY_START_BUDGET_WINDOW, + Stage.ECONOMY_POLL_BUDGET_WINDOW, + ), + ), + } +) + +HOUSEHOLD_STAGES = RUN_STAGE_REGISTRY[RunConfiguration.HOUSEHOLD] +ECONOMY_ANNUAL_STAGES = RUN_STAGE_REGISTRY[RunConfiguration.ECONOMY_ANNUAL] +ECONOMY_BUDGET_WINDOW_STAGES = RUN_STAGE_REGISTRY[ + RunConfiguration.ECONOMY_BUDGET_WINDOW +] diff --git a/policyengine_api/request_context.py b/policyengine_api/request_context.py index feff6194f..08c88d1d2 100644 --- a/policyengine_api/request_context.py +++ b/policyengine_api/request_context.py @@ -6,6 +6,10 @@ from contextvars import ContextVar import flask +from policyengine_api.observability.identifiers import ( + generate_observability_id, + normalize_observability_id, +) REQUEST_ID_HEADER = "X-PolicyEngine-Request-Id" @@ -13,6 +17,14 @@ "policyengine_api_request_id", default=None, ) +_asgi_observability_id: ContextVar[str | None] = ContextVar( + "policyengine_api_observability_id", + default=None, +) +_asgi_incoming_observability_id: ContextVar[str | None] = ContextVar( + "policyengine_api_incoming_observability_id", + default=None, +) def generate_request_id() -> str: @@ -27,3 +39,60 @@ def current_request_id() -> str | None: if flask.has_request_context(): return getattr(flask.g, "request_id", None) return _asgi_request_id.get() + + +def current_observability_id() -> str | None: + """Return the diagnostic correlation identifier for the current request.""" + + if flask.has_request_context(): + return getattr(flask.g, "observability_id", None) + return _asgi_observability_id.get() + + +def incoming_observability_id() -> str | None: + """Return the validated identifier candidate supplied by the caller.""" + + if flask.has_request_context(): + return getattr(flask.g, "incoming_observability_id", None) + return _asgi_incoming_observability_id.get() + + +def _bind_observability_id(observability_id: str) -> str: + if flask.has_request_context(): + flask.g.observability_id = observability_id + else: + _asgi_observability_id.set(observability_id) + try: + from policyengine_api.observability import get_runtime + + get_runtime().set_context(observability_id=observability_id) + except Exception: + pass + return observability_id + + +def start_observability_id(value: object = None) -> str: + """Bind the identifier for a newly accepted calculation or report.""" + + current = current_observability_id() + if current is not None: + return current + observability_id = resolve_observability_id( + normalize_observability_id(value) or incoming_observability_id() + ) + return _bind_observability_id(observability_id) + + +def restore_observability_id(value: object) -> str | None: + """Bind a valid identifier read from durable functional state.""" + + observability_id = normalize_observability_id(value) + if observability_id is None: + return current_observability_id() + return _bind_observability_id(observability_id) + + +def resolve_observability_id(value: object) -> str: + """Use a valid caller value or create a new diagnostic identifier.""" + + return normalize_observability_id(value) or generate_observability_id() diff --git a/policyengine_api/routes/household_routes.py b/policyengine_api/routes/household_routes.py index f62d559aa..5f2c65dd9 100644 --- a/policyengine_api/routes/household_routes.py +++ b/policyengine_api/routes/household_routes.py @@ -16,7 +16,10 @@ get_v1_household_read_source, get_v1_household_write_source, ) -from policyengine_api.request_context import current_request_id +from policyengine_api.request_context import ( + current_request_id, + start_observability_id, +) from policyengine_api.response_factory import _make_error_response from policyengine_api.services.household_mirroring import ( HouseholdMirrorUnavailableError, @@ -168,26 +171,60 @@ def _requested_spm(country_id: str, payload: dict): return selection -def _validate_calculation_spm(func): - """Validate current certification before an HTTP cache can satisfy a request.""" - - @wraps(func) - def wrapped(country_id, *args, **kwargs): - payload = request.get_json() - if not isinstance(payload, dict): - raise BadRequest("Calculation payload must be a JSON object.") - try: - selection = _requested_spm(country_id, payload) - g.spm_requested = selection is not None - g.spm = normalize_spm_selection(country_id, selection) - except ValueError as error: - response = _spm_error_response(error) - if response is not None: - return response - raise - return func(country_id, *args, **kwargs) - - return wrapped +def _validate_calculation_request(*, add_missing: bool): + """Validate calculation inputs before an HTTP cache can satisfy a request.""" + + def decorator(func): + @wraps(func) + def wrapped(country_id, *args, **kwargs): + payload = request.get_json() + if not isinstance(payload, dict): + raise BadRequest("Calculation payload must be a JSON object.") + try: + selection = _requested_spm(country_id, payload) + g.spm_requested = selection is not None + g.spm = normalize_spm_selection(country_id, selection) + g.prepared_household_calculation = ( + household_calculation_service.prepare_household_calculation( + country_id, + payload.get("household", {}), + payload.get("policy", {}), + add_missing=add_missing, + **( + { + "spm": g.spm, + "spm_requested": g.spm_requested, + } + if g.spm is not None + else {} + ), + ) + ) + except InvalidHouseholdInputsError as error: + return _make_error_response( + format_unrecognized_inputs_message(error.invalid_inputs), + 400, + result=None, + errors=[ + invalid_input.to_dict() + for invalid_input in error.invalid_inputs + ], + ) + except ValueError as error: + response = _spm_error_response(error) + if response is not None: + return response + raise + result = func(country_id, *args, **kwargs) + if not g.get("calculation_view_executed", False) and ( + not isinstance(result, Response) or result.status_code < 400 + ): + start_observability_id() + return result + + return wrapped + + return decorator def _calculation_cache_key(*args, **kwargs): @@ -331,6 +368,7 @@ def get_household_under_policy(country_id: str, household_id: str, policy_id: st country_id, int(household_id), int(policy_id), + on_accepted=start_observability_id, ) except HouseholdNotFoundError: return _make_error_response( @@ -363,22 +401,35 @@ def get_household_under_policy(country_id: str, household_id: str, policy_id: st return _calculation_response(calculation) -def _calculate(country_id: str, *, add_missing: bool) -> dict | Response: - payload = request.json - household_json = payload.get("household", {}) - policy_json = payload.get("policy", {}) +def _calculate() -> dict | Response: + g.calculation_view_executed = True + try: + g.prepared_household_calculation = ( + household_calculation_service.parse_prepared_household( + g.prepared_household_calculation, + on_accepted=start_observability_id, + ) + ) + except SituationParsingError as error: + return _make_error_response( + f"Invalid household payload: {error}", + 400, + result=None, + ) + except Exception as error: + response = _spm_error_response(error) + if response is not None: + return response + start_observability_id() + logging.exception(error) + return _make_error_response( + f"Error calculating household under policy: {error}", + 500, + ) try: - calculation = household_calculation_service.calculate_household( - country_id, - household_json, - policy_json, - add_missing=add_missing, - **( - {"spm": g.spm, "spm_requested": g.get("spm_requested", False)} - if g.get("spm") is not None - else {} - ), + calculation = household_calculation_service.calculate_prepared_household( + g.prepared_household_calculation, ) except InvalidHouseholdInputsError as error: return _make_error_response( @@ -411,17 +462,17 @@ def _calculate(country_id: str, *, add_missing: bool) -> dict | Response: @household_bp.route("//calculate", methods=["POST"]) @validate_country -@_validate_calculation_spm +@_validate_calculation_request(add_missing=False) @cache.cached(make_cache_key=_calculation_cache_key) def get_calculate(country_id: str) -> dict | Response: """Calculate a household without adding omitted yearly variables.""" - return _calculate(country_id, add_missing=False) + return _calculate() @household_bp.route("//calculate-full", methods=["POST"]) @validate_country -@_validate_calculation_spm +@_validate_calculation_request(add_missing=True) @cache.cached(make_cache_key=_calculation_cache_key) def get_calculate_full(country_id: str) -> dict | Response: """Calculate a household after adding omitted yearly variables.""" - return _calculate(country_id, add_missing=True) + return _calculate() diff --git a/policyengine_api/runtime_cache/reform_impacts.py b/policyengine_api/runtime_cache/reform_impacts.py index 920b4f422..571e5b24c 100644 --- a/policyengine_api/runtime_cache/reform_impacts.py +++ b/policyengine_api/runtime_cache/reform_impacts.py @@ -7,9 +7,11 @@ import time from typing import Any +from policyengine_api.observability.identifiers import normalize_observability_id from policyengine_api.runtime_cache.claims import ExpiringClaimStore from policyengine_api.runtime_cache.core import ( CacheBackend, + CacheCoordinationError, CacheNamespace, decode_envelope, encode_envelope, @@ -22,6 +24,36 @@ REFORM_IMPACT_TTL_SECONDS = 2_592_000 REFORM_IMPACT_INDEX_LIMIT = 1_000 REFORM_IMPACT_START_CLAIM_TTL_SECONDS = 300 +REFORM_IMPACT_START_CLAIM_FAMILY = "reform-impact-start-claim" + + +@dataclass(frozen=True) +class ReformImpactStartClaim: + """Atomic ownership and diagnostic context for one job submission.""" + + submission_claim_id: str + observability_id: str + + def to_payload(self) -> dict[str, str]: + return { + "submission_claim_id": self.submission_claim_id, + "observability_id": self.observability_id, + } + + @classmethod + def from_payload(cls, payload: object) -> "ReformImpactStartClaim | None": + if not isinstance(payload, dict): + return None + submission_claim_id = payload.get("submission_claim_id") + if not isinstance(submission_claim_id, str) or not submission_claim_id: + return None + observability_id = normalize_observability_id(payload.get("observability_id")) + if observability_id is None: + return None + return cls( + submission_claim_id=submission_claim_id, + observability_id=observability_id, + ) @dataclass(frozen=True) @@ -43,6 +75,7 @@ class CachedReformImpact: end_time: datetime | None execution_id: str | None error_code: str | None = None + observability_id: str | None = None def _datetime_to_wire(value: datetime | None) -> str | None: @@ -97,6 +130,7 @@ def _impact_from_wire(payload: Any) -> CachedReformImpact | None: else None ), error_code=payload.get("error_code"), + observability_id=payload.get("observability_id"), ) except (KeyError, TypeError, ValueError): return None @@ -147,7 +181,7 @@ def _start_claim_key( target: str, ) -> str: return self.namespace.key( - "reform-impact-start-claim", + REFORM_IMPACT_START_CLAIM_FAMILY, REFORM_IMPACT_SCHEMA_VERSION, { "api_version": api_version, @@ -162,6 +196,14 @@ def _start_claim_key( }, ) + @staticmethod + def _encoded_start_claim(claim: ReformImpactStartClaim) -> str: + return encode_envelope( + REFORM_IMPACT_START_CLAIM_FAMILY, + REFORM_IMPACT_SCHEMA_VERSION, + claim.to_payload(), + ) + def claim_start( self, *, @@ -175,9 +217,20 @@ def claim_start( options_hash: str, target: str, claim_token: str, + observability_id: str, ) -> bool: """Atomically claim ownership of one reform-impact submission.""" + claim = ReformImpactStartClaim.from_payload( + { + "submission_claim_id": claim_token, + "observability_id": observability_id, + } + ) + if claim is None: + raise ValueError( + "a claim token and valid observability identifier are required" + ) return self._start_claims.acquire( self._start_claim_key( country_id=country_id, @@ -190,10 +243,77 @@ def claim_start( options_hash=options_hash, target=target, ), - claim_token, + self._encoded_start_claim(claim), ttl_seconds=REFORM_IMPACT_START_CLAIM_TTL_SECONDS, ) + def get_start_claim( + self, + *, + country_id: str, + reform_policy_id: int, + baseline_policy_id: int, + region: str, + dataset: str, + time_period: str, + api_version: str, + options_hash: str, + target: str, + ) -> ReformImpactStartClaim | None: + """Read the current submission owner and its diagnostic identifier.""" + + started_at = time.perf_counter() + key = self._start_claim_key( + country_id=country_id, + reform_policy_id=reform_policy_id, + baseline_policy_id=baseline_policy_id, + region=region, + dataset=dataset, + time_period=time_period, + api_version=api_version, + options_hash=options_hash, + target=target, + ) + try: + encoded = self.client.get(key) + except Exception as error: + record_cache_event( + family=self.family, + event="coordination-failed", + operation="claim-read", + started_at=started_at, + severity="WARNING", + ) + raise CacheCoordinationError( + "reform-impact submission ownership is unavailable" + ) from error + + payload = decode_envelope( + encoded, + family=REFORM_IMPACT_START_CLAIM_FAMILY, + schema_version=REFORM_IMPACT_SCHEMA_VERSION, + ) + claim = ReformImpactStartClaim.from_payload(payload) + if encoded is not None and claim is None: + record_cache_event( + family=self.family, + event="decode-failed", + operation="claim-read", + started_at=started_at, + severity="WARNING", + ) + raise CacheCoordinationError( + "reform-impact submission ownership is unreadable" + ) + + record_cache_event( + family=self.family, + event="hit" if claim is not None else "miss", + operation="claim-read", + started_at=started_at, + ) + return claim + def release_start( self, *, @@ -207,8 +327,20 @@ def release_start( options_hash: str, target: str, claim_token: str, + observability_id: str, ) -> bool: - """Release a start claim only when its ownership token still matches.""" + """Release a start claim only when its complete state still matches.""" + + claim = ReformImpactStartClaim.from_payload( + { + "submission_claim_id": claim_token, + "observability_id": observability_id, + } + ) + if claim is None: + raise ValueError( + "a claim token and valid observability identifier are required" + ) return self._start_claims.release( self._start_claim_key( @@ -222,7 +354,7 @@ def release_start( options_hash=options_hash, target=target, ), - claim_token, + self._encoded_start_claim(claim), ) def _record_key(self, execution_id: str) -> str: diff --git a/policyengine_api/services/budget_window_cache.py b/policyengine_api/services/budget_window_cache.py index ae2e0cf6b..9402bc952 100644 --- a/policyengine_api/services/budget_window_cache.py +++ b/policyengine_api/services/budget_window_cache.py @@ -1,7 +1,11 @@ -"""Shared, namespaced budget-window result cache and coordination claims.""" +"""Shared, namespaced budget-window state and coordination claims.""" + +from __future__ import annotations import time -from typing import Any +from typing import Any, ClassVar, Literal + +from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator from policyengine_api.runtime_cache.claims import ExpiringClaimStore from policyengine_api.runtime_cache.core import ( @@ -17,15 +21,59 @@ BUDGET_WINDOW_CACHE_FAMILY = "budget-window" -BUDGET_WINDOW_CACHE_SCHEMA_VERSION = 1 -BUDGET_WINDOW_STARTING_PREFIX = "starting:" +BUDGET_WINDOW_CACHE_SCHEMA_VERSION = 2 BUDGET_WINDOW_STARTING_TTL_SECONDS = 300 BUDGET_WINDOW_BATCH_TTL_SECONDS = 86_400 BUDGET_WINDOW_RESULT_TTL_SECONDS = 2_592_000 +BudgetWindowStateStatus = Literal["starting", "submitted", "completed", "failed"] +BudgetWindowFailureType = Literal["spm_validation", "execution"] + + +class BudgetWindowCacheState(BaseModel): + """One atomic cache document for a budget-window report.""" + + model_config = ConfigDict(frozen=True, strict=True) + + _required_fields: ClassVar[dict[BudgetWindowStateStatus, tuple[str, ...]]] = { + "starting": ("submission_claim_id",), + "submitted": ("batch_job_id",), + "completed": ("result",), + "failed": ("failure_type", "error"), + } + + status: BudgetWindowStateStatus + observability_id: str | None = None + submission_claim_id: str | None = Field(default=None, min_length=1) + batch_job_id: str | None = Field(default=None, min_length=1) + result: dict[str, Any] | None = None + failure_type: BudgetWindowFailureType | None = None + error: dict[str, Any] | None = None + + def to_payload(self) -> dict[str, Any]: + return self.model_dump(exclude_none=True) + + @model_validator(mode="after") + def require_fields_for_status(self) -> BudgetWindowCacheState: + missing = [ + field + for field in self._required_fields[self.status] + if getattr(self, field) is None + ] + if missing: + raise ValueError(f"{self.status} state requires: {', '.join(missing)}") + return self + + @classmethod + def from_payload(cls, payload: object) -> BudgetWindowCacheState | None: + try: + return cls.model_validate(payload) + except ValidationError: + return None + class BudgetWindowCache: - """Recoverable results plus fail-closed expensive-work coordination.""" + """Atomic report state plus fail-closed expensive-work coordination.""" def __init__( self, @@ -71,16 +119,16 @@ def build_key( ) @staticmethod - def _result_key(cache_key: str) -> str: - return f"{cache_key}:result" - - @staticmethod - def _error_key(cache_key: str) -> str: - return f"{cache_key}:terminal-error" + def _state_key(cache_key: str) -> str: + return f"{cache_key}:state" @staticmethod - def _batch_key(cache_key: str) -> str: - return f"{cache_key}:batch-job-id" + def _encoded_state(state: BudgetWindowCacheState) -> str: + return encode_envelope( + BUDGET_WINDOW_CACHE_FAMILY, + BUDGET_WINDOW_CACHE_SCHEMA_VERSION, + state.to_payload(), + ) @staticmethod def _handle_cache_error( @@ -97,148 +145,116 @@ def _handle_cache_error( severity="WARNING", ) - def get_completed_result(self, cache_key: str) -> dict[str, Any] | None: - return self._get_payload(self._result_key(cache_key), "result") - - def get_terminal_error(self, cache_key: str) -> dict[str, str] | None: - """Replay a typed failure independently of completed success payloads.""" - error = self._get_payload(self._error_key(cache_key), "terminal-error") - if ( - error is not None - and set(error) == {"code", "message"} - and isinstance(error["code"], str) - and isinstance(error["message"], str) - ): - return error - return None - - def _get_payload(self, key: str, kind: str) -> dict[str, Any] | None: + def get_state(self, cache_key: str) -> BudgetWindowCacheState | None: + """Read the complete report state or fail closed on cache outage.""" + started_at = time.perf_counter() + state_key = self._state_key(cache_key) try: - payload = self.client.get(key) - except Exception: + encoded = self.client.get(state_key) + except Exception as error: self._handle_cache_error( - f"read-{kind}", - event="connection-failed", + "read-state", + event="coordination-failed", started_at=started_at, ) - return None - result = decode_envelope( - payload, + raise CacheCoordinationError( + "budget-window coordination state is unavailable" + ) from error + + payload = decode_envelope( + encoded, family=BUDGET_WINDOW_CACHE_FAMILY, schema_version=BUDGET_WINDOW_CACHE_SCHEMA_VERSION, ) - if payload is not None and result is None: + state = BudgetWindowCacheState.from_payload(payload) + if encoded is not None and state is None: self._handle_cache_error( - f"decode-{kind}", + "decode-state", event="decode-failed", started_at=started_at, ) - else: - record_cache_event( - family=BUDGET_WINDOW_CACHE_FAMILY, - event="hit" if isinstance(result, dict) else "miss", - operation=f"read-{kind}", - started_at=started_at, - ) - return result if isinstance(result, dict) else None - - def set_completed_result( - self, - cache_key: str, - result: dict[str, Any], - ) -> bool: - return self._set_payload(self._result_key(cache_key), result, "result") - - def set_terminal_error(self, cache_key: str, error: dict[str, str]) -> bool: - """Retain deterministic typed failures for the existing result lifetime.""" - return self._set_payload(self._error_key(cache_key), error, "terminal-error") + self._clear_invalid_state(state_key, encoded) + return None - def _set_payload(self, key: str, result: dict[str, Any], kind: str) -> bool: - started_at = time.perf_counter() - try: - stored = self.client.set( - key, - encode_envelope( - BUDGET_WINDOW_CACHE_FAMILY, - BUDGET_WINDOW_CACHE_SCHEMA_VERSION, - result, - ), - ex=jittered_ttl(BUDGET_WINDOW_RESULT_TTL_SECONDS), - ) - except Exception: - self._handle_cache_error( - f"write-{kind}", - event="write-failed", - started_at=started_at, - ) - return False record_cache_event( family=BUDGET_WINDOW_CACHE_FAMILY, - event="write", - operation=f"write-{kind}", + event="hit" if state is not None else "miss", + operation="read-state", started_at=started_at, ) - return bool(stored) + return state - def get_batch_job_id(self, cache_key: str) -> str | None: - started_at = time.perf_counter() + def _clear_invalid_state(self, state_key: str, encoded: object) -> None: + if isinstance(encoded, bytes): + try: + encoded = encoded.decode("utf-8") + except UnicodeDecodeError: + return + if not isinstance(encoded, str): + return try: - value = self.client.get(self._batch_key(cache_key)) - except Exception as error: - self._handle_cache_error( - "read-batch-id", - event="coordination-failed", - started_at=started_at, - ) - raise CacheCoordinationError( - "budget-window coordination state is unavailable" - ) from error - if not isinstance(value, str) or not value: - record_cache_event( - family=BUDGET_WINDOW_CACHE_FAMILY, - event="coordination-miss", - operation="read-batch-id", - started_at=started_at, - ) - return None - if value.startswith(BUDGET_WINDOW_STARTING_PREFIX): - record_cache_event( - family=BUDGET_WINDOW_CACHE_FAMILY, - event="claim-contended", - operation="read-batch-id", - started_at=started_at, - ) - return None - record_cache_event( - family=BUDGET_WINDOW_CACHE_FAMILY, - event="coordination-hit", - operation="read-batch-id", - started_at=started_at, + self._claims.release(state_key, encoded) + except CacheCoordinationError: + return + + def claim_batch_start( + self, + cache_key: str, + claim_token: str, + observability_id: str | None, + ) -> bool: + state = BudgetWindowCacheState( + status="starting", + observability_id=observability_id, + submission_claim_id=claim_token, + ) + return self._claims.acquire( + self._state_key(cache_key), + self._encoded_state(state), + ttl_seconds=BUDGET_WINDOW_STARTING_TTL_SECONDS, ) - return value - def claim_batch_start(self, cache_key: str, claim_token: str) -> bool: + def clear_starting_claim( + self, + cache_key: str, + claim_token: str, + observability_id: str | None, + ) -> None: + state = BudgetWindowCacheState( + status="starting", + observability_id=observability_id, + submission_claim_id=claim_token, + ) try: - return self._claims.acquire( - self._batch_key(cache_key), - f"{BUDGET_WINDOW_STARTING_PREFIX}{claim_token}", - ttl_seconds=BUDGET_WINDOW_STARTING_TTL_SECONDS, + self._claims.release( + self._state_key(cache_key), + self._encoded_state(state), ) except CacheCoordinationError: - raise + return - def store_batch_job_id(self, cache_key: str, batch_job_id: str) -> None: + def store_submitted( + self, + cache_key: str, + batch_job_id: str, + observability_id: str | None, + ) -> None: + state = BudgetWindowCacheState( + status="submitted", + observability_id=observability_id, + batch_job_id=batch_job_id, + ) started_at = time.perf_counter() try: stored = self.client.set( - self._batch_key(cache_key), - batch_job_id, + self._state_key(cache_key), + self._encoded_state(state), ex=BUDGET_WINDOW_BATCH_TTL_SECONDS, ) except Exception as error: self._handle_cache_error( - "write-batch-id", + "write-submitted-state", event="coordination-failed", started_at=started_at, ) @@ -247,43 +263,95 @@ def store_batch_job_id(self, cache_key: str, batch_job_id: str) -> None: ) from error if not stored: self._handle_cache_error( - "write-batch-id", + "write-submitted-state", event="coordination-failed", started_at=started_at, ) raise CacheCoordinationError( - "budget-window batch identifier could not be stored" + "budget-window submitted state could not be stored" ) record_cache_event( family=BUDGET_WINDOW_CACHE_FAMILY, event="coordination-write", - operation="write-batch-id", + operation="write-submitted-state", started_at=started_at, ) - def clear_starting_claim(self, cache_key: str, claim_token: str) -> None: - try: - self._claims.release( - self._batch_key(cache_key), - f"{BUDGET_WINDOW_STARTING_PREFIX}{claim_token}", - ) - except CacheCoordinationError: - return + def set_completed_result( + self, + cache_key: str, + result: dict[str, Any], + observability_id: str | None, + ) -> bool: + return self._set_recoverable_state( + cache_key, + BudgetWindowCacheState( + status="completed", + observability_id=observability_id, + result=result, + ), + operation="write-completed-state", + ) + + def set_terminal_error( + self, + cache_key: str, + error: dict[str, str], + observability_id: str | None, + ) -> bool: + return self._set_recoverable_state( + cache_key, + BudgetWindowCacheState( + status="failed", + observability_id=observability_id, + failure_type="spm_validation", + error=error, + ), + operation="write-spm-failure-state", + ) + + def set_execution_failure( + self, + cache_key: str, + result: dict[str, Any], + observability_id: str | None, + ) -> bool: + return self._set_recoverable_state( + cache_key, + BudgetWindowCacheState( + status="failed", + observability_id=observability_id, + failure_type="execution", + error=result, + ), + operation="write-execution-failure-state", + ) - def clear_batch_job_id(self, cache_key: str) -> None: + def _set_recoverable_state( + self, + cache_key: str, + state: BudgetWindowCacheState, + *, + operation: str, + ) -> bool: started_at = time.perf_counter() try: - self.client.delete(self._batch_key(cache_key)) + stored = self.client.set( + self._state_key(cache_key), + self._encoded_state(state), + ex=jittered_ttl(BUDGET_WINDOW_RESULT_TTL_SECONDS), + ) except Exception: self._handle_cache_error( - "clear-batch-id", - event="coordination-failed", + operation, + event="write-failed", started_at=started_at, ) - return + return False record_cache_event( family=BUDGET_WINDOW_CACHE_FAMILY, - event="coordination-cleared", - operation="clear-batch-id", + event="write", + operation=operation, started_at=started_at, ) + return bool(stored) diff --git a/policyengine_api/services/economy_service.py b/policyengine_api/services/economy_service.py index 9f3112655..9cec6e90c 100644 --- a/policyengine_api/services/economy_service.py +++ b/policyengine_api/services/economy_service.py @@ -6,7 +6,6 @@ from typing import Any, Literal, Optional import httpx -import numpy as np from dotenv import load_dotenv from policyengine_api.constants import ( COUNTRY_PACKAGE_VERSIONS, @@ -25,6 +24,22 @@ from policyengine_api.data.places import validate_place_code from policyengine_api.gcp_logging import logger from policyengine_api.libs.simulation_entrypoint import simulation_entrypoint +from policyengine_api.observability import ( + runtime as observability_runtime, + set_runtime_context, +) +from policyengine_api.observability.stages import ( + ECONOMY_ANNUAL_STAGES, + ECONOMY_BUDGET_WINDOW_STAGES, + Stage, +) +from policyengine_api.request_context import ( + incoming_observability_id, + resolve_observability_id, + restore_observability_id, + start_observability_id, +) +from policyengine_api.runtime_cache.core import CacheCoordinationError from policyengine_api.services.budget_window_cache import BudgetWindowCache from policyengine_api.services.policy_service import PolicyService from policyengine_api.services.reform_impacts_service import ( @@ -68,6 +83,8 @@ class ImpactStatus(Enum): BUDGET_WINDOW_MAX_YEARS = budget_window_utils.BUDGET_WINDOW_MAX_YEARS BUDGET_WINDOW_MAX_END_YEAR = budget_window_utils.BUDGET_WINDOW_MAX_END_YEAR BUDGET_WINDOW_SUBMISSION_VALIDATION_ERROR_STATUS_CODES = {400, 422} +BUDGET_WINDOW_CLAIM_ATTEMPTS = 3 +REFORM_IMPACT_START_CLAIM_ATTEMPTS = 3 class SimulationOptions(BaseModel): @@ -84,7 +101,8 @@ class SimulationOptions(BaseModel): class EconomicImpactSetupOptions(BaseModel): - process_id: str + submission_claim_id: str + observability_id: str | None = None country_id: str reform_policy_id: int baseline_policy_id: int @@ -273,6 +291,7 @@ def _budget_window_cache(self) -> BudgetWindowCache: def _simulation_gateway(self): return self._injected_simulation_entrypoint or simulation_entrypoint + @observability_runtime.span(ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_LOAD_POLICIES)) def _get_policy_jsons( self, country_id: str, @@ -301,6 +320,7 @@ def _parse_json_object(value: dict[str, Any] | str) -> dict[str, Any]: raise TypeError("Expected a JSON object") return parsed + @observability_runtime.span(ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_REQUEST)) def get_economic_impact( self, country_id: str, @@ -322,6 +342,13 @@ def get_economic_impact( the status is "computing" or "error". """ + set_runtime_context( + country_id=country_id, + policy_id=policy_id, + baseline_policy_id=baseline_policy_id, + simulation_year=time_period, + ) + economic_impact_setup_options: EconomicImpactSetupOptions | None = None try: # Normalize region early for US; this allows us to accommodate legacy # regions that don't contain a region prefix. @@ -345,9 +372,29 @@ def get_economic_impact( ) except Exception as e: - print(f"Error getting economic impact: {str(e)}") - raise e + logger.log_struct( + { + "message": "Error getting economic impact", + "error_type": type(e).__name__, + }, + severity="ERROR", + ) + raise + finally: + if ( + economic_impact_setup_options is not None + and economic_impact_setup_options.observability_id is not None + ): + # The report identifier is selected inside a nested stage span. + # Reapply it here while the outer economy-request span is current + # so an identifier query returns that complete request duration. + set_runtime_context( + observability_id=(economic_impact_setup_options.observability_id) + ) + @observability_runtime.span( + ECONOMY_BUDGET_WINDOW_STAGES.name(Stage.ECONOMY_BUDGET_WINDOW_REQUEST) + ) def get_budget_window_economic_impact( self, country_id: str, @@ -362,6 +409,13 @@ def get_budget_window_economic_impact( target: Literal["general", "cliff"] = "general", max_active_years: int = BUDGET_WINDOW_MAX_ACTIVE_YEARS, ) -> BudgetWindowEconomicImpactResult: + set_runtime_context( + country_id=country_id, + policy_id=policy_id, + baseline_policy_id=baseline_policy_id, + start_year=start_year, + window_size=window_size, + ) try: if country_id == "us": region = normalize_us_region(region) @@ -386,12 +440,30 @@ def get_budget_window_economic_impact( ) cache_key = self._build_budget_window_cache_key(setup_options) - cached_error = self._budget_window_cache.get_terminal_error(cache_key) - if cached_error is not None: - raise SPMValidationError(**cached_error) + cached_state = self._budget_window_cache.get_state(cache_key) + if cached_state is not None: + self._restore_budget_window_observability_id( + cached_state.observability_id, + setup_options=setup_options, + ) + + if cached_state is not None and cached_state.status == "failed": + if cached_state.failure_type == "spm_validation": + cached_error = cached_state.error or {} + raise SPMValidationError( + code=str(cached_error.get("code", "SPM_VALIDATION_ERROR")), + message=str( + cached_error.get( + "message", "Stored budget-window validation failed" + ) + ), + ) + return BudgetWindowEconomicImpactResult.model_validate( + cached_state.error + ).model_copy(update={"cache_status": "failure-hit"}) - cached_result = self._budget_window_cache.get_completed_result(cache_key) - if cached_result is not None: + if cached_state is not None and cached_state.status == "completed": + cached_result = cached_state.result or {} try: validate_worker_result( cached_result, @@ -406,7 +478,9 @@ def get_budget_window_economic_impact( # and the read above replays it, instead of re-deriving the # same failure from the same payload on every later poll. self._budget_window_cache.set_terminal_error( - cache_key, error.to_dict() + cache_key, + error.to_dict(), + setup_options.observability_id, ) raise return BudgetWindowEconomicImpactResult.completed( @@ -414,20 +488,58 @@ def get_budget_window_economic_impact( cache_status="result-hit", ) - batch_job_id = self._budget_window_cache.get_batch_job_id(cache_key) - if batch_job_id: + if cached_state is not None and cached_state.status == "submitted": return self._get_budget_window_result_from_batch_job_id( - batch_job_id=batch_job_id, + batch_job_id=cached_state.batch_job_id or "", spm=setup_options.options.get("spm"), cache_key=cache_key, total_years=len(years), queued_years_on_submit=years, cache_status="batch-id-hit", + observability_id=setup_options.observability_id, + ) + + if cached_state is not None and cached_state.status == "starting": + return self._build_budget_window_computing_result( + total_years=len(years), + completed_years=[], + computing_years=[], + queued_years=years, + progress=0, + cache_status="starting-claim-hit", ) - claim_token = setup_options.process_id + claim_token = setup_options.submission_claim_id cache_status = "starting-claim-hit" - if self._budget_window_cache.claim_batch_start(cache_key, claim_token): + observability_id_candidate = resolve_observability_id( + incoming_observability_id() + ) + owns_claim = False + claimed_state = None + for _attempt in range(BUDGET_WINDOW_CLAIM_ATTEMPTS): + owns_claim = self._budget_window_cache.claim_batch_start( + cache_key, + claim_token, + observability_id_candidate, + ) + if owns_claim: + setup_options.observability_id = start_observability_id( + observability_id_candidate + ) + break + claimed_state = self._budget_window_cache.get_state(cache_key) + if claimed_state is not None: + self._restore_budget_window_observability_id( + claimed_state.observability_id, + setup_options=setup_options, + ) + break + else: + raise CacheCoordinationError( + "budget-window submission ownership changed repeatedly" + ) + + if owns_claim: cache_status = "miss" try: batch_execution = self._start_budget_window_batch( @@ -436,29 +548,40 @@ def get_budget_window_economic_impact( window_size=window_size, max_parallel=max_active_years, ) - self._budget_window_cache.store_batch_job_id( - cache_key, batch_execution.batch_job_id + self._budget_window_cache.store_submitted( + cache_key, + batch_execution.batch_job_id, + setup_options.observability_id, ) except httpx.HTTPStatusError as error: - self._budget_window_cache.clear_starting_claim( - cache_key, claim_token - ) if ( error.response.status_code in BUDGET_WINDOW_SUBMISSION_VALIDATION_ERROR_STATUS_CODES ): - return BudgetWindowEconomicImpactResult.failed( + failed_result = BudgetWindowEconomicImpactResult.failed( self._build_budget_window_submission_error_message(error), queued_years=years, cache_status=cache_status, ) + self._budget_window_cache.set_execution_failure( + cache_key, + failed_result.model_dump(mode="json"), + setup_options.observability_id, + ) + return failed_result + self._budget_window_cache.clear_starting_claim( + cache_key, + claim_token, + setup_options.observability_id, + ) raise except Exception: self._budget_window_cache.clear_starting_claim( - cache_key, claim_token + cache_key, + claim_token, + setup_options.observability_id, ) raise - return self._build_budget_window_computing_result( total_years=len(years), completed_years=[], @@ -468,8 +591,14 @@ def get_budget_window_economic_impact( cache_status=cache_status, ) except Exception as e: - print(f"Error getting budget-window economic impact: {str(e)}") - raise e + logger.log_struct( + { + "message": "Error getting budget-window economic impact", + "error_type": type(e).__name__, + }, + severity="ERROR", + ) + raise def _build_budget_window_cache_key( self, @@ -486,6 +615,18 @@ def _build_budget_window_cache_key( api_version=setup_options.api_version, ) + @staticmethod + def _restore_budget_window_observability_id( + value: str | None, + *, + setup_options: EconomicImpactSetupOptions, + ) -> str | None: + resolved = restore_observability_id(value) + if resolved is None: + return setup_options.observability_id + setup_options.observability_id = resolved + return resolved + def _build_budget_window_batch_payload( self, *, @@ -519,6 +660,9 @@ def _build_budget_window_batch_payload( sim_params["target"] = setup_options.target return sim_params + @observability_runtime.span( + ECONOMY_BUDGET_WINDOW_STAGES.name(Stage.ECONOMY_START_BUDGET_WINDOW) + ) def _start_budget_window_batch( self, *, @@ -567,6 +711,9 @@ def _build_budget_window_submission_error_message( return str(error) + @observability_runtime.span( + ECONOMY_BUDGET_WINDOW_STAGES.name(Stage.ECONOMY_POLL_BUDGET_WINDOW) + ) def _get_budget_window_result_from_batch_job_id( self, *, @@ -576,7 +723,9 @@ def _get_budget_window_result_from_batch_job_id( queued_years_on_submit: list[str], spm: dict | None = None, cache_status: Optional[str] = None, + observability_id: str | None, ) -> BudgetWindowEconomicImpactResult: + resolved_observability_id = observability_id try: batch_execution = self._simulation_gateway.get_budget_window_batch_by_id( batch_job_id @@ -588,26 +737,34 @@ def _get_budget_window_result_from_batch_job_id( result, spm, expected_years=queued_years_on_submit ) except SPMValidationError as error: - if self._budget_window_cache.set_terminal_error(cache_key, error.to_dict()): - self._budget_window_cache.clear_batch_job_id(cache_key) + self._budget_window_cache.set_terminal_error( + cache_key, + error.to_dict(), + resolved_observability_id, + ) raise if batch_execution.status in EXECUTION_STATUSES_SUCCESS: result = batch_execution.result if not isinstance(result, dict) or not result: - self._budget_window_cache.clear_batch_job_id(cache_key) - return BudgetWindowEconomicImpactResult.failed( + failed_result = BudgetWindowEconomicImpactResult.failed( "Budget-window batch completed without a result", completed_years=batch_execution.completed_years, computing_years=batch_execution.running_years, queued_years=batch_execution.queued_years or queued_years_on_submit, cache_status=cache_status, ) - result_stored = self._budget_window_cache.set_completed_result( - cache_key, result + self._budget_window_cache.set_execution_failure( + cache_key, + failed_result.model_dump(mode="json"), + resolved_observability_id, + ) + return failed_result + self._budget_window_cache.set_completed_result( + cache_key, + result, + resolved_observability_id, ) - if result_stored: - self._budget_window_cache.clear_batch_job_id(cache_key) return BudgetWindowEconomicImpactResult.completed( result, cache_status=cache_status, @@ -615,14 +772,19 @@ def _get_budget_window_result_from_batch_job_id( if batch_execution.status in EXECUTION_STATUSES_FAILURE: error_message = batch_execution.error or "Budget-window batch failed" - self._budget_window_cache.clear_batch_job_id(cache_key) - return BudgetWindowEconomicImpactResult.failed( + failed_result = BudgetWindowEconomicImpactResult.failed( error_message, completed_years=batch_execution.completed_years, computing_years=batch_execution.running_years, queued_years=batch_execution.queued_years or queued_years_on_submit, cache_status=cache_status, ) + self._budget_window_cache.set_execution_failure( + cache_key, + failed_result.model_dump(mode="json"), + resolved_observability_id, + ) + return failed_result if batch_execution.status in EXECUTION_STATUSES_PENDING: return self._build_budget_window_computing_result( @@ -692,7 +854,7 @@ def _build_economic_impact_setup_options( ) if resolved_spm is not None: options = {**options, "spm": resolved_spm} - process_id: str = self._create_process_id() + submission_claim_id = self._create_submission_claim_id() cache_version = get_economy_impact_cache_version(country_id, api_version) country_package_version = COUNTRY_PACKAGE_VERSIONS.get(country_id) resolved_dataset = "default" @@ -710,7 +872,8 @@ def _build_economic_impact_setup_options( return EconomicImpactSetupOptions.model_validate( { - "process_id": process_id, + "submission_claim_id": submission_claim_id, + "observability_id": None, "country_id": country_id, "reform_policy_id": policy_id, "baseline_policy_id": baseline_policy_id, @@ -728,6 +891,9 @@ def _build_economic_impact_setup_options( } ) + @observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_RESOLVE_CACHED_OR_NEW) + ) def _get_or_create_economic_impact( self, setup_options: EconomicImpactSetupOptions ) -> EconomicImpactResult: @@ -742,7 +908,6 @@ def _get_or_create_economic_impact( most_recent_impact: dict | None = self._get_most_recent_impact( setup_options=setup_options ) - if most_recent_impact and self._should_refresh_cached_impact( setup_options=setup_options, most_recent_impact=most_recent_impact, @@ -757,6 +922,15 @@ def _get_or_create_economic_impact( impact_action: ImpactAction = self._determine_impact_action( most_recent_impact=most_recent_impact ) + if most_recent_impact is not None: + setup_options.observability_id = restore_observability_id( + getattr(most_recent_impact, "observability_id", None) + ) + if impact_action != ImpactAction.CREATE: + observability_runtime.event( + "economy.cache_decision", + attributes={"cache_event": impact_action.value}, + ) if impact_action == ImpactAction.COMPLETED: logger.log_struct( @@ -788,8 +962,45 @@ def _get_or_create_economic_impact( ) if impact_action == ImpactAction.CREATE: - self._resolve_runtime_bundle_for_setup_options(setup_options) - if not self._claim_reform_impact_start(setup_options): + existing_claim = None + with observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_RESOLVE_RUNTIME_BUNDLE) + ): + self._resolve_runtime_bundle_for_setup_options(setup_options) + observability_id_candidate = resolve_observability_id( + incoming_observability_id() + ) + for _attempt in range(REFORM_IMPACT_START_CLAIM_ATTEMPTS): + if self._claim_reform_impact_start( + setup_options, + observability_id_candidate, + ): + setup_options.observability_id = start_observability_id( + observability_id_candidate + ) + break + existing_claim = self._get_reform_impact_start_claim(setup_options) + if existing_claim is None: + continue + setup_options.observability_id = restore_observability_id( + existing_claim.observability_id + ) + break + else: + raise CacheCoordinationError( + "reform-impact submission ownership changed repeatedly" + ) + + # The identifier above was selected while the runtime-resolution + # stage was current. Reapply it now so the containing cache-decision + # span carries the same report identifier. + if setup_options.observability_id is not None: + set_runtime_context(observability_id=setup_options.observability_id) + observability_runtime.event( + "economy.cache_decision", + attributes={"cache_event": impact_action.value}, + ) + if existing_claim is not None: logger.log_struct( { "message": "Another request owns this reform-impact submission", @@ -845,7 +1056,7 @@ def _resolve_runtime_bundle_for_setup_options( runtime_app_name=setup_options.runtime_app_name, ) - def _reform_impact_start_claim_arguments( + def _reform_impact_start_claim_scope_arguments( self, setup_options: EconomicImpactSetupOptions, ) -> dict[str, Any]: @@ -861,15 +1072,40 @@ def _reform_impact_start_claim_arguments( "options_hash": setup_options.options_hash, "api_version": setup_options.api_version, "target": setup_options.target, - "claim_token": setup_options.process_id, + } + + def _reform_impact_start_claim_arguments( + self, + setup_options: EconomicImpactSetupOptions, + observability_id: str | None = None, + ) -> dict[str, Any]: + resolved_observability_id = observability_id or setup_options.observability_id + if resolved_observability_id is None: + raise ValueError("reform-impact observability identifier is required") + return { + **self._reform_impact_start_claim_scope_arguments(setup_options), + "claim_token": setup_options.submission_claim_id, + "observability_id": resolved_observability_id, } def _claim_reform_impact_start( self, setup_options: EconomicImpactSetupOptions, + observability_id: str, ) -> bool: return self._reform_impacts.claim_reform_impact_start( - **self._reform_impact_start_claim_arguments(setup_options) + **self._reform_impact_start_claim_arguments( + setup_options, + observability_id, + ) + ) + + def _get_reform_impact_start_claim( + self, + setup_options: EconomicImpactSetupOptions, + ): + return self._reform_impacts.get_reform_impact_start_claim( + **self._reform_impact_start_claim_scope_arguments(setup_options) ) def _release_reform_impact_start( @@ -976,6 +1212,9 @@ def _determine_impact_action( else: raise ValueError(f"Unknown impact status: {status}") + @observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_HANDLE_EXECUTION_STATE) + ) def _handle_execution_state( self, setup_options: EconomicImpactSetupOptions, @@ -1044,6 +1283,9 @@ def _handle_execution_state( else: raise ValueError(f"Unexpected sim API execution state: {execution_state}") + @observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_READ_COMPLETED) + ) def _handle_completed_impact( self, setup_options: EconomicImpactSetupOptions, @@ -1103,6 +1345,7 @@ def _record_uncertifiable_stored_impact( except Exception: pass + @observability_runtime.span(ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_READ_FAILED)) def _handle_failed_impact( self, most_recent_impact: ReformImpact, @@ -1123,6 +1366,7 @@ def _handle_failed_impact( ) ) + @observability_runtime.span(ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_POLL_ACTIVE)) def _handle_computing_impact( self, setup_options: EconomicImpactSetupOptions, @@ -1148,6 +1392,7 @@ def _handle_computing_impact( ) raise + @observability_runtime.span(ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_SUBMIT)) def _handle_create_impact( self, setup_options: EconomicImpactSetupOptions, @@ -1180,16 +1425,16 @@ def _handle_create_impact( logger.log_struct( { "message": "Setting up sim API job", - "run_id": telemetry["run_id"], + "observability_id": setup_options.observability_id, **setup_options.model_dump(), } ) - # Preserve both legacy metadata and the new telemetry envelope. + # Preserve execution metadata and non-identity simulation telemetry. sim_params["_metadata"] = { "reform_policy_id": setup_options.reform_policy_id, "baseline_policy_id": setup_options.baseline_policy_id, - "process_id": setup_options.process_id, + "submission_claim_id": setup_options.submission_claim_id, "model_version": setup_options.model_version, "policyengine_version": setup_options.policyengine_version, "data_version": setup_options.data_version, @@ -1213,15 +1458,11 @@ def _handle_create_impact( entrypoint_execution ) - run_id = ( - getattr(entrypoint_execution, "run_id", None) or telemetry["run_id"] - ) - progress_log = { **setup_options.model_dump(), "message": "Sim API job started", "execution_id": execution_id, - "run_id": run_id, + "observability_id": setup_options.observability_id, } logger.log_struct(progress_log, severity="INFO") @@ -1442,9 +1683,7 @@ def _build_simulation_telemetry( ) return { - "run_id": str(uuid.uuid4()), - "process_id": setup_options.process_id, - "traceparent": self._get_current_traceparent(), + "submission_claim_id": setup_options.submission_claim_id, "requested_at": datetime.datetime.now(datetime.UTC).isoformat(), "simulation_kind": simulation_kind, "geography_code": geography_code, @@ -1479,27 +1718,13 @@ def _stable_config_hash(self, payload: dict[str, Any]) -> str: ).encode("utf-8") return f"sha256:{hashlib.sha256(encoded).hexdigest()}" - def _get_current_traceparent(self) -> str | None: - try: - from opentelemetry import trace - except Exception: - return None - - span = trace.get_current_span() - span_context = span.get_span_context() - if not getattr(span_context, "is_valid", False): - return None - - trace_flags = int(getattr(span_context, "trace_flags", 0)) - return ( - f"00-{span_context.trace_id:032x}-" - f"{span_context.span_id:016x}-{trace_flags:02x}" - ) - # Note: The following methods that interface with the ReformImpactsService # are written separately because the service relies upon mutating an original # 'computing' record to 'ok' or 'error' status, rather than creating a new record. # This should be addressed in the future. + @observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_PERSIST_COMPUTING) + ) def _set_reform_impact_computing( self, setup_options: EconomicImpactSetupOptions, @@ -1523,6 +1748,7 @@ def _set_reform_impact_computing( reform_impact_json={}, start_time=datetime.datetime.now(), execution_id=execution_id, + observability_id=setup_options.observability_id, ) except Exception as e: logger.log_struct( @@ -1533,6 +1759,9 @@ def _set_reform_impact_computing( ) raise e + @observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_PERSIST_COMPLETED) + ) def _set_reform_impact_complete( self, setup_options: EconomicImpactSetupOptions, @@ -1563,6 +1792,9 @@ def _set_reform_impact_complete( ) raise e + @observability_runtime.span( + ECONOMY_ANNUAL_STAGES.name(Stage.ECONOMY_PERSIST_FAILED) + ) def _set_reform_impact_error( self, setup_options: EconomicImpactSetupOptions, @@ -1595,11 +1827,7 @@ def _set_reform_impact_error( ) raise e - def _create_process_id(self) -> str: - """ - Generate a unique process ID based on the current timestamp and a random number. - This is used to track the process in the database and logs. - """ - timestamp = datetime.datetime.now().strftime("%Y%m%d%H%M%S") - random_number = np.random.randint(1000, 9999) - return f"job_{timestamp}_{random_number}" + def _create_submission_claim_id(self) -> str: + """Create an opaque token for one submission ownership claim.""" + + return str(uuid.uuid4()) diff --git a/policyengine_api/services/household_calculation_service.py b/policyengine_api/services/household_calculation_service.py index 8eb26e11f..3d6284567 100644 --- a/policyengine_api/services/household_calculation_service.py +++ b/policyengine_api/services/household_calculation_service.py @@ -1,7 +1,7 @@ from __future__ import annotations from copy import deepcopy -from dataclasses import dataclass +from dataclasses import dataclass, replace from datetime import date import time from typing import Any, Callable @@ -10,6 +10,8 @@ from sqlalchemy.orm import Session, sessionmaker from policyengine_api.constants import COUNTRY_PACKAGE_VERSIONS, POLICYENGINE_VERSION +from policyengine_api.observability import runtime as observability_runtime +from policyengine_api.observability.stages import HOUSEHOLD_STAGES, Stage from policyengine_api.data.orm import get_v1_session_factory from policyengine_api.data.v1_models import ( Household, @@ -44,6 +46,19 @@ class HouseholdCalculationResult: spm_provenance: dict | None = None +@dataclass(frozen=True) +class PreparedHouseholdCalculation: + """Validated request inputs ready for calculation after cache lookup.""" + + country: Any + household_json: dict + policy_json: dict + spm: dict | None + spm_requested: bool + warnings: tuple[str, ...] + country_calculation: Any | None = None + + class HouseholdNotFoundError(LookupError): pass @@ -165,6 +180,7 @@ def _get_inputs( ) return household, policy + @observability_runtime.span(HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_CACHE_WRITE)) def _store_result( self, identity: HouseholdCalculationIdentity, @@ -186,18 +202,32 @@ def calculate_stored_household( country_id: str, household_id: int, policy_id: int, + *, + on_accepted: Callable[[], object] | None = None, ) -> HouseholdCalculationResult: api_version = COUNTRY_PACKAGE_VERSIONS[country_id] - household, policy = self._get_inputs(country_id, household_id, policy_id) - if household is None: - raise HouseholdNotFoundError(household_id) - if policy is None: - raise PolicyNotFoundError(policy_id) - household_inputs = deepcopy(household.household_json) - # A household saved without a selection never chose a measurement, so its - # replay keeps the historical output set rather than failing closed. - saved_spm = household_inputs.pop("spm", None) - spm = normalize_spm_selection(country_id, saved_spm, stored=True) + with observability_runtime.span( + HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_LOAD_INPUTS) + ): + household, policy = self._get_inputs( + country_id, + household_id, + policy_id, + ) + if household is None: + raise HouseholdNotFoundError(household_id) + if policy is None: + raise PolicyNotFoundError(policy_id) + household_inputs = deepcopy(household.household_json) + # A household saved without a selection never chose a measurement, + # so its replay keeps the historical output set rather than failing + # closed. + saved_spm = household_inputs.pop("spm", None) + spm = normalize_spm_selection(country_id, saved_spm, stored=True) + if on_accepted is not None: + # Bind before this stage ends so its span carries the same + # identifier as the cache and calculation stages that follow. + on_accepted() cache_identity = self._cache_identity( country_id, household, @@ -205,7 +235,10 @@ def calculate_stored_household( api_version, spm, ) - cached = self._cache.get(cache_identity) + with observability_runtime.span( + HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_CACHE_LOOKUP) + ): + cached = self._cache.get(cache_identity) if cached is not None: return HouseholdCalculationResult( household=cached.household, @@ -215,34 +248,40 @@ def calculate_stored_household( spm_provenance=cached.spm_provenance, ) - countries = self._countries() - country = countries.get(country_id) - household_json = add_yearly_variables( - household_inputs, - country_id, - countries, - ) - deprecated_inputs = drop_deprecated_inputs(household_json) - household_json = deprecated_inputs.household - invalid_inputs = find_unrecognized_inputs( - household_json, - policy.policy_json, - country.metadata, - ) + with observability_runtime.span( + HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_INPUT_NORMALIZATION) + ): + countries = self._countries() + country = countries.get(country_id) + household_json = add_yearly_variables( + household_inputs, + country_id, + countries, + ) + deprecated_inputs = drop_deprecated_inputs(household_json) + household_json = deprecated_inputs.household + invalid_inputs = find_unrecognized_inputs( + household_json, + policy.policy_json, + country.metadata, + ) if invalid_inputs: raise InvalidHouseholdInputsError(invalid_inputs) calculation_started_at = time.perf_counter() try: - raw_calculation = country.calculate( - household_json, - policy.policy_json, - **( - {"spm": spm, "spm_requested": saved_spm is not None} - if spm is not None - else {} - ), - ) + with observability_runtime.span( + HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_CALCULATION) + ): + raw_calculation = country.calculate( + household_json, + policy.policy_json, + **( + {"spm": spm, "spm_requested": saved_spm is not None} + if spm is not None + else {} + ), + ) except Exception: record_cache_event( family="household-calculation", @@ -287,7 +326,7 @@ def calculate_stored_household( spm_provenance=calculation.spm_provenance, ) - def calculate_household( + def prepare_household_calculation( self, country_id: str, household_json: dict, @@ -296,8 +335,9 @@ def calculate_household( add_missing: bool = False, spm: dict | None = None, spm_requested: bool = False, - ) -> HouseholdCalculationResult: - """Validate and calculate request-provided household and policy data.""" + ) -> PreparedHouseholdCalculation: + """Validate request inputs before accepting a calculation.""" + countries = self._countries() country = countries.get(country_id) spm = normalize_spm_selection(country_id, spm) @@ -319,11 +359,39 @@ def calculate_household( if invalid_inputs: raise InvalidHouseholdInputsError(invalid_inputs) - raw_calculation = country.calculate( - household_json, - policy_json, - **({"spm": spm, "spm_requested": spm_requested} if spm is not None else {}), + return PreparedHouseholdCalculation( + country=country, + household_json=household_json, + policy_json=policy_json, + spm=spm, + spm_requested=spm_requested, + warnings=tuple(warning.message for warning in deprecated_inputs.warnings), ) + + def calculate_prepared_household( + self, + prepared: PreparedHouseholdCalculation, + ) -> HouseholdCalculationResult: + """Calculate inputs already accepted by ``prepare_household_calculation``.""" + + with observability_runtime.span( + HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_CALCULATION) + ): + calculation_options = ( + { + "spm": prepared.spm, + "spm_requested": prepared.spm_requested, + } + if prepared.spm is not None + else {} + ) + if prepared.country_calculation is not None: + calculation_options["prepared"] = prepared.country_calculation + raw_calculation = prepared.country.calculate( + prepared.household_json, + prepared.policy_json, + **calculation_options, + ) if isinstance(raw_calculation, dict): household = raw_calculation calculation_warnings = () @@ -332,10 +400,61 @@ def calculate_household( calculation_warnings = tuple(getattr(raw_calculation, "warnings", ())) return HouseholdCalculationResult( household=household, - warnings=( - tuple(warning.message for warning in deprecated_inputs.warnings) - + calculation_warnings - ), + warnings=(prepared.warnings + calculation_warnings), spm_config=getattr(raw_calculation, "spm_config", None), spm_provenance=getattr(raw_calculation, "spm_provenance", None), ) + + @observability_runtime.span( + HOUSEHOLD_STAGES.name(Stage.HOUSEHOLD_INPUT_NORMALIZATION) + ) + def parse_prepared_household( + self, + prepared: PreparedHouseholdCalculation, + *, + on_accepted: Callable[[], object] | None = None, + ) -> PreparedHouseholdCalculation: + """Parse a prepared situation without performing requested calculations.""" + + prepare_calculation = getattr(prepared.country, "prepare_calculation", None) + country_calculation = prepared.country_calculation + if callable(prepare_calculation): + country_calculation = prepare_calculation( + prepared.household_json, + prepared.policy_json, + **( + { + "spm": prepared.spm, + "spm_requested": prepared.spm_requested, + } + if prepared.spm is not None + else {} + ), + ) + if on_accepted is not None: + # Situation parsing succeeded. Bind before this stage span ends so + # it participates in the accepted calculation's identifier query. + on_accepted() + return replace(prepared, country_calculation=country_calculation) + + def calculate_household( + self, + country_id: str, + household_json: dict, + policy_json: dict, + *, + add_missing: bool = False, + spm: dict | None = None, + spm_requested: bool = False, + ) -> HouseholdCalculationResult: + """Validate and calculate request-provided household and policy data.""" + + prepared = self.prepare_household_calculation( + country_id, + household_json, + policy_json, + add_missing=add_missing, + spm=spm, + spm_requested=spm_requested, + ) + return self.calculate_prepared_household(prepared) diff --git a/policyengine_api/services/reform_impacts_service.py b/policyengine_api/services/reform_impacts_service.py index 81e66cad2..510895636 100644 --- a/policyengine_api/services/reform_impacts_service.py +++ b/policyengine_api/services/reform_impacts_service.py @@ -8,6 +8,7 @@ from policyengine_api.runtime_cache.reform_impacts import ( CachedReformImpact, ReformImpactCache, + ReformImpactStartClaim, reform_impact_id, ) @@ -90,6 +91,7 @@ def claim_reform_impact_start( api_version: str, target: str, claim_token: str, + observability_id: str, ) -> bool: """Fail closed unless this request atomically owns job submission.""" @@ -104,6 +106,34 @@ def claim_reform_impact_start( options_hash=options_hash, target=target, claim_token=claim_token, + observability_id=observability_id, + ) + + def get_reform_impact_start_claim( + self, + *, + country_id: str, + policy_id: int, + baseline_policy_id: int, + region: str, + dataset: str, + time_period: str, + options_hash: str, + api_version: str, + target: str, + ) -> ReformImpactStartClaim | None: + """Return the current submission owner and diagnostic identifier.""" + + return self._cache.get_start_claim( + country_id=country_id, + reform_policy_id=policy_id, + baseline_policy_id=baseline_policy_id, + region=region, + dataset=dataset, + time_period=time_period, + api_version=api_version, + options_hash=options_hash, + target=target, ) def release_reform_impact_start( @@ -119,6 +149,7 @@ def release_reform_impact_start( api_version: str, target: str, claim_token: str, + observability_id: str, ) -> None: """Best-effort release; an unavailable cache safely falls back to expiry.""" @@ -134,6 +165,7 @@ def release_reform_impact_start( options_hash=options_hash, target=target, claim_token=claim_token, + observability_id=observability_id, ) except CacheCoordinationError: pass @@ -153,6 +185,7 @@ def set_reform_impact( reform_impact_json: dict[str, Any], start_time, execution_id: str, + observability_id: str | None = None, ) -> CachedReformImpact: impact = CachedReformImpact( reform_impact_id=reform_impact_id(execution_id), @@ -171,6 +204,7 @@ def set_reform_impact( start_time=start_time, end_time=None, execution_id=execution_id, + observability_id=observability_id, ) if not self._cache.set(impact): raise ReformImpactHandoffError( diff --git a/pyproject.toml b/pyproject.toml index 5bf85c01f..595109844 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,6 +39,7 @@ dependencies = [ "microdf_python>=1.0.0", "openai", "packaging>=24,<27", + "policyengine-observability[flask,google,httpx,otlp-grpc]>=3.0.1,<4", "policyengine_canada==0.96.3", "policyengine-ng==0.5.1", "policyengine-il==0.1.0", @@ -74,6 +75,9 @@ dev = [ [tool.hatch.build.targets.wheel] packages = ["policyengine_api"] +[tool.hatch.metadata] +allow-direct-references = true + [tool.ruff.lint.per-file-ignores] "migrations/*/versions/*.py" = ["F401"] diff --git a/tests/contract/test_v1_route_contracts.py b/tests/contract/test_v1_route_contracts.py index dee9fc76c..b4dce5501 100644 --- a/tests/contract/test_v1_route_contracts.py +++ b/tests/contract/test_v1_route_contracts.py @@ -402,7 +402,7 @@ def _patched_route_dependencies(): ) stack.enter_context( patch( - "policyengine_api.routes.household_routes.household_calculation_service.calculate_household", + "policyengine_api.routes.household_routes.household_calculation_service.calculate_prepared_household", return_value=HouseholdCalculationResult( household={ "people": {"you": {"age": {"2026": 40}}}, diff --git a/tests/fixtures/libs/simulation_entrypoint.py b/tests/fixtures/libs/simulation_entrypoint.py index 9a11b4c90..47263c121 100644 --- a/tests/fixtures/libs/simulation_entrypoint.py +++ b/tests/fixtures/libs/simulation_entrypoint.py @@ -18,7 +18,7 @@ # Mock data constants MOCK_MODAL_JOB_ID = "fc-abc123xyz" -MOCK_RUN_ID = "run-abc123xyz" +MOCK_OBSERVABILITY_ID = "00000000-0000-4000-8000-000000000001" MOCK_BATCH_JOB_ID = "fc-batch123xyz" MOCK_MODAL_BASE_URL = "https://test-modal-api.modal.run" @@ -35,8 +35,7 @@ MOCK_SIMULATION_PAYLOAD_WITH_TELEMETRY = { **MOCK_SIMULATION_PAYLOAD, "_telemetry": { - "run_id": MOCK_RUN_ID, - "process_id": "job_20250626120000_1234", + "submission_claim_id": "job_20250626120000_1234", "capture_mode": "disabled", }, } @@ -59,7 +58,6 @@ MOCK_SUBMIT_RESPONSE_SUCCESS = { "job_id": MOCK_MODAL_JOB_ID, - "run_id": MOCK_RUN_ID, "status": MODAL_EXECUTION_STATUS_SUBMITTED, "poll_url": f"/jobs/{MOCK_MODAL_JOB_ID}", "country": "us", @@ -142,6 +140,7 @@ def create_mock_httpx_response( status_code: int = 200, json_data: dict = None, + headers: dict | None = None, ): """ Helper function to create a mock httpx response. @@ -160,6 +159,9 @@ def create_mock_httpx_response( """ mock_response = MagicMock() mock_response.status_code = status_code + mock_response.headers = headers or { + "X-PolicyEngine-Observability-Id": MOCK_OBSERVABILITY_ID + } mock_response.json.return_value = json_data or {} mock_response.text = json.dumps(json_data or {}) mock_response.raise_for_status = MagicMock() diff --git a/tests/fixtures/services/economy_service.py b/tests/fixtures/services/economy_service.py index cd40425b2..666f93880 100644 --- a/tests/fixtures/services/economy_service.py +++ b/tests/fixtures/services/economy_service.py @@ -1,6 +1,7 @@ import datetime import json from unittest.mock import MagicMock, patch +from uuid import UUID import pytest from policyengine_api.constants import ( @@ -35,8 +36,8 @@ ) MOCK_MODAL_JOB_ID = "fc-test123xyz" MOCK_EXECUTION_ID = MOCK_MODAL_JOB_ID # Alias for test compatibility -MOCK_RUN_ID = "run-test123xyz" -MOCK_PROCESS_ID = "job_20250626120000_1234" +MOCK_OBSERVABILITY_ID = "00000000-0000-4000-8000-000000000001" +MOCK_SUBMISSION_CLAIM_ID = "00000000-0000-4000-8000-000000000002" MOCK_MODEL_VERSION = "1.2.3" MOCK_POLICYENGINE_VERSION = "3.4.0" MOCK_RESOLVED_APP_NAME = "policyengine-simulation-us1-2-3-uk2-7-8" @@ -114,6 +115,7 @@ def mock_reform_impacts_service(): mock_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] mock_service.get_all_reform_impacts.return_value = [] mock_service.claim_reform_impact_start.return_value = True + mock_service.get_reform_impact_start_claim.return_value = None mock_service.release_reform_impact_start.return_value = None mock_service.set_reform_impact.return_value = None mock_service.set_complete_reform_impact.return_value = None @@ -159,14 +161,13 @@ def mock_budget_window_cache(): """Mock Redis-backed budget-window cache.""" mock_cache = MagicMock() mock_cache.build_key.return_value = "budget-window-cache-key" - mock_cache.get_terminal_error.return_value = None - mock_cache.get_completed_result.return_value = None - mock_cache.get_batch_job_id.return_value = None + mock_cache.get_state.return_value = None mock_cache.claim_batch_start.return_value = True - mock_cache.store_batch_job_id.return_value = None + mock_cache.store_submitted.return_value = None mock_cache.clear_starting_claim.return_value = None mock_cache.set_completed_result.return_value = True - mock_cache.clear_batch_job_id.return_value = None + mock_cache.set_terminal_error.return_value = True + mock_cache.set_execution_failure.return_value = True with patch( "policyengine_api.services.economy_service.budget_window_cache", @@ -192,11 +193,11 @@ def mock_datetime(): @pytest.fixture -def mock_numpy_random(): - """Mock numpy random integer generation.""" +def mock_submission_claim_id(): + """Return one stable UUID for submission ownership claims.""" with patch( - "policyengine_api.services.economy_service.np.random.randint", - return_value=1234, + "policyengine_api.services.economy_service.uuid.uuid4", + return_value=UUID(MOCK_SUBMISSION_CLAIM_ID), ) as mock: yield mock @@ -223,7 +224,7 @@ def create_mock_reform_impact( }, } ) - return ReformImpact( + impact = ReformImpact( reform_impact_id=1, country_id=MOCK_COUNTRY_ID, reform_policy_id=MOCK_POLICY_ID, @@ -241,6 +242,8 @@ def create_mock_reform_impact( start_time=start_time or datetime.datetime(2025, 6, 26, 12, 0, 0), end_time=(datetime.datetime(2025, 6, 26, 12, 5, 0) if status == "ok" else None), ) + impact.observability_id = MOCK_OBSERVABILITY_ID + return impact def create_mock_modal_execution( @@ -271,7 +274,6 @@ def create_mock_modal_execution( """ mock_execution = MagicMock() mock_execution.job_id = job_id - mock_execution.run_id = MOCK_RUN_ID mock_execution.name = job_id # Alias for compatibility mock_execution.status = status mock_execution.result = result diff --git a/tests/integration/test_runtime_cache_redis.py b/tests/integration/test_runtime_cache_redis.py index acaffa329..913e9193d 100644 --- a/tests/integration/test_runtime_cache_redis.py +++ b/tests/integration/test_runtime_cache_redis.py @@ -173,12 +173,37 @@ def test_reform_submission_claim_is_shared_across_connections(redis_pair) -> Non "options_hash": "resolved-hash", "target": "general", } + writer_observability_id = "00000000-0000-4000-8000-000000000001" + contender_observability_id = "00000000-0000-4000-8000-000000000002" - assert writer.claim_start(**arguments, claim_token="writer") - assert not contender.claim_start(**arguments, claim_token="contender") - assert not contender.release_start(**arguments, claim_token="contender") - assert writer.release_start(**arguments, claim_token="writer") - assert contender.claim_start(**arguments, claim_token="contender") + assert writer.claim_start( + **arguments, + claim_token="writer", + observability_id=writer_observability_id, + ) + assert not contender.claim_start( + **arguments, + claim_token="contender", + observability_id=contender_observability_id, + ) + winning_claim = contender.get_start_claim(**arguments) + assert winning_claim.submission_claim_id == "writer" + assert winning_claim.observability_id == writer_observability_id + assert not contender.release_start( + **arguments, + claim_token="contender", + observability_id=contender_observability_id, + ) + assert writer.release_start( + **arguments, + claim_token="writer", + observability_id=writer_observability_id, + ) + assert contender.claim_start( + **arguments, + claim_token="contender", + observability_id=contender_observability_id, + ) def test_real_reform_indexes_are_cross_connection_bounded_and_expiring( diff --git a/tests/unit/libs/test_simulation_entrypoint.py b/tests/unit/libs/test_simulation_entrypoint.py index ddf51e16b..8e124ec4f 100644 --- a/tests/unit/libs/test_simulation_entrypoint.py +++ b/tests/unit/libs/test_simulation_entrypoint.py @@ -6,18 +6,12 @@ """ import os -import sys -from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import patch import httpx import pytest from flask import Flask, g -sys.modules.setdefault( - "policyengine_api.gcp_logging", - SimpleNamespace(logger=MagicMock()), -) os.environ.setdefault("FLASK_DEBUG", "1") from policyengine_api.constants import ( # noqa: E402 @@ -35,7 +29,13 @@ ) from policyengine_api.request_context import ( # noqa: E402 REQUEST_ID_HEADER, + _asgi_observability_id, _asgi_request_id, + current_observability_id, + start_observability_id, +) +from policyengine_api.observability.identifiers import ( # noqa: E402 + OBSERVABILITY_ID_HEADER, ) from tests.fixtures.libs.simulation_entrypoint import ( # noqa: E402 @@ -52,7 +52,7 @@ MOCK_POLL_RESPONSE_FAILED, MOCK_POLL_RESPONSE_RUNNING, MOCK_RESOLVED_APP_NAME, - MOCK_RUN_ID, + MOCK_OBSERVABILITY_ID, MOCK_SIMULATION_PAYLOAD, MOCK_SIMULATION_PAYLOAD_WITH_TELEMETRY, MOCK_SIMULATION_RESULT, @@ -117,7 +117,15 @@ def _response(self, method, url, json=None): payload = MOCK_HEALTH_RESPONSE status_code = 200 - return httpx.Response(status_code, request=request, json=payload) + response = httpx.Response( + status_code, + request=request, + json=payload, + headers={"X-PolicyEngine-Observability-Id": MOCK_OBSERVABILITY_ID}, + ) + for hook in self.event_hooks.get("response", []): + hook(response) + return response def post(self, url, json=None): return self._response("POST", url, json=json) @@ -366,7 +374,7 @@ def test__given_partial_gateway_auth_env_vars__then_raises( with pytest.raises(GatewayAuthError): SimulationAPIModal() - def test__given_client_initialized__then_installs_one_request_id_hook( + def test__given_client_initialized__then_installs_correlation_hooks( self, mock_httpx_client ): from policyengine_api.libs.simulation_entrypoint import httpx as modal_httpx @@ -377,6 +385,20 @@ def test__given_client_initialized__then_installs_one_request_id_hook( assert list(kwargs["event_hooks"]) == ["request"] assert len(kwargs["event_hooks"]["request"]) == 1 + def test__given_client_initialized__then_instruments_explicit_httpx_client( + self, mock_httpx_client + ): + from policyengine_api.libs import simulation_entrypoint as module + + runtime = object() + with ( + patch.object(module, "get_runtime", return_value=runtime), + patch.object(module, "instrument_httpx") as instrument, + ): + client = SimulationAPIModal() + + instrument.assert_called_once_with(client.client, runtime) + def test__given_flask_request__then_hook_uses_current_request_id( self, mock_httpx_client ): @@ -397,7 +419,7 @@ def test__given_flask_request__then_hook_uses_current_request_id( assert request.headers[REQUEST_ID_HEADER] == "flask-request-id" - def test__given_asgi_request__then_hook_uses_current_request_id( + def test__given_asgi_request__then_hook_uses_current_correlation_ids( self, mock_httpx_client ): from policyengine_api.libs.simulation_entrypoint import httpx as modal_httpx @@ -406,12 +428,15 @@ def test__given_asgi_request__then_hook_uses_current_request_id( hook = modal_httpx.Client.call_args.kwargs["event_hooks"]["request"][0] request = httpx.Request("GET", MOCK_MODAL_BASE_URL) token = _asgi_request_id.set("asgi-request-id") + observability_token = _asgi_observability_id.set(MOCK_OBSERVABILITY_ID) try: hook(request) finally: _asgi_request_id.reset(token) + _asgi_observability_id.reset(observability_token) assert request.headers[REQUEST_ID_HEADER] == "asgi-request-id" + assert request.headers[OBSERVABILITY_ID_HEADER] == MOCK_OBSERVABILITY_ID def test__given_no_request_context__then_hook_omits_request_id( self, monkeypatch, mock_modal_logger @@ -483,6 +508,38 @@ def test__given_request_context__then_all_calls_forward_request_id( for request in requests ) + def test__given_started_calculation__then_client_preserves_one_observability_id( + self, + monkeypatch, + mock_modal_logger, + ): + from policyengine_api.libs import simulation_entrypoint as module + + RequestRecordingHTTPXClient.instances.clear() + monkeypatch.setattr( + module.httpx, + "Client", + RequestRecordingHTTPXClient, + ) + api = SimulationAPIModal() + app = Flask("observability-id-client-lifecycle") + api_observability_id = "00000000-0000-4000-8000-000000000099" + + with app.test_request_context(): + g.request_id = "flask-request-id" + g.incoming_observability_id = api_observability_id + g.observability_id = None + + selected_id = start_observability_id() + api.run(MOCK_SIMULATION_PAYLOAD) + + assert current_observability_id() == selected_id + + request = RequestRecordingHTTPXClient.instances[-1].requests[-1] + assert selected_id == api_observability_id + assert selected_id != MOCK_OBSERVABILITY_ID + assert request.headers[OBSERVABILITY_ID_HEADER] == selected_id + @pytest.mark.parametrize("method", ["run", "run_budget_window_batch"]) @pytest.mark.parametrize( "override", @@ -520,7 +577,7 @@ def test__given_valid_payload__then_returns_execution_with_job_id( # Then assert execution.job_id == MOCK_MODAL_JOB_ID - assert execution.run_id == MOCK_RUN_ID + assert not hasattr(execution, "observability_id") assert execution.status == MODAL_EXECUTION_STATUS_SUBMITTED assert execution.policyengine_bundle == MOCK_POLICYENGINE_BUNDLE assert execution.resolved_app_name == MOCK_RESOLVED_APP_NAME @@ -546,7 +603,7 @@ def test__given_valid_payload__then_posts_to_correct_endpoint( assert "/simulate/economy/comparison" in call_args[0][0] assert call_args[1]["json"] == MOCK_SIMULATION_PAYLOAD - def test__given_telemetry_payload__then_preserves_it_in_post_body( + def test__given_telemetry_payload__then_preserves_non_identity_fields( self, mock_httpx_client, mock_modal_logger, @@ -560,7 +617,10 @@ def test__given_telemetry_payload__then_preserves_it_in_post_body( api.run(MOCK_SIMULATION_PAYLOAD_WITH_TELEMETRY) call_args = mock_httpx_client.post.call_args - assert call_args[1]["json"]["_telemetry"]["run_id"] == MOCK_RUN_ID + assert call_args[1]["json"]["_telemetry"] == { + "submission_claim_id": "job_20250626120000_1234", + "capture_mode": "disabled", + } def test__given_model_and_bundle_versions__then_translates_payload_for_modal( self, @@ -606,7 +666,7 @@ def test__given_api_v1_default_bundle_payload__then_posts_gateway_contract_body( "model_version": "1.729.0", "policyengine_version": "4.18.3", "_metadata": { - "process_id": "job_20260629120000_1234", + "submission_claim_id": "job_20260629120000_1234", "model_version": "1.729.0", "policyengine_version": "4.18.3", "data_version": None, @@ -614,8 +674,7 @@ def test__given_api_v1_default_bundle_payload__then_posts_gateway_contract_body( "resolved_app_name": "policyengine-simulation-py4-18-3", }, "_telemetry": { - "run_id": "run_20260629120000_1234", - "process_id": "job_20260629120000_1234", + "submission_claim_id": "job_20260629120000_1234", "capture_mode": "disabled", }, } @@ -673,12 +732,18 @@ def test__given_network_error__then_raises_exception( api = SimulationAPIModal() # When/Then - with pytest.raises(httpx.RequestError): + with ( + patch( + "policyengine_api.libs.simulation_entrypoint.current_observability_id", + return_value=MOCK_OBSERVABILITY_ID, + ), + pytest.raises(httpx.RequestError), + ): api.run(MOCK_SIMULATION_PAYLOAD_WITH_TELEMETRY) log_payload = mock_modal_logger.log_struct.call_args.args[0] assert "Simulation entrypoint request error" in log_payload["message"] - assert log_payload["run_id"] == MOCK_RUN_ID + assert log_payload["observability_id"] == MOCK_OBSERVABILITY_ID class TestResolveAppName: def test__given_country_and_version__then_returns_registered_app( @@ -811,11 +876,17 @@ def test__given_network_error__then_raises_exception( mock_httpx_client.post.side_effect = httpx.RequestError("Connection failed") api = SimulationAPIModal() - with pytest.raises(httpx.RequestError): + with ( + patch( + "policyengine_api.libs.simulation_entrypoint.current_observability_id", + return_value=MOCK_OBSERVABILITY_ID, + ), + pytest.raises(httpx.RequestError), + ): api.run_budget_window_batch(MOCK_SIMULATION_PAYLOAD_WITH_TELEMETRY) log_payload = mock_modal_logger.log_struct.call_args.args[0] - assert log_payload["run_id"] == MOCK_RUN_ID + assert log_payload["observability_id"] == MOCK_OBSERVABILITY_ID class TestGetExecutionById: def test__given_running_job__then_returns_running_status( diff --git a/tests/unit/routes/test_calculate_deprecated_inputs.py b/tests/unit/routes/test_calculate_deprecated_inputs.py index aeb0af3e8..67ebe09f5 100644 --- a/tests/unit/routes/test_calculate_deprecated_inputs.py +++ b/tests/unit/routes/test_calculate_deprecated_inputs.py @@ -1,5 +1,6 @@ from flask import Flask import pytest +from unittest.mock import patch from policyengine_api.routes import household_routes from policyengine_api.extensions import cache @@ -141,12 +142,17 @@ def test__calculate__omits_warnings_without_deprecated_input(calculate_client): } } - response = client.post("/us/calculate", json={"household": household}) + with patch.object( + household_routes, + "start_observability_id", + ) as start_observability_id: + response = client.post("/us/calculate", json={"household": household}) assert response.status_code == 200 payload = response.get_json() assert "warnings" not in payload assert country.household == household + start_observability_id.assert_called_once_with() def test__calculate__returns_400_for_unrecognized_household_variable( @@ -162,7 +168,11 @@ def test__calculate__returns_400_for_unrecognized_household_variable( } } - response = client.post("/us/calculate", json={"household": household}) + with patch.object( + household_routes, + "start_observability_id", + ) as start_observability_id: + response = client.post("/us/calculate", json={"household": household}) assert response.status_code == 400 payload = response.get_json() @@ -181,6 +191,7 @@ def test__calculate__returns_400_for_unrecognized_household_variable( } ] assert country.household is None + start_observability_id.assert_not_called() def test__calculate__returns_400_for_variable_on_wrong_entity(calculate_client): diff --git a/tests/unit/routes/test_calculate_error_statuses.py b/tests/unit/routes/test_calculate_error_statuses.py index 6f8f6cd19..206d4dddd 100644 --- a/tests/unit/routes/test_calculate_error_statuses.py +++ b/tests/unit/routes/test_calculate_error_statuses.py @@ -1,5 +1,6 @@ from flask import Flask from policyengine_core.errors import SituationParsingError +from unittest.mock import patch from policyengine_api.extensions import cache from policyengine_api.routes import household_routes @@ -20,7 +21,7 @@ class ParsingErrorCountry(DummyCountry): - def calculate(self, household, policy): + def prepare_calculation(self, household, policy): raise SituationParsingError( ["people", "you", "employment_income", "2026"], "Can't deal with value: expected type number, received '{}'.", @@ -50,28 +51,40 @@ def make_client(monkeypatch, country, add_missing=False): return app.test_client() -def test__calculate__returns_400_on_situation_parsing_error(monkeypatch): +def test__calculate__returns_400_without_accepting_situation_parsing_error(monkeypatch): client = make_client(monkeypatch, ParsingErrorCountry()) - response = client.post("/us/calculate", json={"household": HOUSEHOLD}) + with patch.object( + household_routes, + "start_observability_id", + ) as start_observability_id: + response = client.post("/us/calculate", json={"household": HOUSEHOLD}) assert response.status_code == 400 payload = response.get_json() assert payload["status"] == "error" assert payload["result"] is None assert payload["message"].startswith("Invalid household payload") + start_observability_id.assert_not_called() -def test__calculate_full__returns_400_on_situation_parsing_error(monkeypatch): +def test__calculate_full__returns_400_without_accepting_situation_parsing_error( + monkeypatch, +): client = make_client(monkeypatch, ParsingErrorCountry(), add_missing=True) - response = client.post("/us/calculate-full", json={"household": HOUSEHOLD}) + with patch.object( + household_routes, + "start_observability_id", + ) as start_observability_id: + response = client.post("/us/calculate-full", json={"household": HOUSEHOLD}) assert response.status_code == 400 payload = response.get_json() assert payload["status"] == "error" assert payload["result"] is None assert payload["message"].startswith("Invalid household payload") + start_observability_id.assert_not_called() def test__calculate__returns_500_on_unexpected_error(monkeypatch): diff --git a/tests/unit/routes/test_canonical_spm.py b/tests/unit/routes/test_canonical_spm.py index e3c55a015..2dd51f9d3 100644 --- a/tests/unit/routes/test_canonical_spm.py +++ b/tests/unit/routes/test_canonical_spm.py @@ -2,6 +2,7 @@ from copy import deepcopy from types import SimpleNamespace +from unittest.mock import patch from flask import Flask import pytest @@ -282,6 +283,24 @@ def test_http_cache_varies_with_measurement_settings(certified, harness, selecti assert first.json["spm_config"] != second.json["spm_config"] +def test_successful_http_cache_hit_still_starts_observability(certified, harness): + client, country = harness + payload = {"household": HOUSEHOLD} + + with patch.object( + household_routes, + "start_observability_id", + ) as start_observability_id: + first = client.post("/us/calculate", json=payload) + start_observability_id.reset_mock() + cached = client.post("/us/calculate", json=payload) + + assert first.status_code == 200 + assert cached.status_code == 200 + assert len(country.calls) == 1 + start_observability_id.assert_called_once_with() + + def test_certification_checked_before_cached_response(certified, harness): client, country = harness payload = {"household": HOUSEHOLD, "spm": {"geography_kind": "national"}} diff --git a/tests/unit/routes/test_migration_context_logging.py b/tests/unit/routes/test_migration_context_logging.py index 89dff5b14..0d685954d 100644 --- a/tests/unit/routes/test_migration_context_logging.py +++ b/tests/unit/routes/test_migration_context_logging.py @@ -1,5 +1,5 @@ from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import Mock, patch from fastapi.testclient import TestClient from flask import Flask, Response @@ -12,8 +12,16 @@ from policyengine_api.migration_logging import log_migration_request 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" def _app(): @@ -28,6 +36,15 @@ def readiness_check(): def request_id(): return Response(current_request_id(), status=200, mimetype="text/plain") + @app.route("/calculation") + def calculation(): + start_observability_id() + return Response( + current_observability_id(), + status=200, + mimetype="text/plain", + ) + register_migration_request_logging(app) return app @@ -94,6 +111,126 @@ def test_request_logging_includes_migration_context(): assert log_payload["migration"]["route_impl"] == "flask_fallback" +def test_instrumented_flask_request_enriches_single_adapter_record(): + app = Flask(__name__) + app.config["TESTING"] = True + runtime = Mock() + runtime.capture_context.return_value = {"request_id": "request-123"} + + @app.route("//metadata") + def metadata(country_id): + return Response(country_id, status=200, mimetype="text/plain") + + register_migration_request_logging(app, runtime=runtime) + + with patch("policyengine_api.migration_logging.logger") as mock_logger: + response = app.test_client().get("/us/metadata") + + assert response.status_code == 200 + assert response.headers[REQUEST_ID_HEADER] == "request-123" + 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, + ) + mock_logger.log_struct.assert_not_called() + + +def test_observability_runtime_failure_does_not_reject_flask_request(): + app = Flask(__name__) + runtime = Mock() + runtime.capture_context.side_effect = RuntimeError("runtime unavailable") + runtime.set_context.side_effect = RuntimeError("runtime unavailable") + register_migration_request_logging(app, runtime=runtime) + + @app.get("/health") + def health(): + return {"status": "ok"} + + response = app.test_client().get("/health") + + assert response.status_code == 200 + assert response.json == {"status": "ok"} + assert response.headers[REQUEST_ID_HEADER] + assert "X-PolicyEngine-Observability-Id" not in response.headers + + +def test_flask_reapplies_calculation_identifier_to_server_request_span(): + app = Flask(__name__) + runtime = Mock() + runtime.capture_context.return_value = {"request_id": "request-123"} + + @app.get("//calculation") + def calculation(country_id): + start_observability_id() + return {"country_id": country_id} + + register_migration_request_logging(app, runtime=runtime) + + with patch( + "policyengine_api.observability.get_runtime", + return_value=runtime, + ): + response = app.test_client().get( + "/us/calculation", + headers={OBSERVABILITY_ID_HEADER: OBSERVABILITY_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, + } + + +def test_flask_binds_incoming_observability_id_only_when_calculation_starts(): + response = ( + _app() + .test_client() + .get( + "/calculation", + headers={OBSERVABILITY_ID_HEADER: OBSERVABILITY_ID}, + ) + ) + + assert response.status_code == 200 + assert response.text == OBSERVABILITY_ID + assert response.headers[OBSERVABILITY_ID_HEADER] == OBSERVABILITY_ID + + +def test_flask_does_not_echo_incoming_observability_id_on_non_calculation_route(): + response = ( + _app() + .test_client() + .get( + "/request-id", + headers={OBSERVABILITY_ID_HEADER: OBSERVABILITY_ID}, + ) + ) + + assert response.status_code == 200 + assert OBSERVABILITY_ID_HEADER not in response.headers + + def test_flask_preserves_policyengine_request_id_in_context_log_and_response(): with patch("policyengine_api.migration_logging.logger") as mock_logger: response = ( diff --git a/tests/unit/routes/test_spm_year_worker_polling.py b/tests/unit/routes/test_spm_year_worker_polling.py index 797a57097..c79f023f4 100644 --- a/tests/unit/routes/test_spm_year_worker_polling.py +++ b/tests/unit/routes/test_spm_year_worker_polling.py @@ -156,7 +156,11 @@ def transport(request): ) if budget_window: cache_key = service._build_budget_window_cache_key(setup) - window_cache.store_batch_job_id(cache_key, job_id) + window_cache.store_submitted( + cache_key, + job_id, + setup.observability_id, + ) url = "/us/economy/1/over/2/budget-window" query = {"start_year": "2035", "window_size": "2"} else: @@ -193,9 +197,12 @@ def transport(request): "errors": [typed_error], } if budget_window: - assert window_cache.get_completed_result(cache_key) is None - assert window_cache.get_batch_job_id(cache_key) is None - assert window_cache.get_terminal_error(cache_key) == typed_error + state = window_cache.get_state(cache_key) + assert state is not None + assert state.status == "failed" + assert state.failure_type == "spm_validation" + assert state.error == typed_error + assert state.observability_id == setup.observability_id else: stored = annual_cache.get_by_execution_id(job_id) assert stored.status == "error" @@ -318,6 +325,7 @@ def transport(request): segmented_result("2036", "reform"), ], }, + setup.observability_id, ) url = "/us/economy/1/over/2/budget-window" query = {"start_year": "2035", "window_size": "2"} @@ -354,8 +362,12 @@ def transport(request): assert response.json == first if budget_window: - terminal = window_cache.get_terminal_error(cache_key) - assert terminal["code"] == "SPM_CONFIGURATION_UNAVAILABLE" + state = window_cache.get_state(cache_key) + assert state is not None + assert state.status == "failed" + assert state.failure_type == "spm_validation" + assert state.error is not None + assert state.error["code"] == "SPM_CONFIGURATION_UNAVAILABLE" else: stored = annual_cache.get_by_execution_id(job_id) assert stored.status == "error" diff --git a/tests/unit/runtime_cache/test_reform_impacts.py b/tests/unit/runtime_cache/test_reform_impacts.py index 98b134b4f..43445b86a 100644 --- a/tests/unit/runtime_cache/test_reform_impacts.py +++ b/tests/unit/runtime_cache/test_reform_impacts.py @@ -4,7 +4,7 @@ import pytest -from policyengine_api.runtime_cache.core import CacheNamespace +from policyengine_api.runtime_cache.core import CacheCoordinationError, CacheNamespace from policyengine_api.runtime_cache.fake import InMemoryCacheBackend from policyengine_api.runtime_cache.reform_impacts import ( REFORM_IMPACT_START_CLAIM_TTL_SECONDS, @@ -59,21 +59,82 @@ def test_reform_impact_start_claim_is_atomic_exact_ttl_and_token_safe() -> None: backend = InMemoryCacheBackend() cache = ReformImpactCache(backend, _namespace()) arguments = _claim_arguments() + owner_observability_id = "00000000-0000-4000-8000-000000000001" + contender_observability_id = "00000000-0000-4000-8000-000000000002" - assert cache.claim_start(**arguments, claim_token="owner") is True - assert cache.claim_start(**arguments, claim_token="contender") is False + assert ( + cache.claim_start( + **arguments, + claim_token="owner", + observability_id=owner_observability_id, + ) + is True + ) + assert ( + cache.claim_start( + **arguments, + claim_token="contender", + observability_id=contender_observability_id, + ) + is False + ) assert ( cache.claim_start( **_claim_arguments(target="cliff"), claim_token="cliff-owner", + observability_id=owner_observability_id, ) is True ) assert set(backend._expires.values()) == {REFORM_IMPACT_START_CLAIM_TTL_SECONDS} - assert cache.release_start(**arguments, claim_token="contender") is False - assert cache.release_start(**arguments, claim_token="owner") is True - assert cache.claim_start(**arguments, claim_token="next-owner") is True + claim = cache.get_start_claim(**arguments) + assert claim is not None + assert claim.submission_claim_id == "owner" + assert claim.observability_id == owner_observability_id + assert ( + cache.release_start( + **arguments, + claim_token="contender", + observability_id=contender_observability_id, + ) + is False + ) + assert ( + cache.release_start( + **arguments, + claim_token="owner", + observability_id=contender_observability_id, + ) + is False + ) + assert ( + cache.release_start( + **arguments, + claim_token="owner", + observability_id=owner_observability_id, + ) + is True + ) + assert ( + cache.claim_start( + **arguments, + claim_token="next-owner", + observability_id=contender_observability_id, + ) + is True + ) + + +def test_reform_impact_start_claim_fails_closed_when_state_is_unreadable() -> None: + backend = InMemoryCacheBackend() + cache = ReformImpactCache(backend, _namespace()) + arguments = _claim_arguments() + key = cache._start_claim_key(**arguments) + backend.set(key, "legacy-or-corrupt-claim", ex=300) + + with pytest.raises(CacheCoordinationError, match="ownership is unreadable"): + cache.get_start_claim(**arguments) def test_reform_impact_indexes_are_bounded_expiring_and_query_compatible( diff --git a/tests/unit/services/test_budget_window_cache.py b/tests/unit/services/test_budget_window_cache.py index 0bb51ee9d..1c9009c5e 100644 --- a/tests/unit/services/test_budget_window_cache.py +++ b/tests/unit/services/test_budget_window_cache.py @@ -8,6 +8,7 @@ BUDGET_WINDOW_BATCH_TTL_SECONDS, BUDGET_WINDOW_STARTING_TTL_SECONDS, BudgetWindowCache, + BudgetWindowCacheState, ) @@ -39,6 +40,21 @@ def eval(self, *_args, **_kwargs): raise RuntimeError("redis unavailable") +@pytest.mark.parametrize( + "payload", + [ + {"status": "starting", "submission_claim_id": ""}, + {"status": "submitted", "batch_job_id": 1}, + {"status": "completed"}, + {"status": "completed", "result": []}, + {"status": "failed", "failure_type": "execution"}, + {"status": "failed", "failure_type": "unknown", "error": {}}, + ], +) +def test_cache_state_rejects_malformed_documents(payload): + assert BudgetWindowCacheState.from_payload(payload) is None + + def test_build_key_is_stable_for_request_identity(): cache = BudgetWindowCache(client=FakeRedis()) @@ -64,38 +80,51 @@ def test_build_key_is_stable_for_request_identity(): ) assert first == second - assert first.startswith("policyengine:test:api:budget-window:v1:") + assert first.startswith("policyengine:test:api:budget-window:v2:") -def test_claim_batch_start_allows_one_starter(): +def test_claim_batch_start_allows_one_starter_and_preserves_identity(): cache = BudgetWindowCache(client=FakeRedis()) + cache_key = "budget_window:v2:us:key" + + assert cache.claim_batch_start(cache_key, "claim-1", "obs-1") is True + assert cache.claim_batch_start(cache_key, "claim-2", "obs-2") is False - assert cache.claim_batch_start("budget_window:v1:us:key", "process-1") is True - assert cache.claim_batch_start("budget_window:v1:us:key", "process-2") is False - assert cache.get_batch_job_id("budget_window:v1:us:key") is None + state = cache.get_state(cache_key) + assert state is not None + assert state.status == "starting" + assert state.submission_claim_id == "claim-1" + assert state.observability_id == "obs-1" -def test_store_batch_job_id_replaces_starting_claim(): +def test_store_submitted_replaces_starting_state(): cache = BudgetWindowCache(client=FakeRedis()) - cache.claim_batch_start("budget_window:v1:us:key", "process-1") + cache_key = "budget_window:v2:us:key" + cache.claim_batch_start(cache_key, "claim-1", "obs-1") - cache.store_batch_job_id("budget_window:v1:us:key", "fc-parent") + cache.store_submitted(cache_key, "fc-parent", "obs-1") - assert cache.get_batch_job_id("budget_window:v1:us:key") == "fc-parent" + state = cache.get_state(cache_key) + assert state is not None + assert state.status == "submitted" + assert state.batch_job_id == "fc-parent" + assert state.observability_id == "obs-1" -def test_completed_result_round_trips(): +def test_completed_result_round_trips_with_identity(): cache = BudgetWindowCache(client=FakeRedis()) result = {"kind": "budgetWindow", "totals": {"budgetaryImpact": 10}} - cache.set_completed_result("budget_window:v1:us:key", result) + cache.set_completed_result("budget_window:v2:us:key", result, "obs-1") - assert cache.get_completed_result("budget_window:v1:us:key") == result + state = cache.get_state("budget_window:v2:us:key") + assert state is not None + assert state.status == "completed" + assert state.result == result + assert state.observability_id == "obs-1" -def test_terminal_error_round_trips_separately_from_success_and_other_selections( - monkeypatch, -): +def test_spm_validation_failure_round_trips_for_only_its_selection(monkeypatch): import policyengine_api.services.budget_window_cache as module monkeypatch.setattr(module, "jittered_ttl", lambda _ttl: 123) @@ -113,17 +142,36 @@ def test_terminal_error_round_trips_separately_from_success_and_other_selections } failed_key = cache.build_key(**identity, options_hash="canonical-selection-a") other_key = cache.build_key(**identity, options_hash="canonical-selection-b") - assert cache.set_terminal_error(failed_key, error) - cache = BudgetWindowCache(client=backend) - assert cache.get_terminal_error(failed_key) == error - assert cache.get_terminal_error(other_key) is None - assert cache.get_completed_result(failed_key) is None + + assert cache.set_terminal_error(failed_key, error, "obs-1") + + stored = BudgetWindowCache(client=backend).get_state(failed_key) + assert stored is not None + assert stored.status == "failed" + assert stored.failure_type == "spm_validation" + assert stored.error == error + assert stored.observability_id == "obs-1" + assert cache.get_state(other_key) is None assert set(backend._expires.values()) == {123} backend.advance(123) - assert cache.get_terminal_error(failed_key) is None + assert cache.get_state(failed_key) is None + + +def test_execution_failure_round_trips_with_identity(): + cache = BudgetWindowCache(client=FakeRedis()) + result = {"status": "error", "error": "simulation failed"} + + assert cache.set_execution_failure("budget_window:v2:us:key", result, "obs-1") + + state = cache.get_state("budget_window:v2:us:key") + assert state is not None + assert state.status == "failed" + assert state.failure_type == "execution" + assert state.error == result + assert state.observability_id == "obs-1" -def test_completed_result_ttl_is_jittered_but_coordination_ttls_are_exact( +def test_recoverable_state_ttl_is_jittered_but_coordination_ttls_are_exact( monkeypatch, ): import policyengine_api.services.budget_window_cache as module @@ -131,166 +179,103 @@ def test_completed_result_ttl_is_jittered_but_coordination_ttls_are_exact( monkeypatch.setattr(module, "jittered_ttl", lambda _ttl: 123) redis_client = FakeRedis() cache = BudgetWindowCache(client=redis_client) - cache_key = "budget_window:v1:us:key" + cache_key = "budget_window:v2:us:key" + state_key = f"{cache_key}:state" - assert cache.set_completed_result(cache_key, {"ok": True}) - assert redis_client._expires[f"{cache_key}:result"] == 123 + assert cache.set_completed_result(cache_key, {"ok": True}, "obs-1") + assert redis_client._expires[state_key] == 123 - assert cache.claim_batch_start(cache_key, "process-1") - assert ( - redis_client._expires[f"{cache_key}:batch-job-id"] - == BUDGET_WINDOW_STARTING_TTL_SECONDS - ) - - cache.store_batch_job_id(cache_key, "batch-1") - assert ( - redis_client._expires[f"{cache_key}:batch-job-id"] - == BUDGET_WINDOW_BATCH_TTL_SECONDS - ) - - -def test_get_completed_result_returns_none_for_empty_payload(): - redis_client = FakeRedis() - redis_client.values["budget_window:v1:us:key:result"] = "" - cache = BudgetWindowCache(client=redis_client) + redis_client.delete(state_key) + assert cache.claim_batch_start(cache_key, "claim-1", "obs-1") + assert redis_client._expires[state_key] == BUDGET_WINDOW_STARTING_TTL_SECONDS - assert cache.get_completed_result("budget_window:v1:us:key") is None + cache.store_submitted(cache_key, "batch-1", "obs-1") + assert redis_client._expires[state_key] == BUDGET_WINDOW_BATCH_TTL_SECONDS -def test_get_completed_result_returns_none_for_invalid_json(monkeypatch): +@pytest.mark.parametrize("invalid_value", ["", "{not-json", "123"]) +def test_get_state_removes_invalid_payload(invalid_value, monkeypatch): mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) + monkeypatch.setattr("policyengine_api.runtime_cache.core.logger", mock_logger) redis_client = FakeRedis() - redis_client.values["budget_window:v1:us:key:result"] = "{not-json" + state_key = "budget_window:v2:us:key:state" + redis_client.values[state_key] = invalid_value cache = BudgetWindowCache(client=redis_client) - assert cache.get_completed_result("budget_window:v1:us:key") is None - assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" + assert cache.get_state("budget_window:v2:us:key") is None + assert state_key not in redis_client.values + assert any( + call.kwargs.get("severity") == "WARNING" + for call in mock_logger.log_struct.call_args_list + ) -def test_get_completed_result_treats_read_errors_as_misses(monkeypatch): +def test_get_state_reraises_read_errors(monkeypatch): mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) + monkeypatch.setattr("policyengine_api.runtime_cache.core.logger", mock_logger) cache = BudgetWindowCache(client=RaisingRedis(method="get")) - assert cache.get_completed_result("budget_window:v1:us:key") is None + with pytest.raises(CacheCoordinationError): + cache.get_state("budget_window:v2:us:key") assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" -def test_set_completed_result_does_not_invalidate_compute_on_write_error(monkeypatch): +def test_completed_result_write_error_does_not_change_returned_result(monkeypatch): mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) + monkeypatch.setattr("policyengine_api.runtime_cache.core.logger", mock_logger) cache = BudgetWindowCache(client=RaisingRedis(method="set")) - assert not cache.set_completed_result("budget_window:v1:us:key", {"ok": True}) - - assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" - - -def test_get_batch_job_id_ignores_empty_non_string_and_starting_values(): - redis_client = FakeRedis() - cache = BudgetWindowCache(client=redis_client) - - redis_client.values["budget_window:v1:us:key:batch-job-id"] = "" - assert cache.get_batch_job_id("budget_window:v1:us:key") is None - - redis_client.values["budget_window:v1:us:key:batch-job-id"] = 123 - assert cache.get_batch_job_id("budget_window:v1:us:key") is None - - redis_client.values["budget_window:v1:us:key:batch-job-id"] = "starting:process-1" - assert cache.get_batch_job_id("budget_window:v1:us:key") is None - - -def test_get_batch_job_id_reraises_read_errors(monkeypatch): - mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, + assert not cache.set_completed_result( + "budget_window:v2:us:key", {"ok": True}, "obs-1" ) - cache = BudgetWindowCache(client=RaisingRedis(method="get")) - - with pytest.raises(CacheCoordinationError): - cache.get_batch_job_id("budget_window:v1:us:key") - assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" def test_claim_batch_start_reraises_claim_errors(monkeypatch): mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) + monkeypatch.setattr("policyengine_api.runtime_cache.core.logger", mock_logger) cache = BudgetWindowCache(client=RaisingRedis(method="set")) with pytest.raises(CacheCoordinationError): - cache.claim_batch_start("budget_window:v1:us:key", "process-1") + cache.claim_batch_start("budget_window:v2:us:key", "claim-1", "obs-1") assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" -def test_store_batch_job_id_reraises_write_errors(monkeypatch): +def test_store_submitted_reraises_write_errors(monkeypatch): mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) + monkeypatch.setattr("policyengine_api.runtime_cache.core.logger", mock_logger) cache = BudgetWindowCache(client=RaisingRedis(method="set")) with pytest.raises(CacheCoordinationError): - cache.store_batch_job_id("budget_window:v1:us:key", "fc-parent") + cache.store_submitted("budget_window:v2:us:key", "fc-parent", "obs-1") assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" -def test_clear_starting_claim_deletes_only_matching_token(): +def test_clear_starting_claim_deletes_only_matching_document(): redis_client = FakeRedis() cache = BudgetWindowCache(client=redis_client) - cache.claim_batch_start("budget_window:v1:us:key", "process-1") - - cache.clear_starting_claim("budget_window:v1:us:key", "process-2") - - assert ( - redis_client.values["budget_window:v1:us:key:batch-job-id"] - == "starting:process-1" - ) - - cache.clear_starting_claim("budget_window:v1:us:key", "process-1") + cache_key = "budget_window:v2:us:key" + state_key = f"{cache_key}:state" + cache.claim_batch_start(cache_key, "claim-1", "obs-1") - assert "budget_window:v1:us:key:batch-job-id" not in redis_client.values + cache.clear_starting_claim(cache_key, "claim-2", "obs-1") + assert state_key in redis_client.values + cache.clear_starting_claim(cache_key, "claim-1", "different-observability-id") + assert state_key in redis_client.values -def test_clear_starting_claim_logs_and_swallows_errors(monkeypatch): - mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) - cache = BudgetWindowCache(client=RaisingRedis(method="get")) + cache.clear_starting_claim(cache_key, "claim-1", "obs-1") + assert state_key not in redis_client.values - cache.clear_starting_claim("budget_window:v1:us:key", "process-1") - assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" - - -def test_clear_batch_job_id_logs_and_swallows_errors(monkeypatch): +def test_clear_starting_claim_swallows_coordination_errors(monkeypatch): mock_logger = MagicMock() - monkeypatch.setattr( - "policyengine_api.runtime_cache.core.logger", - mock_logger, - ) - cache = BudgetWindowCache(client=RaisingRedis(method="delete")) + monkeypatch.setattr("policyengine_api.runtime_cache.core.logger", mock_logger) + cache = BudgetWindowCache(client=RaisingRedis(method="eval")) - cache.clear_batch_job_id("budget_window:v1:us:key") + cache.clear_starting_claim("budget_window:v2:us:key", "claim-1", "obs-1") assert mock_logger.log_struct.call_args.kwargs["severity"] == "WARNING" diff --git a/tests/unit/services/test_economy_service.py b/tests/unit/services/test_economy_service.py index ab0fc7353..60f1a42e1 100644 --- a/tests/unit/services/test_economy_service.py +++ b/tests/unit/services/test_economy_service.py @@ -1,10 +1,12 @@ import json from typing import Literal -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import httpx import pytest from policyengine_api.runtime_cache.core import CacheCoordinationError +from policyengine_api.runtime_cache.reform_impacts import ReformImpactStartClaim +from policyengine_api.services.budget_window_cache import BudgetWindowCacheState from policyengine_api.services.reform_impacts_service import ( ReformImpactHandoffError, ) @@ -16,6 +18,7 @@ EconomyService, ImpactAction, ImpactStatus, + REFORM_IMPACT_START_CLAIM_ATTEMPTS, ) from policyengine_api.services.policy_service import PolicyService from policyengine_api.spm import SPMValidationError @@ -32,12 +35,12 @@ MOCK_OPTIONS_HASH, MOCK_POLICY_ID, MOCK_POLICYENGINE_VERSION, - MOCK_PROCESS_ID, + MOCK_SUBMISSION_CLAIM_ID, MOCK_REFORM_IMPACT_DATA, MOCK_REGION, MOCK_RESOLVED_APP_NAME, MOCK_RESOLVED_DATASET, - MOCK_RUN_ID, + MOCK_OBSERVABILITY_ID, MOCK_TIME_PERIOD, create_mock_budget_window_batch_execution, create_mock_reform_impact, @@ -46,6 +49,25 @@ pytest_plugins = ("tests.fixtures.services.economy_service",) +@pytest.fixture(autouse=True) +def stable_observability_lifecycle(): + with ( + patch( + "policyengine_api.services.economy_service.start_observability_id", + return_value=MOCK_OBSERVABILITY_ID, + ), + patch( + "policyengine_api.services.economy_service.restore_observability_id", + side_effect=lambda value: value, + ), + patch( + "policyengine_api.services.economy_service.resolve_observability_id", + return_value=MOCK_OBSERVABILITY_ID, + ), + ): + yield + + def make_mock_budget_impact_data( *, tax_revenue_impact: int, @@ -115,7 +137,7 @@ def test__given_completed_impact__returns_completed_result( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact(status="ok") mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ @@ -151,7 +173,7 @@ def test__given_orm_decoded_completed_impact__returns_completed_result( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact(status="ok") completed_impact.reform_impact_json = json.loads( @@ -193,7 +215,7 @@ def test__given_cached_error_impact__returns_error_message( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): failed_impact = create_mock_reform_impact( status="error", @@ -222,7 +244,7 @@ def test__given_legacy_completed_impact__refreshes_cache( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact( status="ok", @@ -254,7 +276,7 @@ def test__given_computing_impact_with_succeeded_execution__returns_completed_res mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): computing_impact = create_mock_reform_impact(status="computing") mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ @@ -293,7 +315,7 @@ def test__given_computing_impact_with_failed_execution__returns_error_result( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): computing_impact = create_mock_reform_impact(status="computing") mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ @@ -325,7 +347,7 @@ def test__given_computing_impact_with_active_execution__returns_computing_result mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): computing_impact = create_mock_reform_impact(status="computing") mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ @@ -349,7 +371,7 @@ def test__given_no_previous_impact__creates_new_simulation( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] @@ -374,6 +396,66 @@ def test__given_no_previous_impact__creates_new_simulation( ) assert write_values["options"] == MOCK_OPTIONS assert write_values["reform_impact_json"] == {} + assert write_values["observability_id"] == MOCK_OBSERVABILITY_ID + + def test__given_no_previous_impact__starts_observability_after_claim( + self, + economy_service, + base_params, + mock_country_package_versions, + mock_policyengine_version, + mock_policy_service, + mock_reform_impacts_service, + mock_simulation_entrypoint, + mock_logger, + mock_datetime, + mock_submission_claim_id, + ): + mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] + lifecycle_events = [] + mock_reform_impacts_service.claim_reform_impact_start.side_effect = ( + lambda **_kwargs: lifecycle_events.append("claim") or True + ) + + with patch( + "policyengine_api.services.economy_service.start_observability_id", + side_effect=lambda _value: ( + lifecycle_events.append("start") or MOCK_OBSERVABILITY_ID + ), + ) as start_observability_id: + result = economy_service.get_economic_impact(**base_params) + + assert result.status == ImpactStatus.COMPUTING + start_observability_id.assert_called_once_with(MOCK_OBSERVABILITY_ID) + assert lifecycle_events == ["claim", "start"] + + def test__selected_observability_id_is_applied_to_containing_spans( + self, + economy_service, + base_params, + mock_country_package_versions, + mock_policyengine_version, + mock_policy_service, + mock_reform_impacts_service, + mock_simulation_entrypoint, + mock_logger, + mock_datetime, + mock_submission_claim_id, + ): + mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] + + with patch( + "policyengine_api.services.economy_service.set_runtime_context" + ) as set_runtime_context: + result = economy_service.get_economic_impact(**base_params) + + assert result.status == ImpactStatus.COMPUTING + assert ( + set_runtime_context.call_args_list.count( + call(observability_id=MOCK_OBSERVABILITY_ID) + ) + == 2 + ) def test__given_existing_start_claim__does_not_submit_duplicate_simulation( self, @@ -386,17 +468,102 @@ def test__given_existing_start_claim__does_not_submit_duplicate_simulation( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.claim_reform_impact_start.return_value = False + winning_observability_id = "00000000-0000-4000-8000-000000000003" + mock_reform_impacts_service.get_reform_impact_start_claim.return_value = ( + ReformImpactStartClaim( + submission_claim_id="winning-claim", + observability_id=winning_observability_id, + ) + ) - result = economy_service.get_economic_impact(**base_params) + with patch( + "policyengine_api.services.economy_service.restore_observability_id", + return_value=winning_observability_id, + ) as restore_observability_id: + result = economy_service.get_economic_impact(**base_params) assert result.status is ImpactStatus.COMPUTING + restore_observability_id.assert_called_once_with(winning_observability_id) mock_simulation_entrypoint.run.assert_not_called() mock_reform_impacts_service.set_reform_impact.assert_not_called() mock_reform_impacts_service.release_reform_impact_start.assert_not_called() + def test__given_expired_contended_claim__retries_and_submits( + self, + economy_service, + base_params, + mock_country_package_versions, + mock_policyengine_version, + mock_policy_service, + mock_reform_impacts_service, + mock_simulation_entrypoint, + mock_logger, + mock_datetime, + mock_submission_claim_id, + ): + mock_reform_impacts_service.claim_reform_impact_start.side_effect = [ + False, + True, + ] + mock_reform_impacts_service.get_reform_impact_start_claim.return_value = ( + None + ) + + with patch( + "policyengine_api.services.economy_service.start_observability_id", + return_value=MOCK_OBSERVABILITY_ID, + ) as start_observability_id: + result = economy_service.get_economic_impact(**base_params) + + assert result.status is ImpactStatus.COMPUTING + assert mock_reform_impacts_service.claim_reform_impact_start.call_count == 2 + mock_reform_impacts_service.get_reform_impact_start_claim.assert_called_once() + start_observability_id.assert_called_once_with(MOCK_OBSERVABILITY_ID) + mock_simulation_entrypoint.run.assert_called_once() + + def test__given_repeatedly_expiring_start_claims__fails_before_submission( + self, + economy_service, + base_params, + mock_country_package_versions, + mock_policyengine_version, + mock_policy_service, + mock_reform_impacts_service, + mock_simulation_entrypoint, + mock_logger, + mock_datetime, + mock_submission_claim_id, + ): + mock_reform_impacts_service.claim_reform_impact_start.return_value = False + mock_reform_impacts_service.get_reform_impact_start_claim.return_value = ( + None + ) + + with ( + patch( + "policyengine_api.services.economy_service.start_observability_id" + ) as start_observability_id, + pytest.raises( + CacheCoordinationError, + match="ownership changed repeatedly", + ), + ): + economy_service.get_economic_impact(**base_params) + + assert ( + mock_reform_impacts_service.claim_reform_impact_start.call_count + == REFORM_IMPACT_START_CLAIM_ATTEMPTS + ) + assert ( + mock_reform_impacts_service.get_reform_impact_start_claim.call_count + == REFORM_IMPACT_START_CLAIM_ATTEMPTS + ) + start_observability_id.assert_not_called() + mock_simulation_entrypoint.run.assert_not_called() + def test__given_start_claim_cache_failure__fails_before_submission( self, economy_service, @@ -408,7 +575,7 @@ def test__given_start_claim_cache_failure__fails_before_submission( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.claim_reform_impact_start.side_effect = ( CacheCoordinationError("cache unavailable") @@ -431,7 +598,7 @@ def test__given_gateway_raises_before_returning_execution__releases_start_claim( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_simulation_entrypoint.run.side_effect = RuntimeError( "submission failed" @@ -455,7 +622,7 @@ def test__given_submitted_simulation_handoff_failure__retains_start_claim( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.set_reform_impact.side_effect = ( ReformImpactHandoffError("cache unavailable") @@ -478,7 +645,7 @@ def test__given_submitted_simulation_without_execution_id__retains_start_claim( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_simulation_entrypoint.get_execution_id.side_effect = RuntimeError( "missing execution identifier" @@ -519,7 +686,6 @@ def test__given_policies_created_through_orm__submits_decoded_json( MOCK_MODEL_VERSION, ) simulation_gateway.get_execution_id.return_value = "execution-1" - simulation_gateway.run.return_value.run_id = "run-1" monkeypatch.setattr( "policyengine_api.services.economy_service.logger", MagicMock(), @@ -559,7 +725,7 @@ def test__given_no_previous_impact__includes_metadata_in_simulation_params( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): """Verify that _metadata with policy IDs is passed to simulation API.""" mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] @@ -576,7 +742,10 @@ def test__given_no_previous_impact__includes_metadata_in_simulation_params( assert ( sim_params["_metadata"]["baseline_policy_id"] == MOCK_BASELINE_POLICY_ID ) - assert sim_params["_metadata"]["process_id"] == MOCK_PROCESS_ID + assert ( + sim_params["_metadata"]["submission_claim_id"] + == MOCK_SUBMISSION_CLAIM_ID + ) assert sim_params["_metadata"]["model_version"] == MOCK_MODEL_VERSION assert ( sim_params["_metadata"]["policyengine_version"] @@ -600,7 +769,7 @@ def test__given_no_previous_impact__includes_telemetry_in_simulation_params( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.get_all_reform_impacts.return_value = [] @@ -608,15 +777,18 @@ def test__given_no_previous_impact__includes_telemetry_in_simulation_params( sim_params = mock_simulation_entrypoint.run.call_args[0][0] - assert sim_params["_telemetry"]["run_id"] - assert sim_params["_telemetry"]["process_id"] == MOCK_PROCESS_ID + assert "observability_id" not in sim_params["_telemetry"] + assert ( + sim_params["_telemetry"]["submission_claim_id"] + == MOCK_SUBMISSION_CLAIM_ID + ) assert sim_params["_telemetry"]["simulation_kind"] == "national" assert sim_params["_telemetry"]["geography_type"] == "national" assert sim_params["_telemetry"]["geography_code"] == MOCK_COUNTRY_ID assert sim_params["_telemetry"]["capture_mode"] == "disabled" assert sim_params["_telemetry"]["config_hash"].startswith("sha256:") progress_log = mock_logger.log_struct.call_args_list[-1].args[0] - assert progress_log["run_id"] == MOCK_RUN_ID + assert progress_log["observability_id"] == MOCK_OBSERVABILITY_ID assert ( mock_logger.log_struct.call_args_list[-1].kwargs["severity"] == "INFO" ) @@ -632,7 +804,7 @@ def test__given_runtime_cache_version__uses_versioned_economy_cache_key( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, monkeypatch, ): cache_version = "e1cache01" @@ -669,7 +841,7 @@ def test__given_default_dataset__queries_previous_impacts_with_resolved_bundle( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] @@ -696,7 +868,7 @@ def test__given_completed_impact__uses_resolved_runtime_bundle_for_cache_lookup( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact(status="ok") mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ @@ -712,6 +884,39 @@ def test__given_completed_impact__uses_resolved_runtime_bundle_for_cache_lookup( policyengine_version=MOCK_POLICYENGINE_VERSION, ) + def test__given_existing_impact__restores_stored_observability_id( + self, + economy_service, + base_params, + mock_country_package_versions, + mock_policyengine_version, + mock_policy_service, + mock_reform_impacts_service, + mock_simulation_entrypoint, + mock_logger, + mock_datetime, + mock_submission_claim_id, + ): + completed_impact = create_mock_reform_impact(status="ok") + mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ + completed_impact + ] + + with ( + patch( + "policyengine_api.services.economy_service.restore_observability_id", + return_value=MOCK_OBSERVABILITY_ID, + ) as restore_observability_id, + patch( + "policyengine_api.services.economy_service.start_observability_id" + ) as start_observability_id, + ): + result = economy_service.get_economic_impact(**base_params) + + assert result.status == ImpactStatus.OK + restore_observability_id.assert_called_once_with(MOCK_OBSERVABILITY_ID) + start_observability_id.assert_not_called() + def test__given_cached_impact_and_runtime_lookup_fails__then_returns_cached_result( self, economy_service, @@ -723,7 +928,7 @@ def test__given_cached_impact_and_runtime_lookup_fails__then_returns_cached_resu mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact(status="ok") mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [ @@ -752,7 +957,7 @@ def test__given_legacy_cached_impact_without_resolved_app_name__then_refreshes_c mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact( status="ok", @@ -784,7 +989,7 @@ def test__given_legacy_and_refreshed_cached_impacts__then_reuses_refreshed_entry mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): legacy_impact = create_mock_reform_impact( status="ok", @@ -823,7 +1028,7 @@ def test__given_legacy_cached_impact_and_runtime_lookup_fails__then_returns_cach mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): completed_impact = create_mock_reform_impact( status="ok", @@ -854,7 +1059,7 @@ def test__given_legacy_computing_impact_without_resolved_app_name__then_reuses_e mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): computing_impact = create_mock_reform_impact( status="computing", @@ -882,7 +1087,7 @@ def test__given_exception__raises_error( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.side_effect = Exception( "Database error" @@ -902,7 +1107,7 @@ def test__given_uk_request__preserves_model_version_in_bundle( mock_simulation_entrypoint, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): mock_country_package_versions["uk"] = "2.7.8" mock_reform_impacts_service.get_all_reform_impacts_by_options_hash_prefix.return_value = [] @@ -930,7 +1135,7 @@ def economy_service( mock_policy_service, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, ): return EconomyService() @@ -984,10 +1189,14 @@ def test__given_no_cached_batch__submits_parent_batch_and_returns_queued_result( assert submitted_payload["target"] == "general" assert "time_period" not in submitted_payload mock_budget_window_cache.claim_batch_start.assert_called_once_with( - "budget-window-cache-key", MOCK_PROCESS_ID + "budget-window-cache-key", + MOCK_SUBMISSION_CLAIM_ID, + MOCK_OBSERVABILITY_ID, ) - mock_budget_window_cache.store_batch_job_id.assert_called_once_with( - "budget-window-cache-key", "fc-budget-123" + mock_budget_window_cache.store_submitted.assert_called_once_with( + "budget-window-cache-key", + "fc-budget-123", + MOCK_OBSERVABILITY_ID, ) mock_reform_impacts_service.set_reform_impact.assert_not_called() @@ -1022,8 +1231,10 @@ def test__given_completed_cached_result__returns_completed_batch_result( "budgetaryImpact": 90, }, } - mock_budget_window_cache.get_completed_result.return_value = ( - completed_result + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="completed", + result=completed_result, + observability_id=MOCK_OBSERVABILITY_ID, ) result = economy_service.get_budget_window_economic_impact(**base_params) @@ -1042,7 +1253,11 @@ def test__given_cached_batch_id__returns_running_batch_progress( mock_simulation_entrypoint, mock_budget_window_cache, ): - mock_budget_window_cache.get_batch_job_id.return_value = "fc-budget-123" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="fc-budget-123", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_simulation_entrypoint.get_budget_window_batch_by_id.return_value = ( create_mock_budget_window_batch_execution( batch_job_id="fc-budget-123", @@ -1082,7 +1297,11 @@ def test__given_completed_batch_poll__caches_result_and_returns_completed( "annualImpacts": [], "totals": {}, } - mock_budget_window_cache.get_batch_job_id.return_value = "fc-budget-123" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="fc-budget-123", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_simulation_entrypoint.get_budget_window_batch_by_id.return_value = ( create_mock_budget_window_batch_execution( batch_job_id="fc-budget-123", @@ -1099,10 +1318,9 @@ def test__given_completed_batch_poll__caches_result_and_returns_completed( assert result.data == completed_result assert result.cache_status == "batch-id-hit" mock_budget_window_cache.set_completed_result.assert_called_once_with( - "budget-window-cache-key", completed_result - ) - mock_budget_window_cache.clear_batch_job_id.assert_called_once_with( - "budget-window-cache-key" + "budget-window-cache-key", + completed_result, + MOCK_OBSERVABILITY_ID, ) @pytest.mark.parametrize("malformed_result", [None, {}, []]) @@ -1114,7 +1332,11 @@ def test__given_completed_batch_without_result__returns_error_without_caching( mock_budget_window_cache, malformed_result, ): - mock_budget_window_cache.get_batch_job_id.return_value = "fc-budget-123" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="fc-budget-123", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_simulation_entrypoint.get_budget_window_batch_by_id.return_value = ( create_mock_budget_window_batch_execution( batch_job_id="fc-budget-123", @@ -1136,9 +1358,7 @@ def test__given_completed_batch_without_result__returns_error_without_caching( assert result.queued_years == ["2028"] assert result.cache_status == "batch-id-hit" mock_budget_window_cache.set_completed_result.assert_not_called() - mock_budget_window_cache.clear_batch_job_id.assert_called_once_with( - "budget-window-cache-key" - ) + mock_budget_window_cache.set_execution_failure.assert_called_once() def test__given_completed_batch_cache_write_fails__does_not_clear_batch_id( self, @@ -1155,7 +1375,11 @@ def test__given_completed_batch_cache_write_fails__does_not_clear_batch_id( "annualImpacts": [], "totals": {}, } - mock_budget_window_cache.get_batch_job_id.return_value = "fc-budget-123" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="fc-budget-123", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_budget_window_cache.set_completed_result.return_value = False mock_simulation_entrypoint.get_budget_window_batch_by_id.return_value = ( create_mock_budget_window_batch_execution( @@ -1171,7 +1395,11 @@ def test__given_completed_batch_cache_write_fails__does_not_clear_batch_id( assert result.status == ImpactStatus.OK assert result.data == completed_result - mock_budget_window_cache.clear_batch_job_id.assert_not_called() + mock_budget_window_cache.set_completed_result.assert_called_once_with( + "budget-window-cache-key", + completed_result, + MOCK_OBSERVABILITY_ID, + ) def test__given_failed_batch_poll__returns_failed( self, @@ -1180,7 +1408,11 @@ def test__given_failed_batch_poll__returns_failed( mock_simulation_entrypoint, mock_budget_window_cache, ): - mock_budget_window_cache.get_batch_job_id.return_value = "fc-budget-123" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="fc-budget-123", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_simulation_entrypoint.get_budget_window_batch_by_id.return_value = ( create_mock_budget_window_batch_execution( batch_job_id="fc-budget-123", @@ -1202,9 +1434,7 @@ def test__given_failed_batch_poll__returns_failed( assert result.queued_years == ["2028"] assert result.cache_status == "batch-id-hit" mock_budget_window_cache.set_completed_result.assert_not_called() - mock_budget_window_cache.clear_batch_job_id.assert_called_once_with( - "budget-window-cache-key" - ) + mock_budget_window_cache.set_execution_failure.assert_called_once() def test_typed_error_write_failure_retains_batch_identity( self, @@ -1214,7 +1444,11 @@ def test_typed_error_write_failure_retains_batch_identity( mock_budget_window_cache, ): error = SPMValidationError("SPM_YEAR_UNAVAILABLE", "No forecast for 2036") - mock_budget_window_cache.get_batch_job_id.return_value = "expired-job" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="expired-job", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_budget_window_cache.set_terminal_error.return_value = False mock_simulation_entrypoint.get_budget_window_batch_by_id.side_effect = error @@ -1222,7 +1456,6 @@ def test_typed_error_write_failure_retains_batch_identity( economy_service.get_budget_window_economic_impact(**base_params) assert raised.value is error - mock_budget_window_cache.clear_batch_job_id.assert_not_called() mock_budget_window_cache.set_completed_result.assert_not_called() mock_simulation_entrypoint.run_budget_window_batch.assert_not_called() @@ -1234,7 +1467,11 @@ def test_untyped_poll_error_keeps_existing_retry_behavior( mock_budget_window_cache, ): error = make_http_status_error(422, payload={"detail": "Unknown error"}) - mock_budget_window_cache.get_batch_job_id.return_value = "existing-job" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="existing-job", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_simulation_entrypoint.get_budget_window_batch_by_id.side_effect = error with pytest.raises(httpx.HTTPStatusError) as raised: @@ -1242,7 +1479,6 @@ def test_untyped_poll_error_keeps_existing_retry_behavior( assert raised.value is error mock_budget_window_cache.set_terminal_error.assert_not_called() - mock_budget_window_cache.clear_batch_job_id.assert_not_called() mock_simulation_entrypoint.run_budget_window_batch.assert_not_called() def test__given_existing_start_claim__does_not_submit_duplicate_batch( @@ -1252,14 +1488,143 @@ def test__given_existing_start_claim__does_not_submit_duplicate_batch( mock_simulation_entrypoint, mock_budget_window_cache, ): + winning_observability_id = "00000000-0000-4000-8000-000000000099" mock_budget_window_cache.claim_batch_start.return_value = False + mock_budget_window_cache.get_state.side_effect = [ + None, + BudgetWindowCacheState( + status="starting", + submission_claim_id="winning-claim", + observability_id=winning_observability_id, + ), + ] - result = economy_service.get_budget_window_economic_impact(**base_params) + with ( + patch( + "policyengine_api.services.economy_service.restore_observability_id", + return_value=winning_observability_id, + ) as restore_observability_id, + patch( + "policyengine_api.services.economy_service.start_observability_id" + ) as start_observability_id, + ): + result = economy_service.get_budget_window_economic_impact( + **base_params + ) assert result.status == ImpactStatus.COMPUTING assert result.progress == 0 assert result.queued_years == ["2026", "2027", "2028"] assert result.cache_status == "starting-claim-hit" + restore_observability_id.assert_called_with(winning_observability_id) + start_observability_id.assert_not_called() + mock_simulation_entrypoint.run_budget_window_batch.assert_not_called() + + def test__given_disappearing_start_claim__retries_before_binding_identity( + self, + economy_service, + base_params, + mock_simulation_entrypoint, + mock_budget_window_cache, + ): + lifecycle_events = [] + claim_results = iter([False, True]) + mock_budget_window_cache.claim_batch_start.side_effect = lambda *_args: ( + lifecycle_events.append("claim") or next(claim_results) + ) + mock_budget_window_cache.get_state.side_effect = [None, None] + mock_simulation_entrypoint.run_budget_window_batch.return_value = ( + create_mock_budget_window_batch_execution( + batch_job_id="fc-budget-123", + status="submitted", + ) + ) + + with patch( + "policyengine_api.services.economy_service.start_observability_id", + side_effect=lambda value: lifecycle_events.append("bind") or value, + ) as start_observability_id: + result = economy_service.get_budget_window_economic_impact( + **base_params + ) + + assert result.status == ImpactStatus.COMPUTING + assert result.cache_status == "miss" + assert mock_budget_window_cache.claim_batch_start.call_count == 2 + assert mock_budget_window_cache.claim_batch_start.call_args_list == [ + call( + "budget-window-cache-key", + MOCK_SUBMISSION_CLAIM_ID, + MOCK_OBSERVABILITY_ID, + ), + call( + "budget-window-cache-key", + MOCK_SUBMISSION_CLAIM_ID, + MOCK_OBSERVABILITY_ID, + ), + ] + start_observability_id.assert_called_once_with(MOCK_OBSERVABILITY_ID) + assert lifecycle_events == ["claim", "claim", "bind"] + mock_simulation_entrypoint.run_budget_window_batch.assert_called_once() + + def test__given_repeated_disappearing_claims__fails_without_binding_identity( + self, + economy_service, + base_params, + mock_simulation_entrypoint, + mock_budget_window_cache, + ): + mock_budget_window_cache.claim_batch_start.return_value = False + mock_budget_window_cache.get_state.return_value = None + + with ( + patch( + "policyengine_api.services.economy_service.start_observability_id" + ) as start_observability_id, + pytest.raises( + CacheCoordinationError, + match="submission ownership changed repeatedly", + ), + ): + economy_service.get_budget_window_economic_impact(**base_params) + + assert mock_budget_window_cache.claim_batch_start.call_count == 3 + start_observability_id.assert_not_called() + mock_simulation_entrypoint.run_budget_window_batch.assert_not_called() + + def test__given_cached_execution_failure__replays_failure_and_identity( + self, + economy_service, + base_params, + mock_simulation_entrypoint, + mock_budget_window_cache, + ): + stored_observability_id = "00000000-0000-4000-8000-000000000099" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="failed", + observability_id=stored_observability_id, + failure_type="execution", + error={ + "status": "error", + "error": "Budget window failed for 2027", + "completed_years": ["2026"], + "queued_years": ["2028"], + }, + ) + + with patch( + "policyengine_api.services.economy_service.restore_observability_id", + return_value=stored_observability_id, + ) as restore_observability_id: + result = economy_service.get_budget_window_economic_impact( + **base_params + ) + + assert result.status == ImpactStatus.ERROR + assert result.error == "Budget window failed for 2027" + assert result.cache_status == "failure-hit" + restore_observability_id.assert_called_with(stored_observability_id) + mock_simulation_entrypoint.get_budget_window_batch_by_id.assert_not_called() mock_simulation_entrypoint.run_budget_window_batch.assert_not_called() def test__given_gateway_raises_before_returning_batch__clears_start_claim( @@ -1277,7 +1642,9 @@ def test__given_gateway_raises_before_returning_batch__clears_start_claim( economy_service.get_budget_window_economic_impact(**base_params) mock_budget_window_cache.clear_starting_claim.assert_called_once_with( - "budget-window-cache-key", MOCK_PROCESS_ID + "budget-window-cache-key", + MOCK_SUBMISSION_CLAIM_ID, + MOCK_OBSERVABILITY_ID, ) @pytest.mark.parametrize("status_code", [400, 422]) @@ -1311,10 +1678,8 @@ def test__given_modal_rejects_batch_submission_for_validation__returns_failed_re assert result.computing_years == [] assert result.queued_years == ["2026", "2027", "2028"] assert result.cache_status == "miss" - mock_budget_window_cache.clear_starting_claim.assert_called_once_with( - "budget-window-cache-key", MOCK_PROCESS_ID - ) - mock_budget_window_cache.store_batch_job_id.assert_not_called() + mock_budget_window_cache.clear_starting_claim.assert_not_called() + mock_budget_window_cache.set_execution_failure.assert_called_once() @pytest.mark.parametrize("status_code", [401, 403, 429, 500]) def test__given_modal_non_validation_error_on_batch_submission__raises( @@ -1333,9 +1698,11 @@ def test__given_modal_non_validation_error_on_batch_submission__raises( economy_service.get_budget_window_economic_impact(**base_params) mock_budget_window_cache.clear_starting_claim.assert_called_once_with( - "budget-window-cache-key", MOCK_PROCESS_ID + "budget-window-cache-key", + MOCK_SUBMISSION_CLAIM_ID, + MOCK_OBSERVABILITY_ID, ) - mock_budget_window_cache.store_batch_job_id.assert_not_called() + mock_budget_window_cache.store_submitted.assert_not_called() @pytest.mark.parametrize( ("payload", "expected_message"), @@ -1422,7 +1789,7 @@ def test__given_runtime_cache_version__uses_versioned_cache_key_for_budget_windo mock_budget_window_cache, mock_logger, mock_datetime, - mock_numpy_random, + mock_submission_claim_id, monkeypatch, ): cache_version = "e1cache01" @@ -1445,9 +1812,11 @@ def test__given_reordered_options__uses_same_budget_window_cache_identity( mock_simulation_entrypoint, mock_budget_window_cache, ): - mock_budget_window_cache.get_completed_result.return_value = { - "kind": "budgetWindow" - } + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="completed", + result={"kind": "budgetWindow"}, + observability_id=MOCK_OBSERVABILITY_ID, + ) economy_service.get_budget_window_economic_impact( **{ @@ -1498,7 +1867,11 @@ def test__given_unexpected_batch_status__raises_value_error( mock_simulation_entrypoint, mock_budget_window_cache, ): - mock_budget_window_cache.get_batch_job_id.return_value = "fc-budget-123" + mock_budget_window_cache.get_state.return_value = BudgetWindowCacheState( + status="submitted", + batch_job_id="fc-budget-123", + observability_id=MOCK_OBSERVABILITY_ID, + ) mock_simulation_entrypoint.get_budget_window_batch_by_id.return_value = ( create_mock_budget_window_batch_execution( batch_job_id="fc-budget-123", @@ -1583,7 +1956,8 @@ def economy_service(self): @pytest.fixture def setup_options(self): return EconomicImpactSetupOptions( - process_id=MOCK_PROCESS_ID, + submission_claim_id=MOCK_SUBMISSION_CLAIM_ID, + observability_id=MOCK_OBSERVABILITY_ID, country_id=MOCK_COUNTRY_ID, reform_policy_id=MOCK_POLICY_ID, baseline_policy_id=MOCK_BASELINE_POLICY_ID, @@ -1683,7 +2057,8 @@ def economy_service(self): @pytest.fixture def setup_options(self): return EconomicImpactSetupOptions( - process_id=MOCK_PROCESS_ID, + submission_claim_id=MOCK_SUBMISSION_CLAIM_ID, + observability_id=MOCK_OBSERVABILITY_ID, country_id=MOCK_COUNTRY_ID, reform_policy_id=MOCK_POLICY_ID, baseline_policy_id=MOCK_BASELINE_POLICY_ID, @@ -1899,19 +2274,16 @@ def test__given_modal_submitted_state__then_returns_computing_result( assert result.status == ImpactStatus.COMPUTING assert result.data is None - class TestCreateProcessId: + class TestCreateSubmissionClaimId: @pytest.fixture def economy_service(self): return EconomyService() - def test_given_mocked_datetime_and_random_returns_expected_format( - self, economy_service, mock_datetime, mock_numpy_random - ): - result = economy_service._create_process_id() + def test_returns_uuid_string(self, economy_service, mock_submission_claim_id): + result = economy_service._create_submission_claim_id() - assert result == "job_20250626120000_1234" - mock_datetime.now.assert_called_once() - mock_numpy_random.assert_called_once_with(1000, 9999) + assert result == MOCK_SUBMISSION_CLAIM_ID + mock_submission_claim_id.assert_called_once_with() class TestEconomicImpactResult: @@ -1978,7 +2350,8 @@ def test__given_error__creates_correct_instance_and_logs(self): class TestEconomicImpactSetupOptions: def test__given_valid_data__creates_instance(self): options = EconomicImpactSetupOptions( - process_id=MOCK_PROCESS_ID, + submission_claim_id=MOCK_SUBMISSION_CLAIM_ID, + observability_id=MOCK_OBSERVABILITY_ID, country_id=MOCK_COUNTRY_ID, reform_policy_id=MOCK_POLICY_ID, baseline_policy_id=MOCK_BASELINE_POLICY_ID, @@ -1991,7 +2364,7 @@ def test__given_valid_data__creates_instance(self): options_hash=MOCK_OPTIONS_HASH, ) - assert options.process_id == MOCK_PROCESS_ID + assert options.submission_claim_id == MOCK_SUBMISSION_CLAIM_ID assert options.country_id == MOCK_COUNTRY_ID assert options.reform_policy_id == MOCK_POLICY_ID assert options.baseline_policy_id == MOCK_BASELINE_POLICY_ID diff --git a/tests/unit/services/test_household_calculation_service.py b/tests/unit/services/test_household_calculation_service.py index 0efed99bc..5259abad7 100644 --- a/tests/unit/services/test_household_calculation_service.py +++ b/tests/unit/services/test_household_calculation_service.py @@ -1,8 +1,9 @@ from contextlib import contextmanager from pathlib import Path from types import SimpleNamespace -from unittest.mock import MagicMock +from unittest.mock import MagicMock, Mock +import pytest from policyengine_api.constants import COUNTRY_PACKAGE_VERSIONS, POLICYENGINE_VERSION from policyengine_api.data.v1_models import ( Household, @@ -18,6 +19,8 @@ from policyengine_api.services.household_calculation_service import ( CalculationResult, HouseholdCalculationService, + HouseholdNotFoundError, + PolicyNotFoundError, ) @@ -128,6 +131,104 @@ def calculate(self, household, policy): assert result.warnings == ("employment_income could not be calculated",) +def test_parsed_country_calculation_is_reused_for_calculation(): + parsed_calculation = object() + + class Country: + metadata = { + "variables": {}, + "entities": {"person": {"plural": "people", "roles": {}}}, + "parameters": {}, + } + + def __init__(self): + self.prepare = Mock(return_value=parsed_calculation) + self.calculate = Mock( + return_value=CalculationResult(household={"people": {}}) + ) + + def prepare_calculation(self, household, policy): + return self.prepare(household, policy) + + country = Country() + service = HouseholdCalculationService( + cache=_cache(), + country_provider=lambda: {"us": country}, + ) + prepared = service.prepare_household_calculation( + "us", + {"people": {}}, + {}, + ) + + parsed = service.parse_prepared_household(prepared) + result = service.calculate_prepared_household(parsed) + + country.prepare.assert_called_once_with({"people": {}}, {}) + country.calculate.assert_called_once_with( + {"people": {}}, + {}, + prepared=parsed_calculation, + ) + assert result.household == {"people": {}} + + +def test_successful_situation_parsing_accepts_calculation_before_stage_ends(): + accepted = Mock() + + class Country: + metadata = { + "variables": {}, + "entities": {"person": {"plural": "people", "roles": {}}}, + "parameters": {}, + } + + def prepare_calculation(self, household, policy): + return object() + + service = HouseholdCalculationService( + cache=_cache(), + country_provider=lambda: {"us": Country()}, + ) + prepared = service.prepare_household_calculation("us", {"people": {}}, {}) + + parsed = service.parse_prepared_household( + prepared, + on_accepted=accepted, + ) + + assert parsed.country_calculation is not None + accepted.assert_called_once_with() + + +def test_failed_situation_parsing_does_not_accept_calculation(): + accepted = Mock() + + class Country: + metadata = { + "variables": {}, + "entities": {"person": {"plural": "people", "roles": {}}}, + "parameters": {}, + } + + def prepare_calculation(self, household, policy): + raise RuntimeError("invalid situation") + + service = HouseholdCalculationService( + cache=_cache(), + country_provider=lambda: {"us": Country()}, + ) + prepared = service.prepare_household_calculation("us", {"people": {}}, {}) + + with pytest.raises(RuntimeError, match="invalid situation"): + service.parse_prepared_household( + prepared, + on_accepted=accepted, + ) + + accepted.assert_not_called() + + def test_calculation_closes_reads_before_compute_and_caches_atomic_results( orm_session_factory, monkeypatch, @@ -169,6 +270,53 @@ def calculate(self, household, policy): } +def test_missing_stored_household_is_not_accepted(orm_session_factory): + accepted = Mock() + service = HouseholdCalculationService( + primary_session_factory=orm_session_factory, + cache=_cache(), + ) + + with pytest.raises(HouseholdNotFoundError): + service.calculate_stored_household( + "us", + 1, + 2, + on_accepted=accepted, + ) + + accepted.assert_not_called() + + +def test_missing_stored_policy_is_not_accepted(orm_session_factory): + with orm_session_factory.begin() as session: + session.add( + Household( + id=1, + country_id="us", + label=None, + api_version=COUNTRY_PACKAGE_VERSIONS["us"], + household_json={"people": {"you": {}}}, + household_hash="household-hash", + ) + ) + accepted = Mock() + service = HouseholdCalculationService( + primary_session_factory=orm_session_factory, + cache=_cache(), + ) + + with pytest.raises(PolicyNotFoundError): + service.calculate_stored_household( + "us", + 1, + 2, + on_accepted=accepted, + ) + + accepted.assert_not_called() + + def test_calculation_uses_local_cache_without_recomputing(orm_session_factory): _seed_inputs(orm_session_factory) calculated = {"people": {"you": {"net_income": {"2026": 42}}}} @@ -191,12 +339,19 @@ def test_calculation_uses_local_cache_without_recomputing(orm_session_factory): cache=cache, country_provider=lambda: {"us": country}, ) + accepted = Mock() - result = service.calculate_stored_household("us", 1, 2) + result = service.calculate_stored_household( + "us", + 1, + 2, + on_accepted=accepted, + ) assert result.household == calculated assert result.warnings == ("net_income could not be calculated",) assert result.cached is True + accepted.assert_called_once_with() def test_failed_cache_write_does_not_invalidate_successful_calculation( diff --git a/tests/unit/services/test_reform_impacts_service.py b/tests/unit/services/test_reform_impacts_service.py index 16a6bc943..c4371bffa 100644 --- a/tests/unit/services/test_reform_impacts_service.py +++ b/tests/unit/services/test_reform_impacts_service.py @@ -95,16 +95,31 @@ def test_reform_impact_start_claim_is_exclusive_and_releasable(service): "api_version": "1", "target": "general", } + owner_observability_id = "00000000-0000-4000-8000-000000000001" + contender_observability_id = "00000000-0000-4000-8000-000000000002" - assert service.claim_reform_impact_start(**arguments, claim_token="owner") + assert service.claim_reform_impact_start( + **arguments, + claim_token="owner", + observability_id=owner_observability_id, + ) assert not service.claim_reform_impact_start( **arguments, claim_token="contender", + observability_id=contender_observability_id, + ) + claim = service.get_reform_impact_start_claim(**arguments) + assert claim.submission_claim_id == "owner" + assert claim.observability_id == owner_observability_id + service.release_reform_impact_start( + **arguments, + claim_token="owner", + observability_id=owner_observability_id, ) - service.release_reform_impact_start(**arguments, claim_token="owner") assert service.claim_reform_impact_start( **arguments, claim_token="contender", + observability_id=contender_observability_id, ) diff --git a/tests/unit/test_asgi_factory.py b/tests/unit/test_asgi_factory.py index d4d1aceaa..b4660a92e 100644 --- a/tests/unit/test_asgi_factory.py +++ b/tests/unit/test_asgi_factory.py @@ -17,6 +17,7 @@ RouteImplementation, RouteImplementationSettings, ) +from policyengine_api.observability.identifiers import OBSERVABILITY_ID_HEADER from policyengine_api.request_context import ( REQUEST_ID_HEADER, current_request_id, @@ -49,6 +50,14 @@ def request_echo(): response.headers["X-Echo"] = "present" return response + @app.get("/stored-observability-id") + def stored_observability_id(): + response = make_response("stored", 200) + response.headers[OBSERVABILITY_ID_HEADER] = ( + "00000000-0000-4000-8000-000000000012" + ) + return response + @app.get("/readiness-check") def readiness_check(): return Response("OK", status=200, mimetype="text/plain") @@ -303,6 +312,124 @@ def capture_request(**kwargs): generate_request_id.assert_called_once_with() +def test_native_route_uses_observability_request_lifecycle(): + runtime = Mock() + runtime.begin_request.return_value = "request-123" + runtime.response_headers.return_value = { + "traceparent": "00-00000000000000000000000000000001-0000000000000001-01" + } + observability_id = "00000000-0000-4000-8000-000000000001" + + response = TestClient( + create_asgi_app( + create_test_wsgi_app(), + observability_runtime=runtime, + ) + ).get( + "/health", + headers={ + REQUEST_ID_HEADER: "request-123", + OBSERVABILITY_ID_HEADER: observability_id, + }, + ) + + assert response.status_code == 200 + assert response.headers["traceparent"].startswith("00-") + runtime.begin_request.assert_called_once() + assert runtime.begin_request.call_args.kwargs["method"] == "GET" + assert runtime.begin_request.call_args.kwargs["route"] == "/health" + runtime.set_context.assert_called_once_with(request_id="request-123") + assert OBSERVABILITY_ID_HEADER not in response.headers + runtime.update_request_route.assert_called_once_with("/health") + runtime.update_request_status.assert_called_once_with(200) + runtime.end_request.assert_called_once_with(status_code=200, error=None) + + +def test_native_request_span_starts_with_route_template(): + runtime = Mock() + runtime.begin_request.return_value = "request-123" + runtime.response_headers.return_value = {} + + response = TestClient( + create_asgi_app( + create_test_wsgi_app(), + observability_runtime=runtime, + ) + ).get("/v2/tax-benefit-models/by-country/us") + + assert response.status_code != 404 + runtime.begin_request.assert_called_once() + assert runtime.begin_request.call_args.kwargs["route"] == ( + "/v2/tax-benefit-models/by-country/{country_id}" + ) + runtime.update_request_route.assert_called_once_with( + "/v2/tax-benefit-models/by-country/{country_id}" + ) + + +def test_flask_fallback_does_not_duplicate_observability_request_lifecycle(): + runtime = Mock() + + response = TestClient( + create_asgi_app( + create_test_wsgi_app(), + observability_runtime=runtime, + ) + ).get("/fallback") + + assert response.status_code == 202 + runtime.begin_request.assert_not_called() + runtime.end_request.assert_not_called() + + +def test_native_route_survives_observability_runtime_failures(): + runtime = Mock() + runtime.begin_request.side_effect = RuntimeError("begin unavailable") + runtime.set_context.side_effect = RuntimeError("context unavailable") + runtime.response_headers.side_effect = RuntimeError("headers unavailable") + runtime.update_request_route.side_effect = RuntimeError("route unavailable") + runtime.update_request_status.side_effect = RuntimeError("status unavailable") + runtime.end_request.side_effect = RuntimeError("finish unavailable") + + response = TestClient( + create_asgi_app( + create_test_wsgi_app(), + observability_runtime=runtime, + ) + ).get("/health") + + assert response.status_code == 200 + assert response.json() == {"status": "healthy"} + assert response.headers[REQUEST_ID_HEADER] + assert OBSERVABILITY_ID_HEADER not in response.headers + + +def test_native_route_reapplies_calculation_identifier_to_server_request_span(): + runtime = Mock() + runtime.begin_request.return_value = "request-123" + runtime.response_headers.return_value = {} + + with patch( + "policyengine_api.asgi_factory.current_observability_id", + return_value="00000000-0000-4000-8000-000000000001", + ): + response = TestClient( + create_asgi_app( + create_test_wsgi_app(), + observability_runtime=runtime, + ) + ).get("/health") + + assert response.status_code == 200 + runtime.set_context.assert_any_call( + observability_id="00000000-0000-4000-8000-000000000001" + ) + assert ( + response.headers[OBSERVABILITY_ID_HEADER] + == "00000000-0000-4000-8000-000000000001" + ) + + def test_native_route_does_not_accept_x_request_id_as_an_alias(): with ( patch( @@ -415,6 +542,18 @@ def test_flask_fallback_preserves_status_body_headers_and_cookies(): assert response.headers["content-type"].startswith("text/html") +def test_flask_fallback_preserves_a_stored_observability_id(): + client = TestClient(create_asgi_app(create_test_wsgi_app())) + + response = client.get("/stored-observability-id") + + assert response.status_code == 200 + assert ( + response.headers[OBSERVABILITY_ID_HEADER] + == "00000000-0000-4000-8000-000000000012" + ) + + def test_large_flask_fallback_response_supports_http_gzip(): client = TestClient(create_asgi_app(create_test_wsgi_app())) diff --git a/tests/unit/test_cloud_run_deploy_scripts.py b/tests/unit/test_cloud_run_deploy_scripts.py index 601011366..24cbe8c89 100644 --- a/tests/unit/test_cloud_run_deploy_scripts.py +++ b/tests/unit/test_cloud_run_deploy_scripts.py @@ -19,6 +19,7 @@ TEST_V2_RUNTIME_SECRET_RESOURCE = ( "projects/test-project/secrets/v2-runtime-database-url/versions/latest" ) +TEST_OTEL_ENDPOINT = "https://collector.example.test" CLOUD_RUN_SERVICE_SCRIPTS = ( "scripts/deploy_cloud_run_candidate.sh", "scripts/capture_cloud_run_service_state.sh", @@ -73,6 +74,15 @@ def _v2_target_env() -> dict[str, str]: } +def _observability_env() -> dict[str, str]: + return { + "OBSERVABILITY_SERVICE_NAMESPACE": "policyengine.api-v1", + "OBSERVABILITY_TRACE_PROJECT_ID": "central-observability", + "OTEL_EXPORTER_OTLP_ENDPOINT": TEST_OTEL_ENDPOINT, + "POLICYENGINE_OTEL_GOOGLE_AUDIENCE": TEST_OTEL_ENDPOINT, + } + + def _required_runtime_env() -> dict[str, str]: return { "DEPLOYMENT_ENVIRONMENT": "production", @@ -106,6 +116,7 @@ def _required_runtime_env() -> dict[str, str]: "DB_WRITE_POLICY": "cloud_sql", "DB_READ_HOUSEHOLD": "cloud_sql", "DB_WRITE_HOUSEHOLD": "cloud_sql", + **_observability_env(), **_v2_target_env(), **_gateway_auth_env(), } @@ -610,6 +621,7 @@ def test_validate_cloud_run_deploy_env_accepts_direct_mode_from_environment(): ), **_v2_target_env(), **_gateway_auth_env(), + **_observability_env(), ), ) @@ -766,6 +778,7 @@ def test_validate_cloud_run_deploy_env_requires_only_selected_url( ), **_v2_target_env(), **_gateway_auth_env(), + **_observability_env(), ) missing_result = _run_script( ".github/scripts/validate_cloud_run_deploy_env.sh", @@ -988,6 +1001,15 @@ def test_deploy_cloud_run_candidate_dry_run_preserves_access_and_traffic(): assert "RUNTIME_CACHE_MODE=deployed" in result.stdout assert "RUNTIME_CACHE_ENVIRONMENT=production" in result.stdout assert "RUNTIME_CACHE_SERVICE=api" in result.stdout + assert "APP_ENVIRONMENT=production" in result.stdout + assert f"OTEL_EXPORTER_OTLP_ENDPOINT={TEST_OTEL_ENDPOINT}" in result.stdout + assert "OBSERVABILITY_SERVICE_NAMESPACE=policyengine.api-v1" in result.stdout + assert "OBSERVABILITY_TRACE_PROJECT_ID=central-observability" in result.stdout + assert "OTEL_EXPORTER_OTLP_PROTOCOL=grpc" in result.stdout + assert "OTEL_TRACES_EXPORTER=otlp" in result.stdout + assert "OTEL_METRICS_EXPORTER=otlp" in result.stdout + assert "OTEL_TRACES_SAMPLER_ARG=1.0" in result.stdout + assert f"POLICYENGINE_OTEL_GOOGLE_AUDIENCE={TEST_OTEL_ENDPOINT}" in result.stdout assert ( "RUNTIME_CACHE_URL=policyengine-api-prod-runtime-cache-url:latest" in result.stdout diff --git a/tests/unit/test_gcp_logging.py b/tests/unit/test_gcp_logging.py index 2edceadea..af939614f 100644 --- a/tests/unit/test_gcp_logging.py +++ b/tests/unit/test_gcp_logging.py @@ -1,49 +1,47 @@ from unittest.mock import Mock -from policyengine_api.gcp_logging import _LazyGoogleLogger - - -def test_local_logging_uses_stderr_without_initializing_google(monkeypatch): - monkeypatch.delenv("K_SERVICE", raising=False) - logger = _LazyGoogleLogger("test-local") - logger._fallback_logger = Mock() - payload = {"message": "cache miss"} - - logger.log_struct(payload, severity="WARNING", labels={"cache": "analysis"}) - - assert logger._initialization_failed is True - assert logger._google_logger is None - logger._fallback_logger.log.assert_called_once_with(30, "%s", payload) - - -def test_remote_logging_failure_falls_back_and_disables_retries(monkeypatch): - monkeypatch.setenv("K_SERVICE", "policyengine-api") - remote_logger = Mock() - remote_logger.log_struct.side_effect = ConnectionError("logging unavailable") - logger = _LazyGoogleLogger("test-deployed") - logger._google_logger = remote_logger - logger._fallback_logger = Mock() - payload = {"message": "cache write"} - - logger.log_struct(payload, severity="INFO", labels={"cache": "household"}) - logger.log_struct(payload, severity="INFO", labels={"cache": "household"}) +from policyengine_api import gcp_logging +from policyengine_api.gcp_logging import _RuntimeLogger + + +def test_runtime_logger_flattens_migration_context_and_omits_response_text( + monkeypatch, +): + runtime = Mock() + monkeypatch.setattr(gcp_logging, "get_runtime", lambda: runtime) + logger = _RuntimeLogger() + + logger.log_struct( + { + "message": "API request served", + "request_id": "request-1", + "response_text": "must not be recorded", + "migration": { + "route_impl": "flask_fallback", + "db_read_source": "cloud_sql", + }, + }, + severity="WARNING", + labels={"backend": "simulation_entry"}, + ) - remote_logger.log_struct.assert_called_once_with( - payload, - severity="INFO", - labels={"cache": "household"}, + runtime.log.assert_called_once_with( + "API request served", + severity="WARNING", + attributes={ + "request_id": "request-1", + "route_impl": "flask_fallback", + "db_read_source": "cloud_sql", + "backend": "simulation_entry", + }, ) - assert logger._initialization_failed is True - assert logger._google_logger is None - assert logger._fallback_logger.log.call_count == 2 -def test_fallback_logging_failure_does_not_escape(monkeypatch): - monkeypatch.delenv("K_SERVICE", raising=False) - logger = _LazyGoogleLogger("test-broken-fallback") - logger._fallback_logger = Mock() - logger._fallback_logger.log.side_effect = RuntimeError("logging unavailable") +def test_runtime_logging_failure_does_not_escape(monkeypatch): + runtime = Mock() + runtime.log.side_effect = RuntimeError("logging unavailable") + monkeypatch.setattr(gcp_logging, "get_runtime", lambda: runtime) - logger.log_struct({"message": "operation succeeded"}, severity="INFO") + _RuntimeLogger().log_struct({"message": "operation succeeded"}, severity="INFO") - logger._fallback_logger.log.assert_called_once() + runtime.log.assert_called_once() diff --git a/tests/unit/test_observability_collector_config.py b/tests/unit/test_observability_collector_config.py new file mode 100644 index 000000000..e2bc0e1a7 --- /dev/null +++ b/tests/unit/test_observability_collector_config.py @@ -0,0 +1,16 @@ +from __future__ import annotations + +from pathlib import Path + +ROOT = Path(__file__).parents[2] +DEPLOY = ROOT / "gcp" / "observability" + + +def test_collector_accepts_only_traces_and_metrics() -> None: + config = (DEPLOY / "collector" / "config.yaml").read_text() + assert "telemetry.googleapis.com:443" in config + assert "memory_limiter" in config + assert "googleclientauth" in config + assert " traces:" in config + assert " metrics:" in config + assert " logs:\n receivers:" not in config diff --git a/tests/unit/test_observability_runtime.py b/tests/unit/test_observability_runtime.py new file mode 100644 index 000000000..8eba73d22 --- /dev/null +++ b/tests/unit/test_observability_runtime.py @@ -0,0 +1,43 @@ +from policyengine_observability import ( + GoogleCloudLogFormatter, + StdoutLogDestination, +) + +from policyengine_api.observability import ( + _build_runtime, + runtime, + set_runtime_context, +) + + +def test_runtime_uses_consumer_owned_identity_and_stdout(monkeypatch): + monkeypatch.setenv("OTEL_SDK_DISABLED", "true") + monkeypatch.setenv("OTEL_TRACES_SAMPLER_ARG", "0.01") + monkeypatch.setenv("OBSERVABILITY_SERVICE_NAMESPACE", "example.stack") + monkeypatch.setenv("OBSERVABILITY_TRACE_PROJECT_ID", "trace-project") + + runtime = _build_runtime() + try: + assert runtime.config.service.namespace == "example.stack" + assert runtime.config.otel.sampling_ratio == 1.0 + assert runtime.config.application_attribute_keys is None + assert runtime.config.dispatch_attribute_keys == frozenset({"observability_id"}) + assert len(runtime.config.logging.destinations) == 1 + destination = runtime.config.logging.destinations[0] + assert isinstance(destination, StdoutLogDestination) + assert isinstance(destination.formatter, GoogleCloudLogFormatter) + assert destination.formatter.project_id == "trace-project" + finally: + runtime.shutdown() + + +def test_runtime_context_failure_does_not_escape(monkeypatch): + monkeypatch.setattr( + runtime, + "set_context", + lambda **_attributes: (_ for _ in ()).throw( + RuntimeError("observability unavailable") + ), + ) + + set_runtime_context(country_id="us") diff --git a/tests/unit/test_observability_stage_registry.py b/tests/unit/test_observability_stage_registry.py new file mode 100644 index 000000000..3698be2c9 --- /dev/null +++ b/tests/unit/test_observability_stage_registry.py @@ -0,0 +1,31 @@ +import pytest + +from policyengine_api.observability.stages import ( + RUN_STAGE_REGISTRY, + RunConfiguration, + Stage, +) + + +def test_registry_defines_every_run_configuration_and_stage() -> None: + assert set(RUN_STAGE_REGISTRY) == set(RunConfiguration) + registered = { + stage + for stage_plan in RUN_STAGE_REGISTRY.values() + for stage in stage_plan.stages + } + assert registered == set(Stage) + + +@pytest.mark.parametrize("stage_plan", RUN_STAGE_REGISTRY.values()) +def test_each_stage_plan_is_ordered_and_contains_no_duplicates(stage_plan) -> None: + assert stage_plan.stages + assert len(stage_plan.stages) == len(set(stage_plan.stages)) + assert all(stage_plan.name(stage) == stage.value for stage in stage_plan.stages) + + +def test_stage_plan_rejects_a_stage_from_another_configuration() -> None: + household = RUN_STAGE_REGISTRY[RunConfiguration.HOUSEHOLD] + + with pytest.raises(ValueError, match="not registered"): + household.name(Stage.ECONOMY_SUBMIT) diff --git a/tests/unit/test_request_context.py b/tests/unit/test_request_context.py new file mode 100644 index 000000000..0ccdbd686 --- /dev/null +++ b/tests/unit/test_request_context.py @@ -0,0 +1,77 @@ +from unittest.mock import Mock, patch + +from flask import Flask, g + +from policyengine_api.request_context import ( + current_observability_id, + restore_observability_id, + start_observability_id, +) + + +FIRST_ID = "00000000-0000-4000-8000-000000000001" +SECOND_ID = "00000000-0000-4000-8000-000000000002" + + +def test_start_uses_valid_incoming_identifier_and_binds_runtime(): + app = Flask(__name__) + runtime = Mock() + + with ( + app.test_request_context(), + patch("policyengine_api.observability.get_runtime", return_value=runtime), + ): + g.incoming_observability_id = FIRST_ID + g.observability_id = None + + result = start_observability_id() + + assert result == FIRST_ID + assert current_observability_id() == FIRST_ID + runtime.set_context.assert_called_once_with(observability_id=FIRST_ID) + + +def test_durable_state_restoration_replaces_request_candidate(): + app = Flask(__name__) + runtime = Mock() + + with ( + app.test_request_context(), + patch("policyengine_api.observability.get_runtime", return_value=runtime), + ): + g.observability_id = FIRST_ID + + result = restore_observability_id(SECOND_ID) + + assert result == SECOND_ID + assert current_observability_id() == SECOND_ID + runtime.set_context.assert_called_once_with(observability_id=SECOND_ID) + + +def test_null_durable_state_does_not_bind_incoming_candidate(): + app = Flask(__name__) + + with app.test_request_context(): + g.incoming_observability_id = FIRST_ID + g.observability_id = None + + result = restore_observability_id(None) + + assert result is None + assert current_observability_id() is None + + +def test_runtime_failure_does_not_change_identifier_selection(): + app = Flask(__name__) + runtime = Mock() + runtime.set_context.side_effect = RuntimeError("observability unavailable") + + with ( + app.test_request_context(), + patch("policyengine_api.observability.get_runtime", return_value=runtime), + ): + g.incoming_observability_id = FIRST_ID + g.observability_id = None + + assert start_observability_id() == FIRST_ID + assert current_observability_id() == FIRST_ID diff --git a/uv.lock b/uv.lock index 64d96a293..9d035a443 100644 --- a/uv.lock +++ b/uv.lock @@ -1444,18 +1444,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5d/13/ad7d7ca3808a898b4612b6fe93cde56b53f3034dcde235acb1f0e1df24c6/idna-3.13-py3-none-any.whl", hash = "sha256:892ea0cde124a99ce773decba204c5552b69c3c67ffd5f232eb7696135bc8bb3", size = 68629, upload-time = "2026-04-22T16:42:40.909Z" }, ] -[[package]] -name = "importlib-metadata" -version = "8.7.1" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "zipp" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/f3/49/3b30cad09e7771a4982d9975a8cbf64f00d4a1ececb53297f1d9a7be1b10/importlib_metadata-8.7.1.tar.gz", hash = "sha256:49fef1ae6440c182052f407c8d34a68f72efc36db9ca90dc0113398f2fdde8bb", size = 57107, upload-time = "2025-12-21T10:00:19.278Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/fa/5e/f8e9a1d23b9c20a551a8a02ea3637b4642e22c2626e3a13a9a29cdea99eb/importlib_metadata-8.7.1-py3-none-any.whl", hash = "sha256:5a1f80bf1daa489495071efbb095d75a634cf28a8bc299581244063b53176151", size = 27865, upload-time = "2025-12-21T10:00:18.329Z" }, -] - [[package]] name = "iniconfig" version = "2.3.0" @@ -2448,15 +2436,83 @@ wheels = [ [[package]] name = "opentelemetry-api" -version = "1.41.1" +version = "1.44.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "importlib-metadata" }, { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/fa/fc/b7564cbef36601aef0d6c9bc01f7badb64be8e862c2e1c3c5c3b43b53e4f/opentelemetry_api-1.41.1.tar.gz", hash = "sha256:0ad1814d73b875f84494387dae86ce0b12c68556331ce6ce8fe789197c949621", size = 71416, upload-time = "2026-04-24T13:15:38.262Z" } +sdist = { url = "https://files.pythonhosted.org/packages/ee/8b/aa9e2d8b8dfa7c946f7dec5d1f8f6ba8eca062f43509a06bdb5ce93d26c0/opentelemetry_api-1.44.0.tar.gz", hash = "sha256:67647e5e9566edcf421166fdf022b3537f818635daa852b289e34604dc6fb33a", size = 72406, upload-time = "2026-07-16T15:25:32.678Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/29/59/3e7118ed140f76b0982ba4321bdaed1997a0473f9720de2d10788a577033/opentelemetry_api-1.41.1-py3-none-any.whl", hash = "sha256:a22df900e75c76dc08440710e51f52f1aa6b451b429298896023e60db5b3139f", size = 69007, upload-time = "2026-04-24T13:15:15.662Z" }, + { url = "https://files.pythonhosted.org/packages/ca/6f/a04e900f465ff3221ccc395522503e2d10e79fa21f2723c8e177aae1e0d1/opentelemetry_api-1.44.0-py3-none-any.whl", hash = "sha256:94b98c893a91b88657eaac1e3ba89618cdb85be6918196705354f34728b2cdef", size = 60018, upload-time = "2026-07-16T15:25:11.657Z" }, +] + +[[package]] +name = "opentelemetry-exporter-otlp-proto-common" +version = "1.44.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "opentelemetry-proto" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/61/09/4d717852c1cf3f854b76c7110a5d00883bc3c99288b9b0dbcbeb9e306eb6/opentelemetry_exporter_otlp_proto_common-1.44.0.tar.gz", hash = "sha256:dc87a5a5bc58f149a56d1547e4691588fa12994cdc3bc039a694ccb3375862ac", size = 20202, upload-time = "2026-07-16T15:25:37.658Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/5e/71/65fd9d54c10b860f87c045ccee1264cab7011268895d3528818a29c1172a/opentelemetry_exporter_otlp_proto_common-1.44.0-py3-none-any.whl", hash = "sha256:9a9fe61bba73d802904bc989f1d6b4a7b1ee40f06c40e98d6f85af65aaebb694", size = 17045, upload-time = "2026-07-16T15:25:18.201Z" }, +] + +[[package]] +name = "opentelemetry-exporter-otlp-proto-grpc" +version = "1.44.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "googleapis-common-protos" }, + { name = "grpcio" }, + { name = "opentelemetry-api" }, + { name = "opentelemetry-exporter-otlp-proto-common" }, + { name = "opentelemetry-proto" }, + { name = "opentelemetry-sdk" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/1f/47/80d9e9d468dc5de3af5096f5ccdb065fa4dd1470f74495cc53e59e397f47/opentelemetry_exporter_otlp_proto_grpc-1.44.0.tar.gz", hash = "sha256:40d1ae9e03fcc36de3cbac610cc99f35894938bff9cfd90fc4ec68bd85448463", size = 27225, upload-time = "2026-07-16T15:25:38.308Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/54/29/6ae42ba32b153ae0a44ae125f0caff2188bbe62d99c82d1768da30864e72/opentelemetry_exporter_otlp_proto_grpc-1.44.0-py3-none-any.whl", hash = "sha256:6a1a645ea182a2f59440c51fa8301d309f3324a8f9d65f8395584b064b67ee4e", size = 19624, upload-time = "2026-07-16T15:25:19.096Z" }, +] + +[[package]] +name = "opentelemetry-proto" +version = "1.44.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "protobuf" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/64/01/40ac4ae9a149263cc52c2cee200ddd80cb6d8db1a4610abf8eabce0fe771/opentelemetry_proto-1.44.0.tar.gz", hash = "sha256:c547a79c2f8c0c515d31509154682e5921c7cfd5ca67b70e1f9266e2c3e103f3", size = 46488, upload-time = "2026-07-16T15:25:45.34Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d1/7c/8be563d68e93bbefa5c8affb82ddcff91b3ad858ce49957ba7b16fd3e0ab/opentelemetry_proto-1.44.0-py3-none-any.whl", hash = "sha256:898b155a0e1557afd867478fb6158e8122a46329ca0bb8dc53cc55e98f017f56", size = 72483, upload-time = "2026-07-16T15:25:28.429Z" }, +] + +[[package]] +name = "opentelemetry-sdk" +version = "1.44.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "opentelemetry-api" }, + { name = "opentelemetry-semantic-conventions" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/5d/77/a6592cbc7c8d9bcc9d6757a9df45e04a7c585e3e6e7a13456da522b21109/opentelemetry_sdk-1.44.0.tar.gz", hash = "sha256:cebe7f65dc12f26ead75c6064de12fd2a9052e5060c0272d402cfa203aae123b", size = 208624, upload-time = "2026-07-16T15:25:46.078Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e7/23/ff077e61886ee020a17ce9c8b6fa11c601c8d8345b09ea24f605445df62a/opentelemetry_sdk-1.44.0-py3-none-any.whl", hash = "sha256:df081c4c6bcfdb1211e3e86140376792643128a25f8d72d1d27675936e7e96ad", size = 137221, upload-time = "2026-07-16T15:25:29.534Z" }, +] + +[[package]] +name = "opentelemetry-semantic-conventions" +version = "0.65b0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "opentelemetry-api" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/8f/73/0cbdebcb4cf545fdd328da14f5137e37d0770c3f26185e478b0d15d94f50/opentelemetry_semantic_conventions-0.65b0.tar.gz", hash = "sha256:f9b2b81e9d5b64f11bc952075e7e9c7fb0aab075c7fd1c46d597f1b919852d60", size = 148774, upload-time = "2026-07-16T15:25:46.902Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a6/0e/49df70d9b81fb5cbae4bbf2a49d865b09bcbcbc4eb53f5851b1027738d78/opentelemetry_semantic_conventions-0.65b0-py3-none-any.whl", hash = "sha256:1cacde7b0ad306f84c5ef08c3dbe1bbaf20165bba6f8bff43b670e555a086bcb", size = 204645, upload-time = "2026-07-16T15:25:30.688Z" }, ] [[package]] @@ -2713,7 +2769,7 @@ models = [ [[package]] name = "policyengine-api" -version = "3.56.1" +version = "4.2.2" source = { editable = "." } dependencies = [ { name = "a2wsgi" }, @@ -2738,6 +2794,7 @@ dependencies = [ { name = "policyengine-canada" }, { name = "policyengine-il" }, { name = "policyengine-ng" }, + { name = "policyengine-observability", extra = ["flask", "google", "httpx", "otlp-grpc"] }, { name = "psycopg", extra = ["binary"] }, { name = "pydantic" }, { name = "pymysql" }, @@ -2791,6 +2848,7 @@ requires-dist = [ { name = "policyengine-canada", specifier = "==0.96.3" }, { name = "policyengine-il", specifier = "==0.1.0" }, { name = "policyengine-ng", specifier = "==0.5.1" }, + { name = "policyengine-observability", extras = ["flask", "google", "httpx", "otlp-grpc"], specifier = ">=3.0.1,<4" }, { name = "psycopg", extras = ["binary"], specifier = ">=3.3,<4" }, { name = "pydantic" }, { name = "pymysql" }, @@ -2892,6 +2950,32 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/03/9e/1040e63f72f3857d540e52c0f99115f8db04da5e94231959d4245a119bef/policyengine_ng-0.5.1-py3-none-any.whl", hash = "sha256:21fad6aae8d80a156142ac876cf1b7679e036c1640ca6cb375661701f10b9920", size = 31074, upload-time = "2023-04-19T13:14:28.242Z" }, ] +[[package]] +name = "policyengine-observability" +version = "3.0.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/00/a6/f5d3e49523bf3d3231e9d2297fb03a4ae0704ae246d86153356dcb7eee7c/policyengine_observability-3.0.1.tar.gz", hash = "sha256:35e934b31843545a13570d6ba2e6a4b0cb002488f2c7bb1840351f6f90f30104", size = 124370, upload-time = "2026-09-28T16:48:19.916Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/27/6c/879ecc01a618885e4bc741d8abb35e601d638f878e3145cda677b75c46b7/policyengine_observability-3.0.1-py3-none-any.whl", hash = "sha256:b7630d4d2d91c183b733d2e3470d3c0cbad67fd59a8c5e2f2bcc75872eea3c02", size = 40859, upload-time = "2026-09-28T16:48:18.661Z" }, +] + +[package.optional-dependencies] +flask = [ + { name = "flask" }, +] +google = [ + { name = "google-auth" }, + { name = "google-cloud-logging" }, +] +httpx = [ + { name = "httpx" }, +] +otlp-grpc = [ + { name = "opentelemetry-api" }, + { name = "opentelemetry-exporter-otlp-proto-grpc" }, + { name = "opentelemetry-sdk" }, +] + [[package]] name = "policyengine-uk" version = "2.90.2" @@ -4567,15 +4651,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/69/68/c8739671f5699c7dc470580a4f821ef37c32c4cb0b047ce223a7f115757f/yarl-1.23.0-py3-none-any.whl", hash = "sha256:a2df6afe50dea8ae15fa34c9f824a3ee958d785fd5d089063d960bae1daa0a3f", size = 48288, upload-time = "2026-03-01T22:07:51.388Z" }, ] -[[package]] -name = "zipp" -version = "3.23.1" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/30/21/093488dfc7cc8964ded15ab726fad40f25fd3d788fd741cc1c5a17d78ee8/zipp-3.23.1.tar.gz", hash = "sha256:32120e378d32cd9714ad503c1d024619063ec28aad2248dc6672ad13edfa5110", size = 25965, upload-time = "2026-04-13T23:21:46.6Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/08/8a/0861bec20485572fbddf3dfba2910e38fe249796cb73ecdeb74e07eeb8d3/zipp-3.23.1-py3-none-any.whl", hash = "sha256:0b3596c50a5c700c9cb40ba8d86d9f2cc4807e9bedb06bcdf7fac85633e444dc", size = 10378, upload-time = "2026-04-13T23:21:45.386Z" }, -] - [[package]] name = "zope-interface" version = "8.4"