From 78e380bbd633be3086157efe7544384a45e7e1e8 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sat, 29 Aug 2026 17:46:52 +0400 Subject: [PATCH 01/27] Implement Stage 9 v2 metadata catalog --- .github/scripts/cloud_run_env.sh | 2 + .github/scripts/deploy_cloud_run_candidate.sh | 1 + .github/scripts/migrate_v2_metadata_schema.sh | 8 + .../scripts/validate_cloud_run_deploy_env.sh | 1 + .github/workflows/alembic-v2-check.yml | 10 + .github/workflows/initialize-v2-metadata.yml | 44 + .github/workflows/push.yml | 32 +- docs/engineering/migration-contracts.md | 14 +- docs/generated/migration_contracts.json | 35 +- docs/migration/stage-9-v2-metadata.md | 205 ++++ ...5dc5_version_metadata_catalog_snapshots.py | 218 ++++ policyengine_api/asgi_factory.py | 2 + policyengine_api/data/v2/catalog/__init__.py | 15 + .../data/v2/catalog/extraction.py | 712 +++++++++++++ .../data/v2/catalog/initialization.py | 88 ++ .../data/v2/catalog/publication.py | 932 ++++++++++++++++++ policyengine_api/data/v2/catalog/query.py | 342 +++++++ policyengine_api/data/v2/catalog/records.py | 200 ++++ policyengine_api/data/v2/catalog/schemas.py | 151 +++ policyengine_api/data/v2/models/metadata.py | 57 +- policyengine_api/data/v2/settings.py | 76 +- .../fastapi_routes/dependencies.py | 32 + .../fastapi_routes/v2_metadata.py | 154 +++ policyengine_api/migration_flags.py | 33 +- policyengine_api/migration_logging.py | 11 + pyproject.toml | 1 + scripts/guards/migration_contracts.py | 7 +- scripts/initialize_v2_metadata.py | 8 + tests/contract/registry.py | 38 + .../test_app_v2_workflow_contracts.py | 22 +- tests/contract/test_v1_route_contracts.py | 4 +- tests/fixtures/v2_catalog.py | 234 +++++ .../integration/test_alembic_v2_lifecycle.py | 26 +- .../integration/test_v2_catalog_installed.py | 82 ++ .../test_v2_catalog_publication.py | 342 +++++++ ...st_v2_catalog_publication_qualification.py | 123 +++ tests/integration/test_v2_metadata_routes.py | 246 +++++ .../routes/test_migration_context_logging.py | 23 + tests/unit/test_alembic_workflows.py | 6 + tests/unit/test_cloud_run_deploy_scripts.py | 52 +- .../unit/test_migration_contract_artifacts.py | 5 +- tests/unit/test_migration_flags.py | 2 + tests/unit/v2/test_alembic_v2.py | 30 +- tests/unit/v2/test_catalog_extraction.py | 344 +++++++ tests/unit/v2/test_catalog_initialization.py | 143 +++ tests/unit/v2/test_catalog_publication.py | 221 +++++ tests/unit/v2/test_import_side_effects.py | 6 + tests/unit/v2/test_metadata_deployment.py | 96 ++ tests/unit/v2/test_metadata_query.py | 488 +++++++++ tests/unit/v2/test_metadata_routes.py | 371 +++++++ tests/unit/v2/test_model_persistence.py | 47 +- tests/unit/v2/test_models.py | 72 +- tests/unit/v2/test_report_runs.py | 9 +- tests/unit/v2/test_settings.py | 92 +- uv.lock | 2 + 55 files changed, 6453 insertions(+), 64 deletions(-) create mode 100644 .github/scripts/migrate_v2_metadata_schema.sh create mode 100644 .github/workflows/initialize-v2-metadata.yml create mode 100644 docs/migration/stage-9-v2-metadata.md create mode 100644 migrations/v2/versions/68b4a5ae5dc5_version_metadata_catalog_snapshots.py create mode 100644 policyengine_api/data/v2/catalog/__init__.py create mode 100644 policyengine_api/data/v2/catalog/extraction.py create mode 100644 policyengine_api/data/v2/catalog/initialization.py create mode 100644 policyengine_api/data/v2/catalog/publication.py create mode 100644 policyengine_api/data/v2/catalog/query.py create mode 100644 policyengine_api/data/v2/catalog/records.py create mode 100644 policyengine_api/data/v2/catalog/schemas.py create mode 100644 policyengine_api/fastapi_routes/v2_metadata.py create mode 100644 scripts/initialize_v2_metadata.py create mode 100644 tests/fixtures/v2_catalog.py create mode 100644 tests/integration/test_v2_catalog_installed.py create mode 100644 tests/integration/test_v2_catalog_publication.py create mode 100644 tests/integration/test_v2_catalog_publication_qualification.py create mode 100644 tests/integration/test_v2_metadata_routes.py create mode 100644 tests/unit/v2/test_catalog_extraction.py create mode 100644 tests/unit/v2/test_catalog_initialization.py create mode 100644 tests/unit/v2/test_catalog_publication.py create mode 100644 tests/unit/v2/test_metadata_deployment.py create mode 100644 tests/unit/v2/test_metadata_query.py create mode 100644 tests/unit/v2/test_metadata_routes.py diff --git a/.github/scripts/cloud_run_env.sh b/.github/scripts/cloud_run_env.sh index 396d4babe..37522c53d 100755 --- a/.github/scripts/cloud_run_env.sh +++ b/.github/scripts/cloud_run_env.sh @@ -43,6 +43,7 @@ cloud_run_set_defaults() { CLOUD_RUN_RUNTIME_CACHE_URL_SECRET="${CLOUD_RUN_RUNTIME_CACHE_URL_SECRET:-policyengine-api-prod-runtime-cache-url:latest}" CLOUD_RUN_RUNTIME_CACHE_CA_CERT_SECRET="${CLOUD_RUN_RUNTIME_CACHE_CA_CERT_SECRET:-policyengine-api-prod-runtime-cache-ca:latest}" CLOUD_RUN_RUNTIME_CACHE_ENVIRONMENT="${CLOUD_RUN_RUNTIME_CACHE_ENVIRONMENT:-production}" + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE="${V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE:-}" CLOUD_RUN_VPC_NETWORK="${CLOUD_RUN_VPC_NETWORK:-default}" CLOUD_RUN_VPC_SUBNET="${CLOUD_RUN_VPC_SUBNET:-default}" CLOUD_RUN_VPC_EGRESS="${CLOUD_RUN_VPC_EGRESS:-private-ranges-only}" @@ -79,6 +80,7 @@ cloud_run_set_defaults() { export CLOUD_RUN_RUNTIME_CACHE_URL_SECRET export CLOUD_RUN_RUNTIME_CACHE_CA_CERT_SECRET export CLOUD_RUN_RUNTIME_CACHE_ENVIRONMENT + export V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE export CLOUD_RUN_VPC_NETWORK export CLOUD_RUN_VPC_SUBNET export CLOUD_RUN_VPC_EGRESS diff --git a/.github/scripts/deploy_cloud_run_candidate.sh b/.github/scripts/deploy_cloud_run_candidate.sh index 261048103..4c4b28894 100755 --- a/.github/scripts/deploy_cloud_run_candidate.sh +++ b/.github/scripts/deploy_cloud_run_candidate.sh @@ -29,6 +29,7 @@ env_vars=( "RUNTIME_CACHE_SERVICE=api" "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}" ) if [[ -n "${OLD_SIMULATION_GATEWAY_URL:-}" ]]; then diff --git a/.github/scripts/migrate_v2_metadata_schema.sh b/.github/scripts/migrate_v2_metadata_schema.sh new file mode 100644 index 000000000..edbd5cbce --- /dev/null +++ b/.github/scripts/migrate_v2_metadata_schema.sh @@ -0,0 +1,8 @@ +#!/usr/bin/env bash + +set -euo pipefail +set +x + +uv run alembic -c alembic-v2.ini upgrade head +uv run alembic -c alembic-v2.ini current --check-heads +uv run alembic -c alembic-v2.ini check diff --git a/.github/scripts/validate_cloud_run_deploy_env.sh b/.github/scripts/validate_cloud_run_deploy_env.sh index 700cfd490..03a2f2403 100755 --- a/.github/scripts/validate_cloud_run_deploy_env.sh +++ b/.github/scripts/validate_cloud_run_deploy_env.sh @@ -36,6 +36,7 @@ cloud_run_require_env \ CLOUD_RUN_VPC_EGRESS \ V2_SUPABASE_PROJECT_REF \ V2_SUPABASE_ENVIRONMENT \ + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE \ SIM_ENTRYPOINT \ ROUTE_IMPL_HEALTH \ ROUTE_IMPL_SPECIFICATION \ diff --git a/.github/workflows/alembic-v2-check.yml b/.github/workflows/alembic-v2-check.yml index dd05786f9..19da03199 100644 --- a/.github/workflows/alembic-v2-check.yml +++ b/.github/workflows/alembic-v2-check.yml @@ -52,5 +52,15 @@ jobs: run: uv run alembic -c alembic-v2.ini current --check-heads - name: Require no ungenerated v2 schema or data operations run: uv run alembic -c alembic-v2.ini check + - name: Verify the installed PolicyEngine.py catalog interface + run: uv run pytest -q tests/integration/test_v2_catalog_installed.py + env: + RUN_V2_CATALOG_COMPATIBILITY: "1" + - name: Test v2 metadata publication and preview routes + run: uv run pytest -q tests/integration/test_v2_catalog_publication.py tests/integration/test_v2_metadata_routes.py + - name: Qualify production-scale v2 metadata publication + run: uv run pytest -q tests/integration/test_v2_catalog_publication_qualification.py + env: + RUN_V2_CATALOG_PUBLICATION_QUALIFICATION: "1" - name: Test real Redis cross-instance semantics run: uv run pytest -q tests/integration/test_runtime_cache_redis.py diff --git a/.github/workflows/initialize-v2-metadata.yml b/.github/workflows/initialize-v2-metadata.yml new file mode 100644 index 000000000..e4b242a8d --- /dev/null +++ b/.github/workflows/initialize-v2-metadata.yml @@ -0,0 +1,44 @@ +name: Initialize v2 metadata + +on: + workflow_call: + inputs: + deployment_environment: + description: GitHub environment containing one v2 database target + required: true + type: string + workflow_dispatch: + inputs: + deployment_environment: + description: GitHub environment containing one v2 database target + required: true + type: environment + +jobs: + initialize: + name: Upgrade and initialize v2 metadata + runs-on: ubuntu-latest + timeout-minutes: 30 + environment: ${{ inputs.deployment_environment }} + env: + V2_SUPABASE_PROJECT_REF: ${{ vars.V2_SUPABASE_PROJECT_REF }} + V2_SUPABASE_ENVIRONMENT: ${{ vars.V2_SUPABASE_ENVIRONMENT }} + steps: + - name: Checkout repo + uses: actions/checkout@v4 + - name: Setup Python + uses: actions/setup-python@v5 + with: + python-version: "3.12" + - name: Setup uv + uses: astral-sh/setup-uv@v6 + - name: Install locked dependencies + run: uv sync --frozen + - name: Upgrade and verify the v2 schema + run: bash .github/scripts/migrate_v2_metadata_schema.sh + env: + V2_MIGRATION_DATABASE_URL: ${{ secrets.V2_MIGRATION_DATABASE_URL }} + - name: Publish and validate the v2 metadata catalog + run: uv run python scripts/initialize_v2_metadata.py + env: + V2_DATA_WRITE_DATABASE_URL: ${{ secrets.V2_DATA_WRITE_DATABASE_URL }} diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index d320bf972..6a41b3f96 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -153,6 +153,19 @@ jobs: if: always() run: bash .github/scripts/stop_cloud_sql_proxy.sh + initialize-v2-staging: + name: Initialize staging v2 metadata + needs: + - publish-git-tag + - migrate-v1-cloud-sql + if: | + (github.repository == 'PolicyEngine/policyengine-api') + && (github.event.head_commit.message == 'Update PolicyEngine API') + uses: ./.github/workflows/initialize-v2-metadata.yml + with: + deployment_environment: staging + secrets: inherit + deploy-staging: name: Deploy staging App Engine version runs-on: ubuntu-latest @@ -160,6 +173,7 @@ jobs: - ensure-staging-model-version-aligns-with-sim-api - publish-git-tag - migrate-v1-cloud-sql + - initialize-v2-staging if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -255,6 +269,7 @@ jobs: - ensure-staging-model-version-aligns-with-sim-api - publish-git-tag - migrate-v1-cloud-sql + - initialize-v2-staging if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -272,6 +287,7 @@ jobs: CLOUD_RUN_RUNTIME_CACHE_CA_CERT_SECRET: policyengine-api-staging-runtime-cache-ca:latest V2_SUPABASE_PROJECT_REF: ${{ vars.V2_SUPABASE_PROJECT_REF }} V2_SUPABASE_ENVIRONMENT: ${{ vars.V2_SUPABASE_ENVIRONMENT }} + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: ${{ secrets.V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE }} # Staging stays scale-to-zero, single instance: it exists for per-push # validation, not capacity. Both the revision-level (--min-instances) and # service-level (--min) floors are 0. @@ -491,10 +507,21 @@ jobs: - name: Check simulation API supports PolicyEngine bundle run: bash .github/check-policyengine-bundle-supported.sh + initialize-v2-production: + name: Initialize production v2 metadata + needs: ensure-production-model-version-aligns-with-sim-api + if: | + (github.repository == 'PolicyEngine/policyengine-api') + && (github.event.head_commit.message == 'Update PolicyEngine API') + uses: ./.github/workflows/initialize-v2-metadata.yml + with: + deployment_environment: production + secrets: inherit + deploy-production-candidate: name: Deploy production App Engine candidate runs-on: ubuntu-latest - needs: ensure-production-model-version-aligns-with-sim-api + needs: initialize-v2-production if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -620,7 +647,7 @@ jobs: deploy-cloud-run-candidate: name: Deploy production Cloud Run candidate runs-on: ubuntu-latest - needs: ensure-production-model-version-aligns-with-sim-api + needs: initialize-v2-production if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -638,6 +665,7 @@ jobs: CLOUD_RUN_RUNTIME_CACHE_CA_CERT_SECRET: policyengine-api-prod-runtime-cache-ca:latest V2_SUPABASE_PROJECT_REF: ${{ vars.V2_SUPABASE_PROJECT_REF }} V2_SUPABASE_ENVIRONMENT: ${{ vars.V2_SUPABASE_ENVIRONMENT }} + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: ${{ secrets.V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE }} # Sized by the Stage 2 qualification and the PR 4 host cutover — rationale # and numbers in docs/migration/cloud-run-operations.md ("Runtime shape and # scaling"). Warm capacity is expressed service-level (--min); the diff --git a/docs/engineering/migration-contracts.md b/docs/engineering/migration-contracts.md index 648f79d4d..3f96a9b9f 100644 --- a/docs/engineering/migration-contracts.md +++ b/docs/engineering/migration-contracts.md @@ -7,8 +7,8 @@ Generated from `policyengine_api/migration_registry.py` and `tests/contract/regi | Metric | Count | | --- | ---: | | route group count | 9 | -| workflow count | 7 | -| request count | 14 | +| workflow count | 8 | +| request count | 16 | | db entity count | 6 | | sim flow count | 3 | @@ -69,6 +69,16 @@ Generated from `policyengine_api/migration_registry.py` and `tests/contract/regi | `GET` | `/us/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | | `GET` | `/uk/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | +### `region_selection_v2_preview` + +- Current contract: `typed_v2_preview` +- Future owner: Later metadata read cutover and preview-path removal + +| Method | Path | Status | Route group | Stable response fields | +| --- | --- | ---: | --- | --- | +| `GET` | `/v2/us/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | +| `GET` | `/v2/uk/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | + ### `simulation_submit_poll` - Current contract: `api_v1_compatible` diff --git a/docs/generated/migration_contracts.json b/docs/generated/migration_contracts.json index 3ae30a010..ce2464979 100644 --- a/docs/generated/migration_contracts.json +++ b/docs/generated/migration_contracts.json @@ -1,10 +1,10 @@ { "metadata": { "db_entity_count": 6, - "request_count": 14, + "request_count": 16, "route_group_count": 9, "sim_flow_count": 3, - "workflow_count": 7 + "workflow_count": 8 }, "route_groups": [ { @@ -217,6 +217,37 @@ } ] }, + { + "current_contract": "typed_v2_preview", + "future_owner_pr": "Later metadata read cutover and preview-path removal", + "name": "region_selection_v2_preview", + "requests": [ + { + "expected_status": 200, + "method": "GET", + "path": "/v2/us/metadata", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "result.current_law_id", + "result.economy_options.region", + "result.economy_options.time_period" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/uk/metadata", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "result.current_law_id", + "result.economy_options.region", + "result.economy_options.time_period" + ] + } + ] + }, { "current_contract": "api_v1_compatible", "future_owner_pr": "PR 13: Household Calculation Compute Cutover", diff --git a/docs/migration/stage-9-v2-metadata.md b/docs/migration/stage-9-v2-metadata.md new file mode 100644 index 000000000..3176b270c --- /dev/null +++ b/docs/migration/stage-9-v2-metadata.md @@ -0,0 +1,205 @@ +# Stage 9 v2 metadata deployment runbook + +This runbook intentionally omits database hosts, database URLs, passwords, +project references, environment names, secret-resource names, service-account +identities, and physical dataset locations. Resolve those values from the +approved environment inventory and secret-management system. + +## Scope + +Stage 9 populates the dormant v2 US and UK reference catalogs and exposes +read-only preview endpoints from the Cloud Run ASGI application at +`GET /v2/us/metadata` and `GET /v2/uk/metadata`. Their generated OpenAPI +document is available at `GET /v2/openapi.json`. App Engine continues to run +the Flask v1 application and does not expose these routes. Stage 9 does not +change `GET /us/metadata`, `GET /uk/metadata`, their callers, or their v1 data +source. Existing clients must not be redirected to the preview endpoints. + +The initializer creates only reusable logical input `Dataset` rows. Each row +has `is_output_dataset=false` and a null `storage_path`. It creates no +package-derived `DatasetVersion` rows or simulation/report output datasets and +does not modify `Simulation`, `Report`, or `ReportRun` rows or their dataset +references. + +## Configuration and credential boundaries + +Select the target with `V2_SUPABASE_PROJECT_REF` and +`V2_SUPABASE_ENVIRONMENT`. Confirm both values against the approved inventory +before connecting. + +Use three separately managed database credentials: + +| Operation | Configuration | Required database access | +| --- | --- | --- | +| Alembic schema upgrade | `V2_MIGRATION_DATABASE_URL` | Reviewed v2 schema changes; no application-runtime use | +| Catalog publication | `V2_DATA_WRITE_DATABASE_URL` | Row insertion, update, selection, temporary tables, and transaction-scoped advisory locks; no persistent schema changes | +| Preview GET requests | `V2_RUNTIME_DATABASE_URL` or `V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE` | Catalog selection only during Stage 9; no migration access | + +Do not place more than one database URL in the environment of either the +migration or publication process. Cloud Run API processes receive only the +runtime Secret Manager resource identifier. They resolve the URL lazily when a +v2 preview request is made; application startup and unprefixed v1 requests do +not resolve it. App Engine does not receive the v2 runtime database URL or its +Secret Manager resource identifier. + +Give the Cloud Run runtime identity access only to its approved runtime URL +secret. Do not give an application runtime identity access to the migration or +catalog-publication credentials. + +## Pre-activation sequence + +The release workflow calls `.github/workflows/initialize-v2-metadata.yml` for +the selected GitHub Environment before creating an API candidate. Its steps +must execute in this order: + +1. Install the locked application dependencies from the same revision that + will be deployed. +2. Supply only `V2_MIGRATION_DATABASE_URL` and the target identity, then run + `.github/scripts/migrate_v2_metadata_schema.sh`. This upgrades the v2 + Alembic chain, confirms all heads are current, and detects ungenerated model + changes. +3. Remove the migration URL from the command environment. +4. Supply only `V2_DATA_WRITE_DATABASE_URL` and the same target identity, then + run `uv run python scripts/initialize_v2_metadata.py`. +5. Require the command to finish successfully and retain its non-secret JSON + evidence before creating the candidate revision. +6. Deploy the Cloud Run candidate with the target identity and + `V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE`. Keep all existing route selectors + and traffic rules unchanged. + +Any schema-upgrade, revision-check, extraction, compatibility, publication, or +post-publication validation failure stops the workflow before candidate +creation. The currently serving revision remains unchanged. + +Initialization is an explicit deployment operation, once per target and +deployment artifact. It must never be called by a module import, an individual +web or worker process startup, a readiness check, an HTTP request, or an +ordinary process restart. The release workflow may repeat the explicit command +because publication is serialized and idempotent. + +## Source and regional dataset behavior + +The command obtains the US and UK models and packaged dependency selections +through PolicyEngine.py. `TaxBenefitModelVersion.version` stores the canonical +PolicyEngine.py version once. Variables, parameter nodes, parameters, logical +datasets, and regions reference that model-version row; parameter values obtain +the version through their parameter. Model descriptions, current-law IDs, and +metadata time-period options are stored on the model-version snapshot. +PolicyEngine Core and country-package versions are derived diagnostic evidence +and are not separately persisted. + +Parameter extraction requires the PolicyEngine.py public history contract from +5.0.4 and later: values are ordered from oldest to newest, bounded `end_date` +values are inclusive and fall one day before the following `start_date`, and +the newest value is open-ended. The initializer validates that contract and +does not reorder, deduplicate, or repair package output. A value superseded on +the immediately following day therefore has equal start and end dates and is a +valid one-day interval. + +The US national region uses `populace_us_2024`. A US state, +congressional-district, or place region uses its PolicyEngine.py regional +alternative when one is available. Until a later reviewed change removes the +fallback, a subnational region without such an alternative uses +`populace_us_2024`. The command emits one summary warning containing only the +affected region types and counts. UK regions use the +`enhanced_frs_2024_25` default certified by PolicyEngine.py 5.0.4. + +Review the fallback summary on every new PolicyEngine.py release. A changed +count may identify an upstream catalog change that needs a reviewed dataset +selection even when publication otherwise succeeds. + +## Validation evidence and retry + +Successful JSON evidence contains: + +- the canonical PolicyEngine.py version; +- the dependency versions derived from its packaged manifest; +- catalog entity counts; +- US fallback counts; +- elapsed publication time; and +- a success outcome. + +It must contain no database URL, credential, environment or project identity, +physical dataset location, dataset release identifier, digest, parameter +value, or generated catalog payload. + +The initializer validates the complete normalized source before database +mutation. Publication then uses one PostgreSQL transaction, a transaction-level +advisory lock, private temporary staging tables, bounded `COPY` operations, and +set-based reconciliation. Before commit it compares persisted relationships and +counts, confirms input-only region defaults, confirms canonical parameter-value +uniqueness, confirms no package-derived `DatasetVersion` rows were added, and +confirms existing simulation and report record counts are unchanged. + +A failed attempt rolls back all catalog changes. Retry the same artifact only +after correcting the external failure. A matching retry preserves identifiers, +content, and row counts. Different normalized content under the same +PolicyEngine.py version fails for operator review; it is never silently +replaced. A later PolicyEngine.py version adds a catalog version without +deleting or modifying earlier versioned rows, including their logical datasets, +regions, model descriptions, current-law IDs, and time-period options. + +## Production-scale qualification record + +The Stage 9 implementation qualification used the locked PolicyEngine.py +5.0.4 distribution. Its manifest selected PolicyEngine Core 3.30.1, +PolicyEngine US 1.764.6, and PolicyEngine UK 2.90.2. Extraction produced 2 +models, 2 model versions, 6,649 variables, 27,826 named parameter nodes, 99,006 +parameters, 1,172,130 parameter values, 2 logical input datasets, and 826 +regions. Publication took 33.916 seconds and added 7,979,842 bytes of measured +peak publisher memory. The US fallback summary reported 436 congressional +districts, 333 places, and 51 states. + +Re-run the production-scale test only against disposable Postgres: + +```bash +RUN_V2_CATALOG_PUBLICATION_QUALIFICATION=1 \ +V2_ALEMBIC_DISPOSABLE_TEST=1 \ +V2_MIGRATION_DATABASE_URL="" \ +uv run pytest -q tests/integration/test_v2_catalog_publication_qualification.py +``` + +Do not point this qualification command at a persistent staging or production +target. + +## Preview verification + +After Cloud Run candidate creation, explicitly request both preview GET +endpoints and validate their typed response envelopes. A request without a +`policyengine_version` query parameter selects the exact PolicyEngine.py version +installed in that candidate artifact. It does not select the newest database +row. Also request a known published version with, for example, +`?policyengine_version=5.0.4`, and confirm that the response contains that +exact version's complete snapshot. + +A successful response has HTTP 200, `status: "ok"`, `message: null`, and a +typed `result`. A malformed or noncanonical explicit version returns a typed +HTTP 400 error. A valid explicit version that has not been published for the +country returns a typed HTTP 404 error. An absent or incomplete catalog for the +candidate's installed default version returns a typed HTTP 503 service error. +None of these outcomes reads or repairs v1. Unsupported countries and methods +return typed client errors. + +Request `GET /v2/openapi.json` and confirm that the public document contains +the US, UK, and unsupported-country preview paths and explicit component schema +references for every documented response. + +Also request the unprefixed US and UK metadata endpoints and confirm their +responses still come from v1. Do not modify route selectors, internal callers, +or traffic rules to use `/v2` during Stage 9. + +## Rollback + +If Cloud Run candidate verification or activation fails, keep or restore +traffic on the exact preceding Cloud Run revision. The pre-candidate workflow +leaves that revision untouched. + +Leave successfully published v2 catalog rows and the v2 schema in place. They +are dormant, additive, and safe for an idempotent retry. Do not delete catalog +rows, reset the database, or downgrade Alembic as part of application rollback. +A schema downgrade is a separate reviewed operator operation against the +reconfirmed v2 target. Because every Cloud Run revision selects its own +installed PolicyEngine.py version, a preceding artifact reads its existing +snapshot without a database-wide current-version setting. If that snapshot is +absent, run that artifact's compatible explicit initializer before making its +candidate eligible for traffic. diff --git a/migrations/v2/versions/68b4a5ae5dc5_version_metadata_catalog_snapshots.py b/migrations/v2/versions/68b4a5ae5dc5_version_metadata_catalog_snapshots.py new file mode 100644 index 000000000..2bcdcc453 --- /dev/null +++ b/migrations/v2/versions/68b4a5ae5dc5_version_metadata_catalog_snapshots.py @@ -0,0 +1,218 @@ +"""version metadata catalog snapshots + +Revision ID: 68b4a5ae5dc5 +Revises: f5ef4347cb2a +Create Date: 2026-08-21 02:13:03.270478 +Generation: uv run alembic -c alembic-v2.ini revision --autogenerate +""" + +from typing import Sequence, Union + +from alembic import op +import sqlalchemy as sa +import sqlmodel + + +revision: str = "68b4a5ae5dc5" +down_revision: Union[str, None] = "f5ef4347cb2a" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.add_column( + "datasets", sa.Column("tax_benefit_model_version_id", sa.Uuid(), nullable=False) + ) + op.add_column( + "regions", sa.Column("tax_benefit_model_version_id", sa.Uuid(), nullable=False) + ) + op.drop_constraint( + op.f("fk_regions_default_dataset_model_datasets"), "regions", type_="foreignkey" + ) + op.drop_index(op.f("ix_datasets_tax_benefit_model_id"), table_name="datasets") + op.drop_constraint(op.f("uq_datasets_id_model"), "datasets", type_="unique") + op.drop_constraint(op.f("uq_datasets_model_name"), "datasets", type_="unique") + op.create_index( + op.f("ix_datasets_tax_benefit_model_version_id"), + "datasets", + ["tax_benefit_model_version_id"], + unique=False, + ) + op.create_unique_constraint( + "uq_datasets_id_model_version", + "datasets", + ["id", "tax_benefit_model_version_id"], + ) + op.create_unique_constraint( + "uq_datasets_model_version_name", + "datasets", + ["tax_benefit_model_version_id", "name"], + ) + op.drop_constraint( + op.f("fk_datasets_tax_benefit_model_id_tax_benefit_models"), + "datasets", + type_="foreignkey", + ) + op.create_foreign_key( + op.f("fk_datasets_tax_benefit_model_version_id_tax_benefit_model_versions"), + "datasets", + "tax_benefit_model_versions", + ["tax_benefit_model_version_id"], + ["id"], + ondelete="RESTRICT", + ) + op.drop_column("datasets", "tax_benefit_model_id") + op.create_index( + "uq_parameter_values_canonical_parameter_start_date", + "parameter_values", + ["parameter_id", "start_date"], + unique=True, + postgresql_where=sa.text("policy_id IS NULL AND dynamic_id IS NULL"), + sqlite_where=sa.text("policy_id IS NULL AND dynamic_id IS NULL"), + ) + op.drop_index(op.f("ix_regions_tax_benefit_model_id"), table_name="regions") + op.drop_constraint(op.f("uq_regions_model_code"), "regions", type_="unique") + op.create_index( + op.f("ix_regions_tax_benefit_model_version_id"), + "regions", + ["tax_benefit_model_version_id"], + unique=False, + ) + op.create_unique_constraint( + "uq_regions_model_version_code", + "regions", + ["tax_benefit_model_version_id", "code"], + ) + op.drop_constraint( + op.f("fk_regions_tax_benefit_model_id_tax_benefit_models"), + "regions", + type_="foreignkey", + ) + op.create_foreign_key( + "fk_regions_default_dataset_model_version", + "regions", + "datasets", + ["default_dataset_id", "tax_benefit_model_version_id"], + ["id", "tax_benefit_model_version_id"], + ondelete="RESTRICT", + ) + op.create_foreign_key( + op.f("fk_regions_tax_benefit_model_version_id_tax_benefit_model_versions"), + "regions", + "tax_benefit_model_versions", + ["tax_benefit_model_version_id"], + ["id"], + ondelete="RESTRICT", + ) + op.drop_column("regions", "tax_benefit_model_id") + op.add_column( + "tax_benefit_model_versions", + sa.Column("current_law_id", sa.Integer(), nullable=False), + ) + op.add_column( + "tax_benefit_model_versions", + sa.Column("metadata_time_periods", sa.JSON(), nullable=False), + ) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_column("tax_benefit_model_versions", "metadata_time_periods") + op.drop_column("tax_benefit_model_versions", "current_law_id") + op.add_column( + "datasets", + sa.Column( + "tax_benefit_model_id", sa.UUID(), autoincrement=False, nullable=False + ), + ) + op.add_column( + "regions", + sa.Column( + "tax_benefit_model_id", sa.UUID(), autoincrement=False, nullable=False + ), + ) + op.drop_constraint( + op.f("fk_regions_tax_benefit_model_version_id_tax_benefit_model_versions"), + "regions", + type_="foreignkey", + ) + op.drop_constraint( + "fk_regions_default_dataset_model_version", "regions", type_="foreignkey" + ) + op.drop_constraint("uq_regions_model_version_code", "regions", type_="unique") + op.drop_index(op.f("ix_regions_tax_benefit_model_version_id"), table_name="regions") + op.drop_constraint( + op.f("fk_datasets_tax_benefit_model_version_id_tax_benefit_model_versions"), + "datasets", + type_="foreignkey", + ) + op.drop_constraint("uq_datasets_model_version_name", "datasets", type_="unique") + op.drop_constraint("uq_datasets_id_model_version", "datasets", type_="unique") + op.drop_index( + op.f("ix_datasets_tax_benefit_model_version_id"), table_name="datasets" + ) + op.create_foreign_key( + op.f("fk_datasets_tax_benefit_model_id_tax_benefit_models"), + "datasets", + "tax_benefit_models", + ["tax_benefit_model_id"], + ["id"], + ondelete="RESTRICT", + ) + op.create_unique_constraint( + op.f("uq_datasets_model_name"), + "datasets", + ["tax_benefit_model_id", "name"], + postgresql_nulls_not_distinct=False, + ) + op.create_unique_constraint( + op.f("uq_datasets_id_model"), + "datasets", + ["id", "tax_benefit_model_id"], + postgresql_nulls_not_distinct=False, + ) + op.create_index( + op.f("ix_datasets_tax_benefit_model_id"), + "datasets", + ["tax_benefit_model_id"], + unique=False, + ) + op.create_foreign_key( + op.f("fk_regions_tax_benefit_model_id_tax_benefit_models"), + "regions", + "tax_benefit_models", + ["tax_benefit_model_id"], + ["id"], + ondelete="RESTRICT", + ) + op.create_foreign_key( + op.f("fk_regions_default_dataset_model_datasets"), + "regions", + "datasets", + ["default_dataset_id", "tax_benefit_model_id"], + ["id", "tax_benefit_model_id"], + ondelete="RESTRICT", + ) + op.create_unique_constraint( + op.f("uq_regions_model_code"), + "regions", + ["tax_benefit_model_id", "code"], + postgresql_nulls_not_distinct=False, + ) + op.create_index( + op.f("ix_regions_tax_benefit_model_id"), + "regions", + ["tax_benefit_model_id"], + unique=False, + ) + op.drop_column("regions", "tax_benefit_model_version_id") + op.drop_column("datasets", "tax_benefit_model_version_id") + op.drop_index( + "uq_parameter_values_canonical_parameter_start_date", + table_name="parameter_values", + postgresql_where=sa.text("policy_id IS NULL AND dynamic_id IS NULL"), + sqlite_where=sa.text("policy_id IS NULL AND dynamic_id IS NULL"), + ) + # ### end Alembic commands ### diff --git a/policyengine_api/asgi_factory.py b/policyengine_api/asgi_factory.py index c2929d4a6..89d96143b 100644 --- a/policyengine_api/asgi_factory.py +++ b/policyengine_api/asgi_factory.py @@ -17,6 +17,7 @@ from policyengine_api.fastapi_routes.specification import ( build_specification_router, ) +from policyengine_api.fastapi_routes.v2_metadata import build_v2_metadata_router from policyengine_api.migration_flags import ( BACKEND_RESPONSE_HEADER, RouteImplementation, @@ -145,6 +146,7 @@ def log_native_route(status_code: int) -> None: _asgi_request_id.reset(context_token) app.include_router(build_core_health_router(dependencies)) + app.include_router(build_v2_metadata_router(dependencies)) if route_settings.health is RouteImplementation.FASTAPI_NATIVE: app.include_router(build_readiness_router(dependencies)) if route_settings.specification is RouteImplementation.FASTAPI_NATIVE: diff --git a/policyengine_api/data/v2/catalog/__init__.py b/policyengine_api/data/v2/catalog/__init__.py new file mode 100644 index 000000000..a0d453d12 --- /dev/null +++ b/policyengine_api/data/v2/catalog/__init__.py @@ -0,0 +1,15 @@ +"""PolicyEngine.py-derived v2 reference catalog extraction and publication.""" + +from policyengine_api.data.v2.catalog.extraction import ( + CatalogExtractionError, + extract_catalog, + extract_installed_catalog, +) +from policyengine_api.data.v2.catalog.records import NormalizedCatalog + +__all__ = [ + "CatalogExtractionError", + "NormalizedCatalog", + "extract_catalog", + "extract_installed_catalog", +] diff --git a/policyengine_api/data/v2/catalog/extraction.py b/policyengine_api/data/v2/catalog/extraction.py new file mode 100644 index 000000000..fa82dcb0a --- /dev/null +++ b/policyengine_api/data/v2/catalog/extraction.py @@ -0,0 +1,712 @@ +"""Normalize the installed PolicyEngine.py public model catalog.""" + +from __future__ import annotations + +from collections import Counter +from collections.abc import Callable, Mapping, Sequence +from datetime import date, datetime, timedelta, timezone +from enum import Enum +from importlib import metadata as importlib_metadata +import math +import re +from typing import Any +from urllib.parse import urlsplit +from uuid import NAMESPACE_URL, UUID, uuid5 + +from policyengine_api.data.v2.catalog.records import ( + CountryCatalog, + DatasetRecord, + FallbackSummary, + ModelRecord, + ModelVersionRecord, + NormalizedCatalog, + ParameterNodeRecord, + ParameterRecord, + ParameterValueRecord, + RegionRecord, + VariableRecord, +) + + +SUPPORTED_COUNTRIES = ("us", "uk") +REQUIRED_DEPENDENCIES = ( + "policyengine-core", + "policyengine-us", + "policyengine-uk", +) +REVIEWED_DEFAULT_DATASETS = { + "us": "populace_us_2024", + "uk": "enhanced_frs_2024_25", +} +CURRENT_LAW_IDS = {"us": 2, "uk": 1} +METADATA_TIME_PERIODS = { + "us": tuple(range(2035, 2021, -1)), + "uk": tuple(range(2024, 2031)), +} +SUPPORTED_REGION_TYPES = frozenset( + { + "national", + "country", + "state", + "congressional_district", + "constituency", + "local_authority", + "city", + "place", + } +) +PLACEHOLDER_VERSIONS = frozenset({"", "0", "0.0", "0.0.0", "unknown"}) +DATASET_YEAR_PATTERN = re.compile(r"(?:19|20)\d{2}") + + +class CatalogExtractionError(RuntimeError): + """Raised before database access when source metadata cannot be normalized.""" + + +def _identifier(*parts: object) -> UUID: + value = "/".join(str(part) for part in parts) + return uuid5(NAMESPACE_URL, f"https://api.policyengine.org/v2/catalog/{value}") + + +def _required_text(value: object, *, field_name: str, maximum: int) -> str: + if not isinstance(value, str) or not value.strip(): + raise CatalogExtractionError(f"{field_name} must be a non-empty string") + normalized = value.strip() + if len(normalized) > maximum: + raise CatalogExtractionError(f"{field_name} exceeds {maximum} characters") + return normalized + + +def _optional_text( + value: object, + *, + field_name: str, + maximum: int | None = None, +) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise CatalogExtractionError(f"{field_name} must be a string or null") + if maximum is not None and len(value) > maximum: + raise CatalogExtractionError(f"{field_name} exceeds {maximum} characters") + return value + + +def _type_name(value: object) -> str | None: + if value is None: + return None + if isinstance(value, str): + return value + name = getattr(value, "__name__", None) + if not isinstance(name, str) or not name: + raise CatalogExtractionError(f"unsupported data type {value!r}") + return name + + +def normalize_json_value(value: object) -> Any: + """Convert supported public values into strict JSON-compatible values.""" + + if value is None or isinstance(value, (str, bool, int)): + return value + if isinstance(value, float): + if math.isnan(value): + raise CatalogExtractionError("NaN is not JSON-compatible") + if math.isinf(value): + return "Infinity" if value > 0 else "-Infinity" + return value + if isinstance(value, Enum): + return normalize_json_value(value.value) + if isinstance(value, (date, datetime)): + return value.isoformat() + if isinstance(value, Mapping): + normalized: dict[str, Any] = {} + for key, nested in value.items(): + if not isinstance(key, str): + raise CatalogExtractionError("JSON object keys must be strings") + normalized[key] = normalize_json_value(nested) + return normalized + if isinstance(value, (list, tuple)): + return [normalize_json_value(item) for item in value] + + item = getattr(value, "item", None) + if callable(item): + scalar = item() + if scalar is not value: + return normalize_json_value(scalar) + raise CatalogExtractionError( + f"unsupported JSON value type {type(value).__module__}.{type(value).__name__}" + ) + + +def _aware_datetime(value: object, *, field_name: str) -> datetime: + if isinstance(value, date) and not isinstance(value, datetime): + value = datetime(value.year, value.month, value.day) + if not isinstance(value, datetime): + raise CatalogExtractionError(f"{field_name} must be a date or datetime") + if value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value.astimezone(timezone.utc) + + +def _optional_aware_datetime(value: object, *, field_name: str) -> datetime | None: + if value is None: + return None + return _aware_datetime(value, field_name=field_name) + + +def _dataset_name_from_path(path: str) -> str: + without_revision = path.split("@", maxsplit=1)[0] + parsed_path = urlsplit(without_revision).path or without_revision + filename = parsed_path.rstrip("/").rsplit("/", maxsplit=1)[-1] + for suffix in (".hdf5", ".parquet", ".csv", ".h5"): + if filename.endswith(suffix): + filename = filename[: -len(suffix)] + break + return _required_text(filename, field_name="dataset name", maximum=255) + + +def _dataset_year(name: str) -> int: + matches = DATASET_YEAR_PATTERN.findall(name) + if not matches: + raise CatalogExtractionError(f"dataset {name!r} has no four-digit year") + match = list(DATASET_YEAR_PATTERN.finditer(name))[-1] + return int(match.group()) + + +def _verify_bundle( + *, + bundle: Mapping[str, Any], + policyengine_version: str, + installed_version: Callable[[str], str], +) -> tuple[tuple[str, str], ...]: + if ( + not isinstance(policyengine_version, str) + or policyengine_version.strip().lower() in PLACEHOLDER_VERSIONS + ): + raise CatalogExtractionError("PolicyEngine.py version is a placeholder") + canonical_version = _required_text( + policyengine_version, + field_name="PolicyEngine.py version", + maximum=128, + ) + manifest_version = bundle.get("policyengine_version") or bundle.get( + "bundle_version" + ) + if manifest_version != canonical_version: + raise CatalogExtractionError( + "installed PolicyEngine.py version does not match its packaged manifest" + ) + + packages = bundle.get("packages") + if not isinstance(packages, Mapping): + raise CatalogExtractionError("PolicyEngine.py manifest has no package map") + + observed_dependencies: list[tuple[str, str]] = [] + for package_name in REQUIRED_DEPENDENCIES: + package = packages.get(package_name) + if not isinstance(package, Mapping): + raise CatalogExtractionError( + f"PolicyEngine.py manifest omits {package_name}" + ) + expected = package.get("version") + if not isinstance(expected, str) or expected.lower() in PLACEHOLDER_VERSIONS: + raise CatalogExtractionError( + f"PolicyEngine.py manifest has no valid {package_name} version" + ) + try: + observed = installed_version(package_name) + except importlib_metadata.PackageNotFoundError as error: + raise CatalogExtractionError( + f"required installed distribution {package_name} is absent" + ) from error + if observed != expected: + raise CatalogExtractionError( + f"installed {package_name} version does not match " + "the PolicyEngine.py manifest" + ) + observed_dependencies.append((package_name, expected)) + return tuple(observed_dependencies) + + +def _normalize_variable( + source: object, + *, + model_version_id: UUID, +) -> VariableRecord: + name = _required_text( + getattr(source, "name", None), field_name="variable name", maximum=512 + ) + possible_values = getattr(source, "possible_values", None) + if possible_values is not None: + if not isinstance(possible_values, Sequence) or isinstance( + possible_values, (str, bytes) + ): + raise CatalogExtractionError( + f"variable {name!r} possible_values must be a sequence" + ) + normalized_possible_values = [ + _required_text( + value, field_name=f"variable {name} possible value", maximum=255 + ) + for value in possible_values + ] + else: + normalized_possible_values = None + + def optional_names(attribute: str) -> list[str] | None: + values = getattr(source, attribute, None) + if values is None: + return None + if not isinstance(values, Sequence) or isinstance(values, (str, bytes)): + raise CatalogExtractionError( + f"variable {name!r} {attribute} must be a sequence" + ) + return [ + _required_text( + value, + field_name=f"variable {name} {attribute} item", + maximum=512, + ) + for value in values + ] + + return VariableRecord( + id=_identifier("variable", model_version_id, name), + model_version_id=model_version_id, + name=name, + label=_optional_text( + getattr(source, "label", None), + field_name=f"variable {name} label", + maximum=512, + ), + entity=_required_text( + getattr(source, "entity", None), + field_name=f"variable {name} entity", + maximum=128, + ), + description=_optional_text( + getattr(source, "description", None), + field_name=f"variable {name} description", + ), + data_type=_type_name(getattr(source, "data_type", None)), + possible_values=normalized_possible_values, + default_value=normalize_json_value(getattr(source, "default_value", None)), + adds=optional_names("adds"), + subtracts=optional_names("subtracts"), + ) + + +def _normalize_parameter( + source: object, + *, + model_version_id: UUID, +) -> ParameterRecord: + name = _required_text( + getattr(source, "name", None), + field_name="parameter name", + maximum=512, + ) + parameter_id = _identifier("parameter", model_version_id, name) + source_values: list[tuple[datetime, datetime | None, Any]] = [] + seen_starts: set[datetime] = set() + previous_start: datetime | None = None + for source_value in getattr(source, "parameter_values", ()): + start_date = _aware_datetime( + getattr(source_value, "start_date", None), + field_name=f"parameter {name} value start_date", + ) + end_date = _optional_aware_datetime( + getattr(source_value, "end_date", None), + field_name=f"parameter {name} value end_date", + ) + try: + value_json = normalize_json_value(getattr(source_value, "value", None)) + except CatalogExtractionError as error: + raise CatalogExtractionError( + f"parameter {name!r} has an unsupported JSON value: {error}" + ) from error + if start_date in seen_starts: + raise CatalogExtractionError( + f"parameter {name!r} has a duplicate value start date" + ) + if previous_start is not None and start_date <= previous_start: + raise CatalogExtractionError( + f"parameter {name!r} values are not ordered oldest to newest" + ) + seen_starts.add(start_date) + previous_start = start_date + source_values.append((start_date, end_date, value_json)) + + normalized_values: list[ParameterValueRecord] = [] + for index, (start_date, supplied_end, value_json) in enumerate(source_values): + expected_end = ( + source_values[index + 1][0] - timedelta(days=1) + if index + 1 < len(source_values) + else None + ) + if supplied_end != expected_end: + raise CatalogExtractionError( + f"parameter {name!r} does not expose canonical inclusive intervals" + ) + normalized_values.append( + ParameterValueRecord( + id=_identifier("parameter-value", parameter_id, start_date.isoformat()), + parameter_id=parameter_id, + value_json=value_json, + start_date=start_date, + end_date=expected_end, + ) + ) + + return ParameterRecord( + id=parameter_id, + model_version_id=model_version_id, + name=name, + label=_optional_text( + getattr(source, "label", None), + field_name=f"parameter {name} label", + maximum=512, + ), + description=_optional_text( + getattr(source, "description", None), + field_name=f"parameter {name} description", + ), + data_type=_type_name(getattr(source, "data_type", None)), + unit=_optional_text( + getattr(source, "unit", None), + field_name=f"parameter {name} unit", + maximum=128, + ), + values=tuple(normalized_values), + ) + + +def _public_named_records( + source: object, + *, + attribute: str, + entity_name: str, +) -> tuple[object, ...]: + records = getattr(source, attribute, None) + if not isinstance(records, Mapping): + raise CatalogExtractionError( + f"public model {attribute} must be a name-indexed mapping" + ) + + normalized: list[object] = [] + for key, record in sorted(records.items()): + name = getattr(record, "name", None) + if key != name: + raise CatalogExtractionError( + f"{entity_name} mapping key {key!r} does not match record name {name!r}" + ) + normalized.append(record) + return tuple(normalized) + + +def _normalize_country( + *, + country_id: str, + source: object, + policyengine_version: str, + expected_country_package_version: str, +) -> CountryCatalog: + source_model = getattr(source, "model", None) + model_name = _required_text( + getattr(source_model, "id", None), + field_name=f"{country_id} model name", + maximum=32, + ) + model_id = _identifier("model", model_name) + model_version_id = _identifier("model-version", model_id, policyengine_version) + + model_package = getattr(source, "model_package", None) + if getattr(model_package, "version", None) != expected_country_package_version: + raise CatalogExtractionError( + f"{country_id} public model does not match the PolicyEngine.py manifest" + ) + release_manifest = getattr(source, "release_manifest", None) + if getattr(release_manifest, "country_id", None) != country_id: + raise CatalogExtractionError( + f"{country_id} public model has an inconsistent release manifest" + ) + if getattr(release_manifest, "policyengine_version", None) != policyengine_version: + raise CatalogExtractionError( + f"{country_id} public model was not certified by the installed " + "PolicyEngine.py version" + ) + + variables = tuple( + _normalize_variable(variable, model_version_id=model_version_id) + for variable in _public_named_records( + source, + attribute="variables_by_name", + entity_name="variable", + ) + ) + normalized_nodes: list[ParameterNodeRecord] = [] + for node in _public_named_records( + source, + attribute="parameter_nodes_by_name", + entity_name="parameter node", + ): + source_name = getattr(node, "name", None) + # The public UK model includes one unnamed structural root. It has no + # addressable identity and contains no metadata of its own. + if source_name == "": + continue + node_name = _required_text( + source_name, + field_name="parameter node name", + maximum=512, + ) + normalized_nodes.append( + ParameterNodeRecord( + id=_identifier("parameter-node", model_version_id, node_name), + model_version_id=model_version_id, + name=node_name, + label=_optional_text( + getattr(node, "label", None), + field_name="parameter node label", + maximum=512, + ), + description=_optional_text( + getattr(node, "description", None), + field_name="parameter node description", + ), + ) + ) + parameter_nodes = tuple(normalized_nodes) + parameters = tuple( + _normalize_parameter(parameter, model_version_id=model_version_id) + for parameter in _public_named_records( + source, + attribute="parameters_by_name", + entity_name="parameter", + ) + ) + + reviewed_default = REVIEWED_DEFAULT_DATASETS[country_id] + manifest_default = getattr(release_manifest, "default_dataset", None) + if manifest_default != reviewed_default: + raise CatalogExtractionError( + f"{country_id} default dataset does not match the reviewed selection" + ) + default_dataset_uri = getattr(release_manifest, "default_dataset_uri", None) + registry = getattr(source, "region_registry", None) + if registry is None or getattr(registry, "country_id", None) != country_id: + raise CatalogExtractionError(f"{country_id} public region registry is absent") + + source_regions = tuple(registry) + region_codes = [getattr(region, "code", None) for region in source_regions] + duplicate_region_codes = sorted( + code for code, count in Counter(region_codes).items() if count > 1 + ) + if duplicate_region_codes: + raise CatalogExtractionError( + f"duplicate region natural keys: {duplicate_region_codes[:3]}" + ) + if country_id not in region_codes: + raise CatalogExtractionError(f"{country_id} national region is absent") + + selected_dataset_names: set[str] = {reviewed_default} + region_dataset_names: dict[str, str] = {} + fallback_counts: Counter[str] = Counter() + for region in source_regions: + code = _required_text( + getattr(region, "code", None), + field_name="region code", + maximum=255, + ) + region_type = _required_text( + getattr(region, "region_type", None), + field_name=f"region {code} type", + maximum=64, + ) + if region_type not in SUPPORTED_REGION_TYPES: + raise CatalogExtractionError( + f"region {code!r} has unsupported type {region_type!r}" + ) + dataset_path = getattr(region, "dataset_path", None) + if country_id == "us" and region_type != "national" and dataset_path: + dataset_name = _dataset_name_from_path(dataset_path) + selected_dataset_names.add(dataset_name) + else: + dataset_name = reviewed_default + if country_id == "us" and region_type != "national" and not dataset_path: + fallback_counts[region_type] += 1 + if region_type == "national" and dataset_path not in { + None, + default_dataset_uri, + }: + raise CatalogExtractionError( + f"{country_id} national region does not use the reviewed dataset" + ) + region_dataset_names[code] = dataset_name + + datasets = tuple( + DatasetRecord( + id=_identifier("dataset", model_version_id, name), + model_version_id=model_version_id, + name=name, + description=f"PolicyEngine.py logical input dataset {name}", + year=_dataset_year(name), + ) + for name in sorted(selected_dataset_names) + ) + dataset_ids = {dataset.name: dataset.id for dataset in datasets} + + regions: list[RegionRecord] = [] + for source_region in source_regions: + code = str(getattr(source_region, "code")) + strategy = getattr(source_region, "scoping_strategy", None) + requires_filter = bool(getattr(source_region, "requires_filter", False)) + if requires_filter: + strategy_type = _required_text( + getattr(strategy, "strategy_type", None), + field_name=f"region {code} filter strategy", + maximum=64, + ) + filter_field = _required_text( + getattr(strategy, "variable_name", None), + field_name=f"region {code} filter field", + maximum=128, + ) + filter_value_source = getattr(strategy, "variable_value", None) + if not isinstance(filter_value_source, (str, int, float, bool)): + raise CatalogExtractionError( + f"region {code!r} has an unsupported filter value" + ) + filter_value = str(filter_value_source) + additional_filters = getattr(strategy, "additional_filters", {}) + if additional_filters: + raise CatalogExtractionError( + f"region {code!r} has unsupported additional filters" + ) + else: + strategy_type = None + filter_field = None + filter_value = None + + regions.append( + RegionRecord( + id=_identifier("region", model_version_id, code), + model_version_id=model_version_id, + default_dataset_id=dataset_ids[region_dataset_names[code]], + code=code, + label=_required_text( + getattr(source_region, "label", None), + field_name=f"region {code} label", + maximum=255, + ), + region_type=str(getattr(source_region, "region_type")), + requires_filter=requires_filter, + filter_field=filter_field, + filter_value=filter_value, + filter_strategy=strategy_type, + parent_code=_optional_text( + getattr(source_region, "parent_code", None), + field_name=f"region {code} parent_code", + maximum=255, + ), + state_code=_optional_text( + getattr(source_region, "state_code", None), + field_name=f"region {code} state_code", + maximum=16, + ), + state_name=_optional_text( + getattr(source_region, "state_name", None), + field_name=f"region {code} state_name", + maximum=128, + ), + ) + ) + + return CountryCatalog( + country_id=country_id, + model=ModelRecord( + id=model_id, + country_id=country_id, + name=model_name, + description=_optional_text( + getattr(source_model, "description", None), + field_name=f"{country_id} model description", + ), + ), + model_version=ModelVersionRecord( + id=model_version_id, + model_id=model_id, + version=policyengine_version, + description=_optional_text( + getattr(source_model, "description", None), + field_name=f"{country_id} model-version description", + ), + current_law_id=CURRENT_LAW_IDS[country_id], + metadata_time_periods=METADATA_TIME_PERIODS[country_id], + ), + variables=variables, + parameter_nodes=parameter_nodes, + parameters=parameters, + datasets=datasets, + regions=tuple(sorted(regions, key=lambda record: record.code)), + fallback_summaries=tuple( + FallbackSummary(region_type=region_type, count=count) + for region_type, count in sorted(fallback_counts.items()) + ), + ) + + +def extract_catalog( + *, + bundle: Mapping[str, Any], + policyengine_version: str, + models: Mapping[str, object], + installed_version: Callable[[str], str] = importlib_metadata.version, +) -> NormalizedCatalog: + """Validate and normalize supplied PolicyEngine.py public objects.""" + + dependencies = _verify_bundle( + bundle=bundle, + policyengine_version=policyengine_version, + installed_version=installed_version, + ) + missing = [country for country in SUPPORTED_COUNTRIES if country not in models] + if missing: + raise CatalogExtractionError( + f"PolicyEngine.py public models are absent for: {', '.join(missing)}" + ) + + package_map = bundle["packages"] + countries = tuple( + _normalize_country( + country_id=country_id, + source=models[country_id], + policyengine_version=policyengine_version, + expected_country_package_version=package_map[f"policyengine-{country_id}"][ + "version" + ], + ) + for country_id in SUPPORTED_COUNTRIES + ) + return NormalizedCatalog( + policyengine_version=policyengine_version, + dependency_versions=dependencies, + countries=countries, + ) + + +def extract_installed_catalog() -> NormalizedCatalog: + """Load and normalize the installed PolicyEngine.py certified catalog.""" + + import policyengine + from policyengine.bundle import get_current_bundle + + policyengine_version = importlib_metadata.version("policyengine") + return extract_catalog( + bundle=get_current_bundle(), + policyengine_version=policyengine_version, + models={ + "us": policyengine.us.model, + "uk": policyengine.uk.model, + }, + ) diff --git a/policyengine_api/data/v2/catalog/initialization.py b/policyengine_api/data/v2/catalog/initialization.py new file mode 100644 index 000000000..873ed6a28 --- /dev/null +++ b/policyengine_api/data/v2/catalog/initialization.py @@ -0,0 +1,88 @@ +"""Explicit command boundary for one-time v2 catalog initialization.""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +import json +import sys + +from sqlalchemy import Engine, create_engine +from sqlalchemy.pool import NullPool + +from policyengine_api.data.v2.catalog.extraction import ( + CatalogExtractionError, + extract_installed_catalog, +) +from policyengine_api.data.v2.catalog.publication import ( + CatalogPublicationError, + PublicationEvidence, + publish_catalog, +) +from policyengine_api.data.v2.catalog.records import NormalizedCatalog +from policyengine_api.data.v2.settings import ( + V2ConfigurationError, + V2DatabaseSettings, + load_v2_data_write_database_settings, +) + + +def build_data_write_engine(settings: V2DatabaseSettings) -> Engine: + """Build an isolated, unpooled engine for one initialization attempt.""" + + return create_engine(settings.connection.url, poolclass=NullPool) + + +def initialize_catalog( + environ: Mapping[str, str] | None = None, + *, + extractor: Callable[[], NormalizedCatalog] = extract_installed_catalog, + engine_builder: Callable[[V2DatabaseSettings], Engine] = build_data_write_engine, + publisher: Callable[[Engine, NormalizedCatalog], PublicationEvidence] = ( + publish_catalog + ), +) -> PublicationEvidence: + """Load only row-write settings, extract fully, then publish explicitly.""" + + settings = load_v2_data_write_database_settings(environ) + catalog = extractor() + engine = engine_builder(settings) + try: + return publisher(engine, catalog) + finally: + engine.dispose() + + +def _error_payload(error: Exception) -> dict[str, object]: + safe_errors = ( + V2ConfigurationError, + CatalogExtractionError, + CatalogPublicationError, + ) + message = ( + str(error) + if isinstance(error, safe_errors) + else "catalog initialization failed unexpectedly" + ) + return { + "outcome": "error", + "error": { + "type": type(error).__name__, + "message": message, + }, + } + + +def main() -> int: + """Run explicit initialization and return a shell-compatible status.""" + + try: + evidence = initialize_catalog() + except Exception as error: # noqa: BLE001 - command must return safe evidence + print(json.dumps(_error_payload(error), sort_keys=True), file=sys.stderr) + return 1 + print(json.dumps(evidence.as_dict(), sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/policyengine_api/data/v2/catalog/publication.py b/policyengine_api/data/v2/catalog/publication.py new file mode 100644 index 000000000..3ecf34937 --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication.py @@ -0,0 +1,932 @@ +"""Atomic PostgreSQL publication for a validated PolicyEngine.py catalog.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Iterator, Sequence +from dataclasses import dataclass +import logging +import time + +from psycopg.types.json import Jsonb +from sqlalchemy import Connection, Engine, text + +from policyengine_api.data.v2.catalog.records import ( + CountryCatalog, + NormalizedCatalog, + iter_batches, +) + + +EXPECTED_ALEMBIC_REVISION = "68b4a5ae5dc5" +PUBLICATION_ADVISORY_LOCK_KEY = 8_629_020_026_090_001 +COPY_BATCH_SIZE = 10_000 + +LOGGER = logging.getLogger(__name__) + + +class CatalogPublicationError(RuntimeError): + """Raised when publication cannot prove an atomic, complete result.""" + + +@dataclass(frozen=True, slots=True) +class PublicationEvidence: + """Non-secret facts emitted after a successful publication.""" + + policyengine_version: str + dependency_versions: tuple[tuple[str, str], ...] + entity_counts: dict[str, int] + fallback_summaries: tuple[tuple[str, str, int], ...] + elapsed_seconds: float + + def as_dict(self) -> dict[str, object]: + return { + "outcome": "ok", + "policyengine_version": self.policyengine_version, + "dependency_versions": dict(self.dependency_versions), + "entity_counts": self.entity_counts, + "fallback_summaries": [ + { + "country_id": country_id, + "region_type": region_type, + "count": count, + } + for country_id, region_type, count in self.fallback_summaries + ], + "elapsed_seconds": round(self.elapsed_seconds, 3), + } + + +TEMP_TABLE_STATEMENTS = ( + """ + CREATE TEMP TABLE stage_catalog_models ( + country_id text PRIMARY KEY, + id uuid NOT NULL, + name text NOT NULL, + description text, + version_id uuid NOT NULL, + version text NOT NULL, + version_description text, + current_law_id integer NOT NULL, + metadata_time_periods jsonb NOT NULL + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_variables ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + label text, + entity text NOT NULL, + description text, + data_type text, + possible_values jsonb, + default_value jsonb, + adds jsonb, + subtracts jsonb, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_parameter_nodes ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + label text, + description text, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_parameters ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + label text, + description text, + data_type text, + unit text, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_parameter_values ( + country_id text NOT NULL, + parameter_name text NOT NULL, + id uuid NOT NULL, + value_json jsonb NOT NULL, + start_date timestamptz NOT NULL, + end_date timestamptz, + PRIMARY KEY (country_id, parameter_name, start_date) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_datasets ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + description text, + year integer NOT NULL, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_regions ( + country_id text NOT NULL, + id uuid NOT NULL, + code text NOT NULL, + label text NOT NULL, + region_type text NOT NULL, + requires_filter boolean NOT NULL, + filter_field text, + filter_value text, + filter_strategy text, + parent_code text, + state_code text, + state_name text, + default_dataset_name text NOT NULL, + PRIMARY KEY (country_id, code) + ) ON COMMIT DROP + """, +) + + +COPY_COLUMNS = { + "stage_catalog_models": ( + "country_id", + "id", + "name", + "description", + "version_id", + "version", + "version_description", + "current_law_id", + "metadata_time_periods", + ), + "stage_catalog_variables": ( + "country_id", + "id", + "name", + "label", + "entity", + "description", + "data_type", + "possible_values", + "default_value", + "adds", + "subtracts", + ), + "stage_catalog_parameter_nodes": ( + "country_id", + "id", + "name", + "label", + "description", + ), + "stage_catalog_parameters": ( + "country_id", + "id", + "name", + "label", + "description", + "data_type", + "unit", + ), + "stage_catalog_parameter_values": ( + "country_id", + "parameter_name", + "id", + "value_json", + "start_date", + "end_date", + ), + "stage_catalog_datasets": ( + "country_id", + "id", + "name", + "description", + "year", + ), + "stage_catalog_regions": ( + "country_id", + "id", + "code", + "label", + "region_type", + "requires_filter", + "filter_field", + "filter_value", + "filter_strategy", + "parent_code", + "state_code", + "state_name", + "default_dataset_name", + ), +} + + +def _optional_json(value: object) -> Jsonb | None: + return None if value is None else Jsonb(value) + + +def _catalog_rows( + country: CountryCatalog, +) -> dict[str, Iterator[tuple[object, ...]]]: + dataset_names = {dataset.id: dataset.name for dataset in country.datasets} + + def model_rows() -> Iterator[tuple[object, ...]]: + yield ( + country.country_id, + country.model.id, + country.model.name, + country.model.description, + country.model_version.id, + country.model_version.version, + country.model_version.description, + country.model_version.current_law_id, + Jsonb(country.model_version.metadata_time_periods), + ) + + def variable_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.variables, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.label, + record.entity, + record.description, + record.data_type, + _optional_json(record.possible_values), + Jsonb(record.default_value), + _optional_json(record.adds), + _optional_json(record.subtracts), + ) + + def parameter_node_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches( + country.parameter_nodes, + batch_size=COPY_BATCH_SIZE, + ): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.label, + record.description, + ) + + def parameter_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.parameters, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.label, + record.description, + record.data_type, + record.unit, + ) + + def parameter_value_rows() -> Iterator[tuple[object, ...]]: + parameter_names = { + parameter.id: parameter.name for parameter in country.parameters + } + for batch in country.parameter_value_batches(batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + parameter_names[record.parameter_id], + record.id, + Jsonb(record.value_json), + record.start_date, + record.end_date, + ) + + def dataset_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.datasets, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.description, + record.year, + ) + + def region_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.regions, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.code, + record.label, + record.region_type, + record.requires_filter, + record.filter_field, + record.filter_value, + record.filter_strategy, + record.parent_code, + record.state_code, + record.state_name, + dataset_names[record.default_dataset_id], + ) + + return { + "stage_catalog_models": model_rows(), + "stage_catalog_variables": variable_rows(), + "stage_catalog_parameter_nodes": parameter_node_rows(), + "stage_catalog_parameters": parameter_rows(), + "stage_catalog_parameter_values": parameter_value_rows(), + "stage_catalog_datasets": dataset_rows(), + "stage_catalog_regions": region_rows(), + } + + +def _copy_rows( + connection: Connection, + *, + table_name: str, + columns: Sequence[str], + rows: Iterable[Sequence[object]], +) -> int: + """Write one bounded source stream through Psycopg COPY.""" + + raw_connection = connection.connection.driver_connection + statement = f"COPY {table_name} ({', '.join(columns)}) FROM STDIN" + count = 0 + with raw_connection.cursor() as cursor: + with cursor.copy(statement) as copy: + for row in rows: + copy.write_row(row) + count += 1 + return count + + +def _verify_expected_revision(connection: Connection) -> None: + if connection.dialect.name != "postgresql": + raise CatalogPublicationError("catalog publication requires PostgreSQL") + version_table = connection.execute( + text("SELECT to_regclass('public.alembic_version')") + ).scalar_one() + if version_table is None: + raise CatalogPublicationError("the v2 Alembic revision table is absent") + revisions = set( + connection.execute(text("SELECT version_num FROM alembic_version")).scalars() + ) + if revisions != {EXPECTED_ALEMBIC_REVISION}: + raise CatalogPublicationError( + "the v2 database is not at the expected Alembic revision" + ) + + +def _acquire_publication_lock(connection: Connection) -> None: + connection.execute( + text("SELECT pg_advisory_xact_lock(:lock_key)"), + {"lock_key": PUBLICATION_ADVISORY_LOCK_KEY}, + ).scalar_one() + + +def _create_staging_tables(connection: Connection) -> None: + for statement in TEMP_TABLE_STATEMENTS: + connection.execute(text(statement)) + + +def _stage_catalog( + connection: Connection, + catalog: NormalizedCatalog, + *, + checkpoint: Callable[[str, Connection], None] | None = None, +) -> dict[str, int]: + observed = {table_name: 0 for table_name in COPY_COLUMNS} + for country in catalog.countries: + for table_name, rows in _catalog_rows(country).items(): + observed[table_name] += _copy_rows( + connection, + table_name=table_name, + columns=COPY_COLUMNS[table_name], + rows=rows, + ) + if checkpoint is not None: + checkpoint("during_copy", connection) + expected = catalog.entity_counts() + expected_by_table = { + "stage_catalog_models": expected["models"], + "stage_catalog_variables": expected["variables"], + "stage_catalog_parameter_nodes": expected["parameter_nodes"], + "stage_catalog_parameters": expected["parameters"], + "stage_catalog_parameter_values": expected["parameter_values"], + "stage_catalog_datasets": expected["datasets"], + "stage_catalog_regions": expected["regions"], + } + if observed != expected_by_table: + raise CatalogPublicationError("COPY row counts differ from the catalog") + return observed + + +def _protected_row_counts(connection: Connection) -> tuple[int, int, int, int]: + return tuple( + connection.execute(text(f"SELECT count(*) FROM {table_name}")).scalar_one() + for table_name in ( + "dataset_versions", + "simulations", + "reports", + "report_runs", + ) + ) + + +def _version_exists(connection: Connection, country_id: str) -> bool: + return bool( + connection.execute( + text( + """ + SELECT EXISTS ( + SELECT 1 + FROM stage_catalog_models AS source + JOIN tax_benefit_models AS model + ON model.name = source.name + JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source.version + WHERE source.country_id = :country_id + ) + """ + ), + {"country_id": country_id}, + ).scalar_one() + ) + + +MODEL_DIFFERENCE_SQL = """ +SELECT EXISTS ( + SELECT 1 + FROM stage_catalog_models AS source + JOIN tax_benefit_models AS model ON model.name = source.name + JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source.version + WHERE source.country_id = :country_id + AND ( + model_version.description IS DISTINCT FROM source.version_description + OR model_version.current_law_id IS DISTINCT FROM source.current_law_id + OR model_version.metadata_time_periods::jsonb + IS DISTINCT FROM source.metadata_time_periods + ) +) +""" + + +VERSIONED_DIFFERENCE_SQL = ( + """ + WITH staged AS ( + SELECT name, label, entity, description, data_type, possible_values, + default_value, adds, subtracts + FROM stage_catalog_variables + WHERE country_id = :country_id + ), actual AS ( + SELECT variable.name, variable.label, variable.entity, + variable.description, variable.data_type, + variable.possible_values::jsonb, + variable.default_value::jsonb, + variable.adds::jsonb, variable.subtracts::jsonb + FROM variables AS variable + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = variable.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT name, label, description + FROM stage_catalog_parameter_nodes + WHERE country_id = :country_id + ), actual AS ( + SELECT node.name, node.label, node.description + FROM parameter_nodes AS node + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = node.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT name, label, description, data_type, unit + FROM stage_catalog_parameters + WHERE country_id = :country_id + ), actual AS ( + SELECT parameter.name, parameter.label, parameter.description, + parameter.data_type, parameter.unit + FROM parameters AS parameter + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = parameter.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT parameter_name, value_json, start_date, end_date + FROM stage_catalog_parameter_values + WHERE country_id = :country_id + ), actual AS ( + SELECT parameter.name, parameter_value.value_json::jsonb, + parameter_value.start_date, parameter_value.end_date + FROM parameter_values AS parameter_value + JOIN parameters AS parameter ON parameter.id = parameter_value.parameter_id + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = parameter.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + AND parameter_value.policy_id IS NULL + AND parameter_value.dynamic_id IS NULL + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, +) + + +VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL = ( + """ + WITH staged AS ( + SELECT name, description, year, false AS is_output_dataset, + NULL::text AS storage_path + FROM stage_catalog_datasets + WHERE country_id = :country_id + ), actual AS ( + SELECT dataset.name, dataset.description, dataset.year, + dataset.is_output_dataset, dataset.storage_path + FROM datasets AS dataset + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = dataset.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + AND NOT dataset.is_output_dataset + AND dataset.storage_path IS NULL + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT code, label, region_type, requires_filter, filter_field, + filter_value, filter_strategy, parent_code, state_code, + state_name, default_dataset_name + FROM stage_catalog_regions + WHERE country_id = :country_id + ), actual AS ( + SELECT region.code, region.label, region.region_type::text, + region.requires_filter, region.filter_field, + region.filter_value, region.filter_strategy, + region.parent_code, region.state_code, region.state_name, + dataset.name + FROM regions AS region + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = region.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN datasets AS dataset ON dataset.id = region.default_dataset_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, +) + + +def _assert_country_matches(connection: Connection, country_id: str) -> None: + statements = ( + MODEL_DIFFERENCE_SQL, + *VERSIONED_DIFFERENCE_SQL, + *VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL, + ) + for statement in statements: + differs = connection.execute( + text(statement), + {"country_id": country_id}, + ).scalar_one() + if differs: + raise CatalogPublicationError( + f"persisted {country_id} catalog differs from PolicyEngine.py" + ) + + +def _reject_dataset_role_conflicts( + connection: Connection, + country_id: str, +) -> None: + conflict = connection.execute( + text( + """ + SELECT EXISTS ( + SELECT 1 + FROM stage_catalog_datasets AS source_dataset + JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_dataset.country_id + JOIN tax_benefit_models AS model + ON model.name = source_model.name + JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version + JOIN datasets AS dataset + ON dataset.tax_benefit_model_version_id = model_version.id + AND dataset.name = source_dataset.name + WHERE source_dataset.country_id = :country_id + AND ( + dataset.is_output_dataset + OR dataset.storage_path IS NOT NULL + ) + ) + """ + ), + {"country_id": country_id}, + ).scalar_one() + if conflict: + raise CatalogPublicationError( + f"persisted {country_id} dataset identity is not an input dataset" + ) + + +INSERT_MODEL_SQL = """ +INSERT INTO tax_benefit_models ( + id, created_at, updated_at, name, description +) +SELECT id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, name, description +FROM stage_catalog_models +WHERE country_id = :country_id +ON CONFLICT (name) DO NOTHING +""" + +INSERT_MODEL_VERSION_SQL = """ +INSERT INTO tax_benefit_model_versions ( + id, created_at, model_id, version, description, current_law_id, + metadata_time_periods +) +SELECT source.version_id, CURRENT_TIMESTAMP, model.id, + source.version, source.version_description, source.current_law_id, + source.metadata_time_periods::json +FROM stage_catalog_models AS source +JOIN tax_benefit_models AS model ON model.name = source.name +WHERE source.country_id = :country_id +""" + +INSERT_DATASETS_SQL = """ +INSERT INTO datasets ( + id, created_at, updated_at, name, description, storage_path, year, + is_output_dataset, tax_benefit_model_version_id +) +SELECT source_dataset.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, + source_dataset.name, source_dataset.description, NULL, + source_dataset.year, false, model_version.id +FROM stage_catalog_datasets AS source_dataset +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_dataset.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_dataset.country_id = :country_id +""" + +INSERT_VARIABLES_SQL = """ +INSERT INTO variables ( + id, created_at, name, label, entity, description, data_type, + possible_values, default_value, adds, subtracts, + tax_benefit_model_version_id +) +SELECT source_variable.id, CURRENT_TIMESTAMP, source_variable.name, + source_variable.label, source_variable.entity, + source_variable.description, source_variable.data_type, + source_variable.possible_values::json, + source_variable.default_value::json, + source_variable.adds::json, source_variable.subtracts::json, + model_version.id +FROM stage_catalog_variables AS source_variable +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_variable.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_variable.country_id = :country_id +""" + +INSERT_PARAMETER_NODES_SQL = """ +INSERT INTO parameter_nodes ( + id, created_at, name, label, description, tax_benefit_model_version_id +) +SELECT source_node.id, CURRENT_TIMESTAMP, source_node.name, + source_node.label, source_node.description, model_version.id +FROM stage_catalog_parameter_nodes AS source_node +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_node.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_node.country_id = :country_id +""" + +INSERT_PARAMETERS_SQL = """ +INSERT INTO parameters ( + id, created_at, name, label, description, data_type, unit, + tax_benefit_model_version_id +) +SELECT source_parameter.id, CURRENT_TIMESTAMP, source_parameter.name, + source_parameter.label, source_parameter.description, + source_parameter.data_type, source_parameter.unit, model_version.id +FROM stage_catalog_parameters AS source_parameter +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_parameter.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_parameter.country_id = :country_id +""" + +INSERT_PARAMETER_VALUES_SQL = """ +INSERT INTO parameter_values ( + id, created_at, parameter_id, value_json, start_date, end_date, + policy_id, dynamic_id +) +SELECT source_value.id, CURRENT_TIMESTAMP, parameter.id, + source_value.value_json::json, source_value.start_date, + source_value.end_date, NULL, NULL +FROM stage_catalog_parameter_values AS source_value +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_value.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +JOIN parameters AS parameter + ON parameter.tax_benefit_model_version_id = model_version.id + AND parameter.name = source_value.parameter_name +WHERE source_value.country_id = :country_id +""" + +INSERT_REGIONS_SQL = """ +INSERT INTO regions ( + id, created_at, updated_at, code, label, region_type, requires_filter, + filter_field, filter_value, filter_strategy, parent_code, state_code, + state_name, tax_benefit_model_version_id, default_dataset_id +) +SELECT source_region.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, + source_region.code, source_region.label, + source_region.region_type::v2_region_type, + source_region.requires_filter, source_region.filter_field, + source_region.filter_value, source_region.filter_strategy, + source_region.parent_code, source_region.state_code, + source_region.state_name, model_version.id, dataset.id +FROM stage_catalog_regions AS source_region +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_region.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +JOIN datasets AS dataset + ON dataset.tax_benefit_model_version_id = model_version.id + AND dataset.name = source_region.default_dataset_name +WHERE source_region.country_id = :country_id +""" + + +SET_BASED_INSERT_SQL = ( + INSERT_MODEL_SQL, + INSERT_MODEL_VERSION_SQL, + INSERT_DATASETS_SQL, + INSERT_VARIABLES_SQL, + INSERT_PARAMETER_NODES_SQL, + INSERT_PARAMETERS_SQL, + INSERT_PARAMETER_VALUES_SQL, + INSERT_REGIONS_SQL, +) + + +def _publish_new_country(connection: Connection, country_id: str) -> None: + _reject_dataset_role_conflicts(connection, country_id) + for statement in SET_BASED_INSERT_SQL: + connection.execute(text(statement), {"country_id": country_id}) + + +def _assert_canonical_value_uniqueness(connection: Connection) -> None: + duplicates = connection.execute( + text( + """ + SELECT EXISTS ( + SELECT 1 + FROM parameter_values + WHERE policy_id IS NULL AND dynamic_id IS NULL + GROUP BY parameter_id, start_date + HAVING count(*) > 1 + ) + """ + ) + ).scalar_one() + if duplicates: + raise CatalogPublicationError( + "canonical parameter-value uniqueness validation failed" + ) + + +def publish_catalog( + engine: Engine, + catalog: NormalizedCatalog, + *, + checkpoint: Callable[[str, Connection], None] | None = None, +) -> PublicationEvidence: + """Publish one complete catalog atomically and return non-secret evidence.""" + + started = time.monotonic() + with engine.begin() as connection: + _verify_expected_revision(connection) + _acquire_publication_lock(connection) + before = _protected_row_counts(connection) + _create_staging_tables(connection) + _stage_catalog(connection, catalog, checkpoint=checkpoint) + if checkpoint is not None: + checkpoint("after_copy", connection) + + existing: set[str] = set() + for country in catalog.countries: + if _version_exists(connection, country.country_id): + _assert_country_matches(connection, country.country_id) + existing.add(country.country_id) + + for country in catalog.countries: + if country.country_id not in existing: + _publish_new_country(connection, country.country_id) + if checkpoint is not None: + checkpoint("after_reconciliation", connection) + + for country in catalog.countries: + _assert_country_matches(connection, country.country_id) + _assert_canonical_value_uniqueness(connection) + if _protected_row_counts(connection) != before: + raise CatalogPublicationError( + "publication changed simulation, report, or dataset-version rows" + ) + if checkpoint is not None: + checkpoint("after_validation", connection) + + fallback_summaries = tuple( + (country.country_id, summary.region_type, summary.count) + for country in catalog.countries + for summary in country.fallback_summaries + ) + _log_fallback_warning(fallback_summaries) + return PublicationEvidence( + policyengine_version=catalog.policyengine_version, + dependency_versions=catalog.dependency_versions, + entity_counts=catalog.entity_counts(), + fallback_summaries=fallback_summaries, + elapsed_seconds=time.monotonic() - started, + ) + + +def _log_fallback_warning( + fallback_summaries: tuple[tuple[str, str, int], ...], +) -> None: + if not fallback_summaries: + return + LOGGER.warning( + "PolicyEngine.py regional dataset fallback summary: %s", + fallback_summaries, + ) diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py new file mode 100644 index 000000000..bce8a4084 --- /dev/null +++ b/policyengine_api/data/v2/catalog/query.py @@ -0,0 +1,342 @@ +"""Read-only v2 catalog queries and deterministic response serialization.""" + +from __future__ import annotations + +from collections import defaultdict + +from packaging.version import InvalidVersion, Version +from sqlalchemy.exc import SQLAlchemyError +from sqlmodel import Session, select + +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDataset, + MetadataDatasetOption, + MetadataEconomyOptions, + MetadataModel, + MetadataModelVersion, + MetadataParameter, + MetadataParameterNode, + MetadataParameterValue, + MetadataRegion, + MetadataRegionOption, + MetadataResult, + MetadataTimePeriodOption, + MetadataVariable, +) +from policyengine_api.data.v2.models import ( + Dataset, + Parameter, + ParameterNode, + ParameterValue, + Region, + TaxBenefitModel, + TaxBenefitModelVersion, + Variable, +) + + +SUPPORTED_PREVIEW_COUNTRIES = frozenset({"us", "uk"}) + + +class MetadataCatalogUnavailableError(RuntimeError): + """Raised when a complete initialized catalog cannot be read.""" + + +class UnsupportedPreviewCountryError(ValueError): + """Raised when a country has no Stage 9 preview catalog.""" + + +class InvalidPolicyEngineVersionError(ValueError): + """Raised when an explicit version is not a canonical package version.""" + + +class MetadataCatalogVersionNotFoundError(LookupError): + """Raised when an explicitly selected catalog version is absent.""" + + +def validate_policyengine_version(value: str) -> str: + """Return one bounded canonical PEP 440 version string.""" + + if not isinstance(value, str) or not value or value != value.strip(): + raise InvalidPolicyEngineVersionError( + "policyengine_version must be a non-empty canonical version" + ) + if len(value) > 128: + raise InvalidPolicyEngineVersionError( + "policyengine_version must be at most 128 characters" + ) + try: + parsed = Version(value) + except InvalidVersion as error: + raise InvalidPolicyEngineVersionError( + "policyengine_version must be a canonical PEP 440 version" + ) from error + if str(parsed) != value or parsed == Version("0.0.0"): + raise InvalidPolicyEngineVersionError( + "policyengine_version must be a canonical non-placeholder version" + ) + return value + + +class V2MetadataQueryService: + """Assemble preview metadata using only an injected v2 read session.""" + + def __init__(self, session: Session, *, running_policyengine_version: str): + self._session = session + self._running_policyengine_version = validate_policyengine_version( + running_policyengine_version + ) + + def close(self) -> None: + """Close the request-owned read session.""" + + self._session.close() + + def get_metadata( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataResult: + if country_id not in SUPPORTED_PREVIEW_COUNTRIES: + raise UnsupportedPreviewCountryError(country_id) + explicit_version = policyengine_version is not None + selected_version = ( + validate_policyengine_version(policyengine_version) + if explicit_version + else self._running_policyengine_version + ) + try: + return self._read_metadata( + country_id, + selected_version, + explicit_version=explicit_version, + ) + except ( + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, + ): + raise + except SQLAlchemyError as error: + raise MetadataCatalogUnavailableError( + "the v2 metadata catalog cannot be queried" + ) from error + + def _read_metadata( + self, + country_id: str, + policyengine_version: str, + *, + explicit_version: bool, + ) -> MetadataResult: + model = self._session.exec( + select(TaxBenefitModel).where( + TaxBenefitModel.name == f"policyengine-{country_id}" + ) + ).one_or_none() + if model is None: + if explicit_version: + raise MetadataCatalogVersionNotFoundError( + f"PolicyEngine.py {policyengine_version} is not published " + f"for {country_id}" + ) + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 metadata catalog is not initialized" + ) + + model_version = self._session.exec( + select(TaxBenefitModelVersion).where( + TaxBenefitModelVersion.model_id == model.id, + TaxBenefitModelVersion.version == policyengine_version, + ) + ).one_or_none() + if model_version is None: + if explicit_version: + raise MetadataCatalogVersionNotFoundError( + f"PolicyEngine.py {policyengine_version} is not published " + f"for {country_id}" + ) + raise MetadataCatalogUnavailableError( + f"the running PolicyEngine.py {policyengine_version} catalog " + f"is absent for {country_id}" + ) + + variables = self._session.exec( + select(Variable) + .where(Variable.tax_benefit_model_version_id == model_version.id) + .order_by(Variable.name) + ).all() + nodes = self._session.exec( + select(ParameterNode) + .where(ParameterNode.tax_benefit_model_version_id == model_version.id) + .order_by(ParameterNode.name) + ).all() + parameters = self._session.exec( + select(Parameter) + .where(Parameter.tax_benefit_model_version_id == model_version.id) + .order_by(Parameter.name) + ).all() + parameter_values = self._session.exec( + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + Parameter.tax_benefit_model_version_id == model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ) + .order_by(Parameter.name, ParameterValue.start_date) + ).all() + regions = self._session.exec( + select(Region) + .where(Region.tax_benefit_model_version_id == model_version.id) + .order_by(Region.code) + ).all() + + if not variables or not nodes or not parameters or not regions: + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 metadata catalog is incomplete" + ) + + dataset_ids = {region.default_dataset_id for region in regions} + datasets = self._session.exec( + select(Dataset) + .where( + Dataset.id.in_(dataset_ids), + Dataset.tax_benefit_model_version_id == model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ) + .order_by(Dataset.name) + ).all() + if {dataset.id for dataset in datasets} != dataset_ids: + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 region datasets are incomplete" + ) + + values_by_parameter = defaultdict(list) + for value in parameter_values: + values_by_parameter[value.parameter_id].append( + MetadataParameterValue( + id=value.id, + value=value.value_json, + start_date=value.start_date, + end_date=value.end_date, + ) + ) + + national_region = next( + (region for region in regions if region.code == country_id), + None, + ) + if national_region is None: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 region is absent" + ) + datasets_by_id = {dataset.id: dataset for dataset in datasets} + national_dataset = datasets_by_id[national_region.default_dataset_id] + time_periods = model_version.metadata_time_periods + if ( + not isinstance(model_version.current_law_id, int) + or not isinstance(time_periods, list) + or not time_periods + or any(not isinstance(year, int) for year in time_periods) + ): + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 model-version options are incomplete" + ) + + return MetadataResult( + current_law_id=model_version.current_law_id, + model=MetadataModel( + id=model.id, + name=model.name, + description=model_version.description, + ), + model_version=MetadataModelVersion( + id=model_version.id, + model_id=model.id, + version=model_version.version, + description=model_version.description, + ), + variables=[ + MetadataVariable( + id=variable.id, + name=variable.name, + label=variable.label, + entity=variable.entity, + description=variable.description, + data_type=variable.data_type, + possible_values=variable.possible_values, + default_value=variable.default_value, + adds=variable.adds, + subtracts=variable.subtracts, + ) + for variable in variables + ], + parameter_nodes=[ + MetadataParameterNode( + id=node.id, + name=node.name, + label=node.label, + description=node.description, + ) + for node in nodes + ], + parameters=[ + MetadataParameter( + id=parameter.id, + name=parameter.name, + label=parameter.label, + description=parameter.description, + data_type=parameter.data_type, + unit=parameter.unit, + values=values_by_parameter[parameter.id], + ) + for parameter in parameters + ], + datasets=[ + MetadataDataset( + id=dataset.id, + name=dataset.name, + description=dataset.description, + year=dataset.year, + ) + for dataset in datasets + ], + regions=[ + MetadataRegion( + id=region.id, + code=region.code, + label=region.label, + region_type=region.region_type.value, + requires_filter=region.requires_filter, + filter_field=region.filter_field, + filter_value=region.filter_value, + filter_strategy=region.filter_strategy, + parent_code=region.parent_code, + state_code=region.state_code, + state_name=region.state_name, + default_dataset_id=region.default_dataset_id, + ) + for region in regions + ], + economy_options=MetadataEconomyOptions( + region=[ + MetadataRegionOption( + name=region.code, + label=region.label, + type=region.region_type.value, + ) + for region in regions + ], + time_period=[ + MetadataTimePeriodOption(name=year, label=str(year)) + for year in time_periods + ], + datasets=[ + MetadataDatasetOption( + name=national_dataset.name, + label=national_dataset.description or national_dataset.name, + ) + ], + ), + ) diff --git a/policyengine_api/data/v2/catalog/records.py b/policyengine_api/data/v2/catalog/records.py new file mode 100644 index 000000000..bead12892 --- /dev/null +++ b/policyengine_api/data/v2/catalog/records.py @@ -0,0 +1,200 @@ +"""Immutable database-independent records for the v2 metadata catalog.""" + +from __future__ import annotations + +from collections.abc import Iterator, Sequence +from dataclasses import dataclass, fields +from datetime import datetime +from typing import Any, TypeVar +from uuid import UUID + + +@dataclass(frozen=True, slots=True) +class ModelRecord: + id: UUID + country_id: str + name: str + description: str | None + + +@dataclass(frozen=True, slots=True) +class ModelVersionRecord: + id: UUID + model_id: UUID + version: str + description: str | None + current_law_id: int + metadata_time_periods: tuple[int, ...] + + +@dataclass(frozen=True, slots=True) +class VariableRecord: + id: UUID + model_version_id: UUID + name: str + label: str | None + entity: str + description: str | None + data_type: str | None + possible_values: list[str] | None + default_value: Any + adds: list[str] | None + subtracts: list[str] | None + + +@dataclass(frozen=True, slots=True) +class ParameterNodeRecord: + id: UUID + model_version_id: UUID + name: str + label: str | None + description: str | None + + +@dataclass(frozen=True, slots=True) +class ParameterValueRecord: + id: UUID + parameter_id: UUID + value_json: Any + start_date: datetime + end_date: datetime | None + + +@dataclass(frozen=True, slots=True) +class ParameterRecord: + id: UUID + model_version_id: UUID + name: str + label: str | None + description: str | None + data_type: str | None + unit: str | None + values: tuple[ParameterValueRecord, ...] + + +@dataclass(frozen=True, slots=True) +class DatasetRecord: + id: UUID + model_version_id: UUID + name: str + description: str | None + year: int + storage_path: None = None + is_output_dataset: bool = False + + +@dataclass(frozen=True, slots=True) +class RegionRecord: + id: UUID + model_version_id: UUID + default_dataset_id: UUID + code: str + label: str + region_type: str + requires_filter: bool + filter_field: str | None + filter_value: str | None + filter_strategy: str | None + parent_code: str | None + state_code: str | None + state_name: str | None + + +@dataclass(frozen=True, slots=True) +class FallbackSummary: + region_type: str + count: int + + +RecordT = TypeVar("RecordT") + + +def iter_batches( + records: Sequence[RecordT], + *, + batch_size: int, +) -> Iterator[tuple[RecordT, ...]]: + """Yield bounded immutable slices without constructing a flattened copy.""" + + if batch_size < 1: + raise ValueError("batch_size must be positive") + for start in range(0, len(records), batch_size): + yield tuple(records[start : start + batch_size]) + + +@dataclass(frozen=True, slots=True) +class CountryCatalog: + country_id: str + model: ModelRecord + model_version: ModelVersionRecord + variables: tuple[VariableRecord, ...] + parameter_nodes: tuple[ParameterNodeRecord, ...] + parameters: tuple[ParameterRecord, ...] + datasets: tuple[DatasetRecord, ...] + regions: tuple[RegionRecord, ...] + fallback_summaries: tuple[FallbackSummary, ...] + + def parameter_value_batches( + self, + *, + batch_size: int, + ) -> Iterator[tuple[ParameterValueRecord, ...]]: + """Yield canonical parameter values in bounded batches.""" + + if batch_size < 1: + raise ValueError("batch_size must be positive") + batch: list[ParameterValueRecord] = [] + for parameter in self.parameters: + for value in parameter.values: + batch.append(value) + if len(batch) == batch_size: + yield tuple(batch) + batch.clear() + if batch: + yield tuple(batch) + + def entity_counts(self) -> dict[str, int]: + """Return non-secret record counts for validation and deployment evidence.""" + + return { + "models": 1, + "model_versions": 1, + "variables": len(self.variables), + "parameter_nodes": len(self.parameter_nodes), + "parameters": len(self.parameters), + "parameter_values": sum( + len(parameter.values) for parameter in self.parameters + ), + "datasets": len(self.datasets), + "regions": len(self.regions), + } + + +@dataclass(frozen=True, slots=True) +class NormalizedCatalog: + policyengine_version: str + dependency_versions: tuple[tuple[str, str], ...] + countries: tuple[CountryCatalog, ...] + + def country(self, country_id: str) -> CountryCatalog: + """Return one supported country catalog.""" + + for catalog in self.countries: + if catalog.country_id == country_id: + return catalog + raise KeyError(country_id) + + def entity_counts(self) -> dict[str, int]: + """Return aggregate counts without catalog content.""" + + counts: dict[str, int] = {} + for country in self.countries: + for name, count in country.entity_counts().items(): + counts[name] = counts.get(name, 0) + count + return counts + + +def record_content(record: object) -> tuple[Any, ...]: + """Return a deterministic field-ordered representation for comparisons.""" + + return tuple(getattr(record, item.name) for item in fields(record)) diff --git a/policyengine_api/data/v2/catalog/schemas.py b/policyengine_api/data/v2/catalog/schemas.py new file mode 100644 index 000000000..e3bb198d0 --- /dev/null +++ b/policyengine_api/data/v2/catalog/schemas.py @@ -0,0 +1,151 @@ +"""Typed response schemas for the dormant v2 metadata preview.""" + +from __future__ import annotations + +from datetime import datetime +from enum import StrEnum +from typing import Annotated, Literal +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, Field, JsonValue, StringConstraints + + +class StrictResponseModel(BaseModel): + model_config = ConfigDict(extra="forbid", allow_inf_nan=False) + + +class MetadataRegionType(StrEnum): + NATIONAL = "national" + COUNTRY = "country" + STATE = "state" + CONGRESSIONAL_DISTRICT = "congressional_district" + CONSTITUENCY = "constituency" + LOCAL_AUTHORITY = "local_authority" + CITY = "city" + PLACE = "place" + + +class MetadataModel(StrictResponseModel): + id: UUID + name: str + description: str | None + + +class MetadataModelVersion(StrictResponseModel): + id: UUID + model_id: UUID + version: str + description: str | None + + +class MetadataVariable(StrictResponseModel): + id: UUID + name: str + label: str | None + entity: str + description: str | None + data_type: str | None + possible_values: list[str] | None + default_value: JsonValue + adds: list[str] | None + subtracts: list[str] | None + + +class MetadataParameterNode(StrictResponseModel): + id: UUID + name: str + label: str | None + description: str | None + + +class MetadataParameterValue(StrictResponseModel): + id: UUID + value: JsonValue + start_date: datetime + end_date: datetime | None + + +class MetadataParameter(StrictResponseModel): + id: UUID + name: str + label: str | None + description: str | None + data_type: str | None + unit: str | None + values: list[MetadataParameterValue] + + +class MetadataDataset(StrictResponseModel): + id: UUID + name: str + description: str | None + year: int + storage_path: None = None + is_output_dataset: Literal[False] = False + + +class MetadataRegion(StrictResponseModel): + id: UUID + code: str + label: str + region_type: MetadataRegionType + requires_filter: bool + filter_field: str | None + filter_value: str | None + filter_strategy: str | None + parent_code: str | None + state_code: str | None + state_name: str | None + default_dataset_id: UUID + + +class MetadataRegionOption(StrictResponseModel): + name: str + label: str + type: MetadataRegionType + + +class MetadataTimePeriodOption(StrictResponseModel): + name: int + label: str + + +class MetadataDatasetOption(StrictResponseModel): + name: str + label: str + default: Literal[True] = True + + +class MetadataEconomyOptions(StrictResponseModel): + region: list[MetadataRegionOption] + time_period: list[MetadataTimePeriodOption] + datasets: list[MetadataDatasetOption] + + +class MetadataResult(StrictResponseModel): + current_law_id: int + model: MetadataModel + model_version: MetadataModelVersion + variables: list[MetadataVariable] + parameter_nodes: list[MetadataParameterNode] + parameters: list[MetadataParameter] + datasets: list[MetadataDataset] + regions: list[MetadataRegion] + economy_options: MetadataEconomyOptions + + +class MetadataSuccessResponse(StrictResponseModel): + status: Literal["ok"] = "ok" + message: None = None + result: MetadataResult + + +class MetadataErrorResponse(StrictResponseModel): + status: Literal["error"] = "error" + message: Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)] + + +MetadataPreviewResponse = Annotated[ + MetadataSuccessResponse | MetadataErrorResponse, + Field(discriminator="status"), +] diff --git a/policyengine_api/data/v2/models/metadata.py b/policyengine_api/data/v2/models/metadata.py index 4eaa3d7bd..ab452d5e2 100644 --- a/policyengine_api/data/v2/models/metadata.py +++ b/policyengine_api/data/v2/models/metadata.py @@ -42,11 +42,9 @@ class TaxBenefitModel(TimestampedModel, table=True): back_populates="model", cascade_delete=True, ) - datasets: list["Dataset"] = Relationship(back_populates="tax_benefit_model") dataset_versions: list["DatasetVersion"] = Relationship( back_populates="tax_benefit_model" ) - regions: list["Region"] = Relationship(back_populates="tax_benefit_model") policies: list["Policy"] = Relationship(back_populates="tax_benefit_model") reports: list["Report"] = Relationship(back_populates="tax_benefit_model") @@ -68,6 +66,8 @@ class TaxBenefitModelVersion(IdentifiedModel, table=True): ) version: str = Field(max_length=128) description: str | None = None + current_law_id: int + metadata_time_periods: list[int] = Field(sa_type=sa.JSON) model: TaxBenefitModel = Relationship(back_populates="versions") variables: list["Variable"] = Relationship( @@ -85,15 +85,17 @@ class TaxBenefitModelVersion(IdentifiedModel, table=True): simulations: list["Simulation"] = Relationship( back_populates="tax_benefit_model_version" ) + datasets: list["Dataset"] = Relationship(back_populates="tax_benefit_model_version") + regions: list["Region"] = Relationship(back_populates="tax_benefit_model_version") class Region(TimestampedModel, table=True): __tablename__ = "regions" __table_args__ = ( sa.UniqueConstraint( - "tax_benefit_model_id", + "tax_benefit_model_version_id", "code", - name="uq_regions_model_code", + name="uq_regions_model_version_code", ), sa.CheckConstraint( "NOT requires_filter OR " @@ -101,11 +103,11 @@ class Region(TimestampedModel, table=True): name="ck_regions_required_filter_values", ), # SQLModel does not expose table-level composite foreign keys. This - # keeps a seeded region and its one default dataset in the same model. + # keeps a region and its default dataset in the same model version. sa.ForeignKeyConstraint( - ["default_dataset_id", "tax_benefit_model_id"], - ["datasets.id", "datasets.tax_benefit_model_id"], - name="fk_regions_default_dataset_model_datasets", + ["default_dataset_id", "tax_benefit_model_version_id"], + ["datasets.id", "datasets.tax_benefit_model_version_id"], + name="fk_regions_default_dataset_model_version", ondelete="RESTRICT", ), ) @@ -120,17 +122,19 @@ class Region(TimestampedModel, table=True): parent_code: str | None = Field(default=None, max_length=255) state_code: str | None = Field(default=None, max_length=16) state_name: str | None = Field(default=None, max_length=128) - tax_benefit_model_id: UUID = Field( - foreign_key="tax_benefit_models.id", + tax_benefit_model_version_id: UUID = Field( + foreign_key="tax_benefit_model_versions.id", ondelete="RESTRICT", index=True, ) default_dataset_id: UUID = Field(index=True) - tax_benefit_model: TaxBenefitModel = Relationship(back_populates="regions") + tax_benefit_model_version: TaxBenefitModelVersion = Relationship( + back_populates="regions" + ) default_dataset: "Dataset" = Relationship( back_populates="default_for_regions", - sa_relationship_kwargs={"overlaps": "regions,tax_benefit_model"}, + sa_relationship_kwargs={"viewonly": True}, ) simulations: list["Simulation"] = Relationship(back_populates="region") reports: list["Report"] = Relationship(back_populates="region") @@ -140,14 +144,14 @@ class Dataset(TimestampedModel, table=True): __tablename__ = "datasets" __table_args__ = ( sa.UniqueConstraint( - "tax_benefit_model_id", + "tax_benefit_model_version_id", "name", - name="uq_datasets_model_name", + name="uq_datasets_model_version_name", ), sa.UniqueConstraint( "id", - "tax_benefit_model_id", - name="uq_datasets_id_model", + "tax_benefit_model_version_id", + name="uq_datasets_id_model_version", ), sa.CheckConstraint( "year BETWEEN 1900 AND 2200", @@ -164,20 +168,22 @@ class Dataset(TimestampedModel, table=True): storage_path: str | None = Field(default=None, max_length=1024) year: int is_output_dataset: bool = False - tax_benefit_model_id: UUID = Field( - foreign_key="tax_benefit_models.id", + tax_benefit_model_version_id: UUID = Field( + foreign_key="tax_benefit_model_versions.id", ondelete="RESTRICT", index=True, ) - tax_benefit_model: TaxBenefitModel = Relationship(back_populates="datasets") + tax_benefit_model_version: TaxBenefitModelVersion = Relationship( + back_populates="datasets" + ) versions: list["DatasetVersion"] = Relationship( back_populates="dataset", cascade_delete=True, ) default_for_regions: list[Region] = Relationship( back_populates="default_dataset", - sa_relationship_kwargs={"overlaps": "regions,tax_benefit_model"}, + sa_relationship_kwargs={"viewonly": True}, ) input_simulations: list["Simulation"] = Relationship( back_populates="dataset", @@ -314,6 +320,17 @@ class ParameterValue(IdentifiedModel, table=True): "start_date", "end_date", ), + # SQLModel does not expose dialect-specific partial unique indexes. + # Canonical values have neither an owning policy nor dynamic, while + # those owners may each store their own value for the same period. + sa.Index( + "uq_parameter_values_canonical_parameter_start_date", + "parameter_id", + "start_date", + unique=True, + postgresql_where=sa.text("policy_id IS NULL AND dynamic_id IS NULL"), + sqlite_where=sa.text("policy_id IS NULL AND dynamic_id IS NULL"), + ), ) parameter_id: UUID = Field( diff --git a/policyengine_api/data/v2/settings.py b/policyengine_api/data/v2/settings.py index 0cfca4548..a0f26cde3 100644 --- a/policyengine_api/data/v2/settings.py +++ b/policyengine_api/data/v2/settings.py @@ -9,15 +9,19 @@ from collections.abc import Mapping from dataclasses import dataclass, field +from functools import lru_cache import os import re +from typing import Callable from sqlalchemy.engine import URL, make_url from sqlalchemy.exc import ArgumentError V2_RUNTIME_DATABASE_URL = "V2_RUNTIME_DATABASE_URL" +V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE = "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE" V2_MIGRATION_DATABASE_URL = "V2_MIGRATION_DATABASE_URL" +V2_DATA_WRITE_DATABASE_URL = "V2_DATA_WRITE_DATABASE_URL" V2_SUPABASE_PROJECT_REF = "V2_SUPABASE_PROJECT_REF" V2_SUPABASE_ENVIRONMENT = "V2_SUPABASE_ENVIRONMENT" @@ -26,6 +30,7 @@ LOCAL_HOSTS = frozenset({"localhost", "127.0.0.1", "::1"}) PROJECT_REF_PATTERN = re.compile(r"^[a-z0-9]{20}$") ENVIRONMENT_PATTERN = re.compile(r"^[a-z][a-z0-9-]{1,31}$") +SECRET_RESOURCE_PATTERN = re.compile(r"^projects/[^/]+/secrets/[^/]+/versions/[^/]+$") class V2ConfigurationError(RuntimeError): @@ -81,6 +86,53 @@ def _required(environ: Mapping[str, str], name: str) -> str: return value.strip() +@lru_cache(maxsize=None) +def _load_secret_from_secret_manager(resource_name: str) -> str: + """Resolve a v2 runtime URL only when a preview request selects it.""" + + from google.cloud import secretmanager + + client = secretmanager.SecretManagerServiceClient() + response = client.access_secret_version(request={"name": resource_name}) + return response.payload.data.decode("utf-8") + + +def _resolve_runtime_database_url( + environ: Mapping[str, str], + *, + secret_loader: Callable[[str], str], +) -> str: + direct_value = environ.get(V2_RUNTIME_DATABASE_URL, "").strip() + resource = environ.get(V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE, "").strip() + if direct_value and resource: + raise V2ConfigurationError( + f"set exactly one of {V2_RUNTIME_DATABASE_URL} or " + f"{V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE}" + ) + if direct_value: + return direct_value + if not resource: + raise V2ConfigurationError( + f"{V2_RUNTIME_DATABASE_URL} or " + f"{V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE} is required" + ) + if SECRET_RESOURCE_PATTERN.fullmatch(resource) is None: + raise V2ConfigurationError( + f"{V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE} is invalid" + ) + try: + resolved_value = secret_loader(resource).strip() + except Exception as error: + raise V2ConfigurationError( + f"{V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE} could not be resolved" + ) from error + if not resolved_value: + raise V2ConfigurationError( + f"{V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE} is empty" + ) + return resolved_value + + def load_supabase_target_settings( environ: Mapping[str, str] | None = None, ) -> SupabaseTargetSettings: @@ -140,12 +192,18 @@ def parse_persistent_postgres_url( def load_v2_runtime_database_settings( environ: Mapping[str, str] | None = None, + *, + secret_loader: Callable[[str], str] | None = None, ) -> V2DatabaseSettings: """Load the future ordinary-runtime Postgres identity explicitly.""" values = _environment(environ) + raw_url = _resolve_runtime_database_url( + values, + secret_loader=secret_loader or _load_secret_from_secret_manager, + ) connection = parse_persistent_postgres_url( - _required(values, V2_RUNTIME_DATABASE_URL), + raw_url, setting_name=V2_RUNTIME_DATABASE_URL, ) return V2DatabaseSettings( @@ -168,3 +226,19 @@ def load_v2_migration_database_settings( connection=connection, target=load_supabase_target_settings(values), ) + + +def load_v2_data_write_database_settings( + environ: Mapping[str, str] | None = None, +) -> V2DatabaseSettings: + """Load the one-time catalog row-write Postgres identity explicitly.""" + + values = _environment(environ) + connection = parse_persistent_postgres_url( + _required(values, V2_DATA_WRITE_DATABASE_URL), + setting_name=V2_DATA_WRITE_DATABASE_URL, + ) + return V2DatabaseSettings( + connection=connection, + target=load_supabase_target_settings(values), + ) diff --git a/policyengine_api/fastapi_routes/dependencies.py b/policyengine_api/fastapi_routes/dependencies.py index b1a6eccf4..5b87e8465 100644 --- a/policyengine_api/fastapi_routes/dependencies.py +++ b/policyengine_api/fastapi_routes/dependencies.py @@ -4,8 +4,11 @@ from collections.abc import Callable from dataclasses import dataclass +from functools import lru_cache +from importlib import metadata as importlib_metadata from typing import Protocol +from policyengine_api.data.v2.catalog.schemas import MetadataResult from policyengine_api.json_types import JSONObject @@ -15,6 +18,18 @@ class MetadataReader(Protocol): def get_metadata(self, country_id: str) -> JSONObject: ... +class V2MetadataReader(Protocol): + """Read one already-initialized v2 metadata catalog.""" + + def get_metadata( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataResult: ... + + def close(self) -> None: ... + + class SimulationGatewayProbe(Protocol): """Minimal simulation-entrypoint health-check interface.""" @@ -45,6 +60,21 @@ def _default_specification_provider() -> JSONObject: return OPENAPI_SPECIFICATION +@lru_cache(maxsize=1) +def _running_policyengine_version() -> str: + return importlib_metadata.version("policyengine") + + +def _default_v2_metadata_reader_factory() -> V2MetadataReader: + from policyengine_api.data.v2.catalog.query import V2MetadataQueryService + from policyengine_api.data.v2.database import get_v2_session_factory + + return V2MetadataQueryService( + get_v2_session_factory()(), + running_policyengine_version=_running_policyengine_version(), + ) + + @dataclass(frozen=True) class NativeRouteDependencies: """Runtime collaborators for native read routes.""" @@ -53,6 +83,7 @@ class NativeRouteDependencies: gateway_client_factory: Callable[[], SimulationGatewayProbe] metadata_reader_factory: Callable[[], MetadataReader] specification_provider: Callable[[], JSONObject] + v2_metadata_reader_factory: Callable[[], V2MetadataReader] | None = None @classmethod def defaults(cls) -> "NativeRouteDependencies": @@ -62,4 +93,5 @@ def defaults(cls) -> "NativeRouteDependencies": gateway_client_factory=_default_gateway_client_factory, metadata_reader_factory=_default_metadata_reader_factory, specification_provider=_default_specification_provider, + v2_metadata_reader_factory=_default_v2_metadata_reader_factory, ) diff --git a/policyengine_api/fastapi_routes/v2_metadata.py b/policyengine_api/fastapi_routes/v2_metadata.py new file mode 100644 index 000000000..410815eb0 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata.py @@ -0,0 +1,154 @@ +"""Dormant, read-only API v2 metadata preview routes.""" + +from __future__ import annotations + +from fastapi import APIRouter, Request +from starlette.responses import JSONResponse + +from policyengine_api.data.v2.catalog.query import ( + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataErrorResponse, + MetadataSuccessResponse, +) +from policyengine_api.data.v2.settings import V2ConfigurationError +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies + + +ERROR_RESPONSES = { + 400: { + "model": MetadataErrorResponse, + "description": "The requested PolicyEngine.py version is invalid.", + }, + 404: { + "model": MetadataErrorResponse, + "description": "The requested PolicyEngine.py catalog is absent.", + }, + 405: { + "model": MetadataErrorResponse, + "description": "The preview supports GET only.", + }, + 503: { + "model": MetadataErrorResponse, + "description": "The initialized v2 catalog is unavailable.", + }, + 500: { + "model": MetadataErrorResponse, + "description": "The preview query failed internally.", + }, +} + + +def _error_response(status_code: int, message: str) -> JSONResponse: + error = MetadataErrorResponse(message=message) + return JSONResponse( + status_code=status_code, + content=error.model_dump(mode="json"), + ) + + +def build_v2_metadata_router( + dependencies: NativeRouteDependencies, +) -> APIRouter: + """Build isolated preview routes without loading v2 configuration.""" + + router = APIRouter() + + @router.get( + "/v2/openapi.json", + include_in_schema=False, + summary="OpenAPI document for dormant v2 preview routes", + ) + def v2_preview_openapi(request: Request) -> JSONResponse: + schema = request.app.openapi() + preview_schema = { + **schema, + "paths": { + path: operation + for path, operation in schema.get("paths", {}).items() + if path.startswith("/v2/") + }, + } + return JSONResponse(preview_schema) + + def read( + country_id: str, + policyengine_version: str | None, + ) -> MetadataSuccessResponse | JSONResponse: + reader = None + try: + factory = dependencies.v2_metadata_reader_factory + if factory is None: + from policyengine_api.fastapi_routes.dependencies import ( + _default_v2_metadata_reader_factory, + ) + + factory = _default_v2_metadata_reader_factory + reader = factory() + result = reader.get_metadata(country_id, policyengine_version) + return MetadataSuccessResponse(result=result) + except InvalidPolicyEngineVersionError as error: + return _error_response(400, str(error)) + except MetadataCatalogVersionNotFoundError as error: + return _error_response(404, str(error)) + except (V2ConfigurationError, MetadataCatalogUnavailableError): + return _error_response(503, "V2 metadata catalog is unavailable") + except Exception: # noqa: BLE001 - preview must return typed errors + return _error_response(500, "V2 metadata query failed") + finally: + if reader is not None: + try: + reader.close() + except Exception: + pass + + @router.get( + "/v2/us/metadata", + response_model=MetadataSuccessResponse, + responses=ERROR_RESPONSES, + summary="Preview US metadata from the v2 catalog", + ) + def us_metadata_preview( + policyengine_version: str | None = None, + ) -> MetadataSuccessResponse | JSONResponse: + return read("us", policyengine_version) + + @router.get( + "/v2/uk/metadata", + response_model=MetadataSuccessResponse, + responses=ERROR_RESPONSES, + summary="Preview UK metadata from the v2 catalog", + ) + def uk_metadata_preview( + policyengine_version: str | None = None, + ) -> MetadataSuccessResponse | JSONResponse: + return read("uk", policyengine_version) + + @router.get( + "/v2/{country_id}/metadata", + response_model=MetadataErrorResponse, + status_code=404, + responses={500: ERROR_RESPONSES[500]}, + summary="Reject an unsupported v2 metadata preview country", + ) + def unsupported_country(country_id: str) -> MetadataErrorResponse: + return MetadataErrorResponse( + message=f"V2 metadata is not available for country {country_id}" + ) + + @router.api_route( + "/v2/{country_id}/metadata", + methods=["POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"], + response_model=MetadataErrorResponse, + status_code=405, + include_in_schema=False, + ) + def unsupported_method(country_id: str) -> MetadataErrorResponse: + return MetadataErrorResponse( + message=f"V2 metadata for country {country_id} supports GET only" + ) + + return router diff --git a/policyengine_api/migration_flags.py b/policyengine_api/migration_flags.py index 0fbfe122b..0d26aed5c 100644 --- a/policyengine_api/migration_flags.py +++ b/policyengine_api/migration_flags.py @@ -92,7 +92,7 @@ def infer_route_group(path: str) -> str: "/readiness-check", }: return "health" - if path == "/specification": + if path in {"/specification", "/v2/openapi.json"}: return "specification" segments = [segment for segment in path.strip("/").split("/") if segment] @@ -106,6 +106,9 @@ def infer_route_group(path: str) -> str: if len(segments) >= 2 and segments[1] in ROUTE_GROUP_BY_SEGMENT: return ROUTE_GROUP_BY_SEGMENT[segments[1]] + if first == "v2" and len(segments) >= 3 and segments[2] in ROUTE_GROUP_BY_SEGMENT: + return ROUTE_GROUP_BY_SEGMENT[segments[2]] + return "unknown" @@ -184,6 +187,9 @@ def get_migration_context( route_impl: RouteImplementation | None = None, db_entity: str | None = None, sim_flow: str | None = None, + use_configured_db_sources: bool = True, + db_write_source: str | None = None, + db_read_source: str | None = None, ) -> MigrationContext: """Return current migration flag values for a request or route group.""" route_config = ROUTE_GROUP_CONFIG_BY_NAME.get(route_group) @@ -192,6 +198,21 @@ def get_migration_context( if sim_flow is None and route_config is not None: sim_flow = route_config.sim_flow + if use_configured_db_sources: + db_write = get_db_write(db_entity) if db_entity else None + db_read = get_db_read(db_entity) if db_entity else None + else: + if db_write_source is not None and db_write_source not in DB_WRITE_SOURCES: + raise ValueError( + f"invalid explicit database write source {db_write_source!r}" + ) + if db_read_source is not None and db_read_source not in DB_READ_SOURCES: + raise ValueError( + f"invalid explicit database read source {db_read_source!r}" + ) + db_write = db_write_source + db_read = db_read_source + return MigrationContext( api_host_backend=_read_choice( "API_HOST_BACKEND", @@ -201,8 +222,8 @@ def get_migration_context( route_group=route_group, route_impl=route_impl or get_route_impl(route_group), db_entity=db_entity, - db_write=get_db_write(db_entity) if db_entity else None, - db_read=get_db_read(db_entity) if db_entity else None, + db_write=db_write, + db_read=db_read, sim_flow=sim_flow, sim_entrypoint=get_sim_entrypoint(), sim_compute=get_sim_compute(sim_flow) if sim_flow else None, @@ -213,12 +234,18 @@ def get_migration_log_context( route_group: str, *, route_impl: RouteImplementation | None = None, + use_configured_db_sources: bool = True, + db_write_source: str | None = None, + db_read_source: str | None = None, ) -> dict: """Best-effort logging context; never raises on invalid flag settings.""" try: return get_migration_context( route_group, route_impl=route_impl, + use_configured_db_sources=use_configured_db_sources, + db_write_source=db_write_source, + db_read_source=db_read_source, ).to_log_dict() except ValueError as error: return { diff --git a/policyengine_api/migration_logging.py b/policyengine_api/migration_logging.py index 511eff10a..d98de5d55 100644 --- a/policyengine_api/migration_logging.py +++ b/policyengine_api/migration_logging.py @@ -19,6 +19,14 @@ ) +V2_METADATA_PREVIEW_READ_PATHS = frozenset( + { + "/v2/us/metadata", + "/v2/uk/metadata", + } +) + + def register_migration_request_logging(app: flask.Flask) -> None: """Register request IDs, backend headers, and migration logging for Flask.""" @@ -72,9 +80,12 @@ def log_migration_request( elapsed_ms = round((time.time() - started_at) * 1000, 2) route_group = infer_route_group(path) + is_v2_metadata_read = method == "GET" and path in V2_METADATA_PREVIEW_READ_PATHS migration_context = get_migration_log_context( route_group, route_impl=route_impl, + use_configured_db_sources=not is_v2_metadata_read, + db_read_source="supabase" if is_v2_metadata_read else None, ) logger.log_struct( diff --git a/pyproject.toml b/pyproject.toml index 3fd93fcaf..2a31bb578 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,6 +38,7 @@ dependencies = [ "markupsafe>=3,<4", "microdf_python>=1.0.0", "openai", + "packaging>=24,<27", "policyengine_canada==0.96.3", "policyengine-ng==0.5.1", "policyengine-il==0.1.0", diff --git a/scripts/guards/migration_contracts.py b/scripts/guards/migration_contracts.py index 2cdec9d31..abf50e29b 100644 --- a/scripts/guards/migration_contracts.py +++ b/scripts/guards/migration_contracts.py @@ -9,6 +9,7 @@ REPO_ROOT = Path(__file__).resolve().parents[2] +ALLOWED_CURRENT_CONTRACTS = frozenset({"api_v1_compatible", "typed_v2_preview"}) def _check_unique_values( @@ -68,10 +69,8 @@ def _check_workflows(payload: dict[str, Any]) -> list[str]: request_keys = [] for workflow in workflows: - if workflow["current_contract"] != "api_v1_compatible": - violations.append( - f"{workflow['name']}: current_contract should be api_v1_compatible" - ) + if workflow["current_contract"] not in ALLOWED_CURRENT_CONTRACTS: + violations.append(f"{workflow['name']}: current_contract is not recognized") if not workflow["future_owner_pr"]: violations.append(f"{workflow['name']}: future_owner_pr is required") if not workflow["requests"]: diff --git a/scripts/initialize_v2_metadata.py b/scripts/initialize_v2_metadata.py new file mode 100644 index 000000000..079176823 --- /dev/null +++ b/scripts/initialize_v2_metadata.py @@ -0,0 +1,8 @@ +#!/usr/bin/env python3 +"""Run the explicit one-time API v2 metadata initialization operation.""" + +from policyengine_api.data.v2.catalog.initialization import main + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/contract/registry.py b/tests/contract/registry.py index 1ba70280d..0592c5e2a 100644 --- a/tests/contract/registry.py +++ b/tests/contract/registry.py @@ -122,6 +122,37 @@ class WorkflowContract: ), ), ), + WorkflowContract( + name="region_selection_v2_preview", + current_contract="typed_v2_preview", + future_owner_pr="Later metadata read cutover and preview-path removal", + requests=( + ContractRequest( + method="GET", + path="/v2/us/metadata", + expected_status=200, + stable_response_fields=( + "status", + "result.current_law_id", + "result.economy_options.region", + "result.economy_options.time_period", + ), + route_group="metadata", + ), + ContractRequest( + method="GET", + path="/v2/uk/metadata", + expected_status=200, + stable_response_fields=( + "status", + "result.current_law_id", + "result.economy_options.region", + "result.economy_options.time_period", + ), + route_group="metadata", + ), + ), + ), WorkflowContract( name="simulation_submit_poll", current_contract="api_v1_compatible", @@ -201,3 +232,10 @@ class WorkflowContract: APP_V2_ROUTE_CONTRACTS = tuple( request for workflow in APP_V2_WORKFLOW_CONTRACTS for request in workflow.requests ) + +APP_V1_COMPATIBLE_ROUTE_CONTRACTS = tuple( + request + for workflow in APP_V2_WORKFLOW_CONTRACTS + if workflow.current_contract == "api_v1_compatible" + for request in workflow.requests +) diff --git a/tests/contract/test_app_v2_workflow_contracts.py b/tests/contract/test_app_v2_workflow_contracts.py index 742287943..6e0aa49e4 100644 --- a/tests/contract/test_app_v2_workflow_contracts.py +++ b/tests/contract/test_app_v2_workflow_contracts.py @@ -1,5 +1,9 @@ from policyengine_api.migration_registry import ROUTE_GROUP_CONFIG_BY_NAME -from tests.contract.registry import APP_V2_ROUTE_CONTRACTS, APP_V2_WORKFLOW_CONTRACTS +from tests.contract.registry import ( + APP_V1_COMPATIBLE_ROUTE_CONTRACTS, + APP_V2_ROUTE_CONTRACTS, + APP_V2_WORKFLOW_CONTRACTS, +) def test_app_v2_workflow_contract_registry_is_complete(): @@ -8,13 +12,19 @@ def test_app_v2_workflow_contract_registry_is_complete(): "household_save_edit_read", "household_calculate", "region_selection", + "region_selection_v2_preview", "simulation_submit_poll", "report_create_poll", "budget_window_submit_poll", } for workflow in APP_V2_WORKFLOW_CONTRACTS: - assert workflow.current_contract == "api_v1_compatible" + expected_contract = ( + "typed_v2_preview" + if workflow.name == "region_selection_v2_preview" + else "api_v1_compatible" + ) + assert workflow.current_contract == expected_contract assert workflow.future_owner_pr assert workflow.requests @@ -24,3 +34,11 @@ def test_app_v2_workflow_contract_registry_is_complete(): assert request.expected_status in {200, 201, 202} assert request.stable_response_fields assert request.route_group in ROUTE_GROUP_CONFIG_BY_NAME + + assert all( + not request.path.startswith("/v2/") + for request in APP_V1_COMPATIBLE_ROUTE_CONTRACTS + ) + assert {request.path for request in APP_V2_ROUTE_CONTRACTS} - { + request.path for request in APP_V1_COMPATIBLE_ROUTE_CONTRACTS + } == {"/v2/us/metadata", "/v2/uk/metadata"} diff --git a/tests/contract/test_v1_route_contracts.py b/tests/contract/test_v1_route_contracts.py index eae014cf1..bc06f381f 100644 --- a/tests/contract/test_v1_route_contracts.py +++ b/tests/contract/test_v1_route_contracts.py @@ -38,7 +38,7 @@ assert_subset, response_json, ) -from tests.contract.registry import APP_V2_ROUTE_CONTRACTS, ContractRequest +from tests.contract.registry import APP_V1_COMPATIBLE_ROUTE_CONTRACTS, ContractRequest class _BudgetWindowEconomicImpactResult: @@ -458,7 +458,7 @@ def _expected_subset(contract: ContractRequest) -> dict: @pytest.mark.parametrize( "contract", - APP_V2_ROUTE_CONTRACTS, + APP_V1_COMPATIBLE_ROUTE_CONTRACTS, ids=lambda contract: f"{contract.method} {contract.path}", ) def test_app_v2_api_v1_route_contract( diff --git a/tests/fixtures/v2_catalog.py b/tests/fixtures/v2_catalog.py new file mode 100644 index 000000000..ad887af56 --- /dev/null +++ b/tests/fixtures/v2_catalog.py @@ -0,0 +1,234 @@ +"""Deterministic PolicyEngine.py-like source objects for v2 catalog tests.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from types import SimpleNamespace + +from policyengine_api.data.v2.catalog.extraction import extract_catalog +from policyengine_api.data.v2.catalog.records import NormalizedCatalog + + +POLICYENGINE_VERSION = "5.0.4" +DEPENDENCY_VERSIONS = { + "policyengine-core": "3.30.1", + "policyengine-us": "1.764.6", + "policyengine-uk": "2.90.2", +} + + +class RegionRegistry(list): + """Small iterable with the public registry's country identifier.""" + + def __init__(self, country_id: str, regions: list[SimpleNamespace]): + super().__init__(regions) + self.country_id = country_id + + +def _strategy(field: str, value: str | int) -> SimpleNamespace: + return SimpleNamespace( + strategy_type="row_filter", + variable_name=field, + variable_value=value, + additional_filters={}, + ) + + +def _region( + *, + code: str, + label: str, + region_type: str, + parent_code: str | None = None, + dataset_path: str | None = None, + strategy: SimpleNamespace | None = None, + state_code: str | None = None, + state_name: str | None = None, +) -> SimpleNamespace: + return SimpleNamespace( + code=code, + label=label, + region_type=region_type, + parent_code=parent_code, + dataset_path=dataset_path, + scoping_strategy=strategy, + requires_filter=strategy is not None, + state_code=state_code, + state_name=state_name, + ) + + +def bundle(*, policyengine_version: str = POLICYENGINE_VERSION) -> dict: + """Return a minimal packaged-bundle manifest.""" + + return { + "bundle_version": policyengine_version, + "policyengine_version": policyengine_version, + "packages": { + "policyengine": { + "name": "policyengine", + "version": policyengine_version, + }, + **{ + name: {"name": name, "version": version} + for name, version in DEPENDENCY_VERSIONS.items() + }, + }, + } + + +def _model_source( + *, + country_id: str, + country_package_version: str, + policyengine_version: str, + regions: list[SimpleNamespace], +) -> SimpleNamespace: + model_name = f"policyengine-{country_id}" + default_dataset = { + "us": "populace_us_2024", + "uk": "enhanced_frs_2024_25", + }[country_id] + default_uri = f"hf://policyengine/{country_id}/{default_dataset}.h5@fixture" + version_holder = SimpleNamespace() + variable = SimpleNamespace( + name="employment_income", + label="Employment income", + entity="person", + description="Employment income before tax", + data_type=float, + possible_values=None, + default_value=0.0, + adds=None, + subtracts=None, + ) + parameter_node = SimpleNamespace( + name="gov.example", + label="Example policy", + description="Example policy parameters", + ) + parameter = SimpleNamespace( + name="gov.example.rate", + label="Example rate", + description="An example rate", + data_type=float, + unit="/1", + parameter_values=[ + SimpleNamespace( + value=0.1, + start_date=datetime(2025, 1, 1, tzinfo=timezone.utc), + end_date=datetime(2025, 12, 31, tzinfo=timezone.utc), + ), + SimpleNamespace( + value=0.2, + start_date=datetime(2026, 1, 1, tzinfo=timezone.utc), + end_date=None, + ), + ], + ) + version_holder.model = SimpleNamespace( + id=model_name, + description=f"Fixture {country_id.upper()} model", + ) + version_holder.version = country_package_version + version_holder.model_package = SimpleNamespace( + name=model_name, + version=country_package_version, + ) + version_holder.release_manifest = SimpleNamespace( + country_id=country_id, + policyengine_version=policyengine_version, + default_dataset=default_dataset, + default_dataset_uri=default_uri, + ) + version_holder.region_registry = RegionRegistry(country_id, regions) + version_holder.variables = [variable] + version_holder.parameter_nodes = [parameter_node] + version_holder.parameters = [parameter] + version_holder.variables_by_name = {variable.name: variable} + version_holder.parameter_nodes_by_name = {parameter_node.name: parameter_node} + version_holder.parameters_by_name = {parameter.name: parameter} + return version_holder + + +def source_models( + *, + policyengine_version: str = POLICYENGINE_VERSION, +) -> dict[str, SimpleNamespace]: + """Return US and UK public-model fixtures.""" + + us_default_uri = "hf://policyengine/us/populace_us_2024.h5@fixture" + uk_default_uri = "hf://policyengine/uk/enhanced_frs_2024_25.h5@fixture" + return { + "us": _model_source( + country_id="us", + country_package_version=DEPENDENCY_VERSIONS["policyengine-us"], + policyengine_version=policyengine_version, + regions=[ + _region( + code="us", + label="United States", + region_type="national", + dataset_path=us_default_uri, + ), + _region( + code="state/ca", + label="California", + region_type="state", + parent_code="us", + dataset_path="hf://policyengine/us/populace_us_ca_2024.h5@fixture", + strategy=_strategy("state_fips", 6), + state_code="CA", + state_name="California", + ), + _region( + code="place/CA-44000", + label="Los Angeles", + region_type="place", + parent_code="state/ca", + state_code="CA", + state_name="California", + ), + ], + ), + "uk": _model_source( + country_id="uk", + country_package_version=DEPENDENCY_VERSIONS["policyengine-uk"], + policyengine_version=policyengine_version, + regions=[ + _region( + code="uk", + label="United Kingdom", + region_type="national", + dataset_path=uk_default_uri, + ), + _region( + code="country/england", + label="England", + region_type="country", + parent_code="uk", + strategy=_strategy("country", "ENGLAND"), + ), + ], + ), + } + + +def installed_version(name: str) -> str: + """Return fixture observed versions.""" + + return DEPENDENCY_VERSIONS[name] + + +def normalized_catalog( + *, + policyengine_version: str = POLICYENGINE_VERSION, +) -> NormalizedCatalog: + """Return the complete deterministic normalized fixture.""" + + return extract_catalog( + bundle=bundle(policyengine_version=policyengine_version), + policyengine_version=policyengine_version, + models=source_models(policyengine_version=policyengine_version), + installed_version=installed_version, + ) diff --git a/tests/integration/test_alembic_v2_lifecycle.py b/tests/integration/test_alembic_v2_lifecycle.py index 69ab2e69f..dd24c82fd 100644 --- a/tests/integration/test_alembic_v2_lifecycle.py +++ b/tests/integration/test_alembic_v2_lifecycle.py @@ -22,7 +22,8 @@ from policyengine_api.data.v2.settings import V2_MIGRATION_DATABASE_URL -HEAD_REVISION = "f5ef4347cb2a" +BASELINE_REVISION = "f5ef4347cb2a" +HEAD_REVISION = "68b4a5ae5dc5" V2_TABLE_NAMES = frozenset(table.name for table in V2_METADATA.tables.values()) @@ -66,6 +67,15 @@ def _assert_head(engine) -> None: ) ).scalar_one() assert (model_count, version_count) == (0, 0) + parameter_value_indexes = { + index["name"]: index + for index in inspect(engine).get_indexes("parameter_values") + } + canonical_index = parameter_value_indexes[ + "uq_parameter_values_canonical_parameter_start_date" + ] + assert canonical_index["unique"] + assert canonical_index["column_names"] == ["parameter_id", "start_date"] def test_empty_upgrade_check_base_downgrade_and_reupgrade() -> None: @@ -83,6 +93,18 @@ def test_empty_upgrade_check_base_downgrade_and_reupgrade() -> None: command.check(config) _assert_head(engine) + command.downgrade(config, BASELINE_REVISION) + with engine.connect() as connection: + context = MigrationContext.configure(connection) + assert context.get_current_revision() == BASELINE_REVISION + assert "uq_parameter_values_canonical_parameter_start_date" not in { + index["name"] for index in inspect(engine).get_indexes("parameter_values") + } + + command.upgrade(config, "head") + command.check(config) + _assert_head(engine) + command.downgrade(config, "base") assert set(inspect(engine).get_table_names(schema="public")) <= { "alembic_version" @@ -222,7 +244,7 @@ def test_baseline_region_default_enforces_same_model_dataset() -> None: second_dataset_id = uuid4() try: - command.upgrade(config, "head") + command.downgrade(config, BASELINE_REVISION) assert "region_datasets" not in inspect(engine).get_table_names(schema="public") default_column = next( column diff --git a/tests/integration/test_v2_catalog_installed.py b/tests/integration/test_v2_catalog_installed.py new file mode 100644 index 000000000..05d79c47a --- /dev/null +++ b/tests/integration/test_v2_catalog_installed.py @@ -0,0 +1,82 @@ +"""Compatibility coverage for the complete installed PolicyEngine.py catalog.""" + +from __future__ import annotations + +from importlib import metadata as importlib_metadata +import os + +import pytest + +from policyengine.bundle import get_current_bundle +from policyengine_api.data.v2.catalog.extraction import extract_installed_catalog + + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_V2_CATALOG_COMPATIBILITY") != "1", + reason=("loads the complete public catalog; set RUN_V2_CATALOG_COMPATIBILITY=1"), +) + + +def test_installed_policyengine_catalog_is_complete_and_bounded() -> None: + catalog = extract_installed_catalog() + bundle = get_current_bundle() + policyengine_version = importlib_metadata.version("policyengine") + expected_dependencies = tuple( + (name, bundle["packages"][name]["version"]) + for name in ( + "policyengine-core", + "policyengine-us", + "policyengine-uk", + ) + ) + + assert catalog.policyengine_version == policyengine_version + assert catalog.dependency_versions == expected_dependencies + assert all( + importlib_metadata.version(name) == expected + for name, expected in expected_dependencies + ) + assert catalog.entity_counts() == { + "models": 2, + "model_versions": 2, + "variables": 6_649, + "parameter_nodes": 27_826, + "parameters": 99_006, + "parameter_values": 1_172_130, + "datasets": 2, + "regions": 826, + } + + for country_id in ("us", "uk"): + country = catalog.country(country_id) + assert country.model.name == f"policyengine-{country_id}" + assert country.model_version.version == policyengine_version + assert country.model_version.version not in dict(expected_dependencies).values() + assert country.variables + assert country.parameter_nodes + assert country.parameters + assert all( + not dataset.is_output_dataset and dataset.storage_path is None + for dataset in country.datasets + ) + + total_values = 0 + for batch in country.parameter_value_batches(batch_size=10_000): + assert 0 < len(batch) <= 10_000 + total_values += len(batch) + assert total_values == country.entity_counts()["parameter_values"] + + assert {dataset.name for dataset in catalog.country("us").datasets} == { + "populace_us_2024" + } + assert {dataset.name for dataset in catalog.country("uk").datasets} == { + "enhanced_frs_2024_25" + } + assert [ + (summary.region_type, summary.count) + for summary in catalog.country("us").fallback_summaries + ] == [ + ("congressional_district", 436), + ("place", 333), + ("state", 51), + ] diff --git a/tests/integration/test_v2_catalog_publication.py b/tests/integration/test_v2_catalog_publication.py new file mode 100644 index 000000000..e5d21f399 --- /dev/null +++ b/tests/integration/test_v2_catalog_publication.py @@ -0,0 +1,342 @@ +"""Disposable-Postgres coverage for Stage 9 catalog publication.""" + +from __future__ import annotations + +from concurrent.futures import ThreadPoolExecutor +from dataclasses import replace +import os +from pathlib import Path + +from alembic import command +from alembic.config import Config +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.engine import Engine, make_url +from sqlalchemy.pool import NullPool +from sqlmodel import Session, select + +from policyengine_api.data.v2.catalog.publication import ( + CatalogPublicationError, + publish_catalog, +) +from policyengine_api.data.v2.settings import V2_MIGRATION_DATABASE_URL +from policyengine_api.data.v2.models import ( + Dataset, + DatasetVersion, + Region, + Report, + ReportRun, + Simulation, + TaxBenefitModel, + TaxBenefitModelVersion, +) +from tests.fixtures.v2_catalog import normalized_catalog + + +REPO = Path(__file__).parents[2] +DISPOSABLE_DATABASE = "policyengine_v2_alembic_test" + + +def _disposable_url() -> str: + database_url = os.environ.get(V2_MIGRATION_DATABASE_URL, "") + if not database_url: + pytest.skip(f"{V2_MIGRATION_DATABASE_URL} is not set") + url = make_url(database_url) + if url.database != DISPOSABLE_DATABASE or url.host not in { + "localhost", + "127.0.0.1", + "::1", + "postgres", + }: + pytest.fail("catalog publication tests require disposable local Postgres") + return database_url + + +@pytest.fixture +def publication_engine() -> Engine: + database_url = _disposable_url() + config = Config(str(REPO / "alembic-v2.ini")) + command.upgrade(config, "head") + engine = create_engine(database_url, poolclass=NullPool) + with engine.begin() as connection: + connection.execute(text("TRUNCATE tax_benefit_models CASCADE")) + yield engine + with engine.begin() as connection: + connection.execute(text("TRUNCATE tax_benefit_models CASCADE")) + engine.dispose() + + +def _counts(engine: Engine) -> dict[str, int]: + with engine.connect() as connection: + return { + table_name: connection.execute( + text(f"SELECT count(*) FROM {table_name}") + ).scalar_one() + for table_name in ( + "tax_benefit_models", + "tax_benefit_model_versions", + "variables", + "parameter_nodes", + "parameters", + "parameter_values", + "datasets", + "dataset_versions", + "regions", + "simulations", + "reports", + "report_runs", + ) + } + + +def _identifiers(engine: Engine) -> dict[str, tuple[str, ...]]: + with engine.connect() as connection: + return { + table_name: tuple( + str(value) + for value in connection.execute( + text(f"SELECT id FROM {table_name} ORDER BY id") + ).scalars() + ) + for table_name in ( + "tax_benefit_models", + "tax_benefit_model_versions", + "variables", + "parameter_nodes", + "parameters", + "parameter_values", + "datasets", + "regions", + ) + } + + +def test_publish_and_republish_preserve_complete_catalog( + publication_engine: Engine, +) -> None: + catalog = normalized_catalog() + + first = publish_catalog(publication_engine, catalog) + first_counts = _counts(publication_engine) + first_ids = _identifiers(publication_engine) + second = publish_catalog(publication_engine, catalog) + + assert first.as_dict()["outcome"] == "ok" + assert second.policyengine_version == first.policyengine_version + assert ( + _counts(publication_engine) + == first_counts + == { + "tax_benefit_models": 2, + "tax_benefit_model_versions": 2, + "variables": 2, + "parameter_nodes": 2, + "parameters": 2, + "parameter_values": 4, + "datasets": 3, + "dataset_versions": 0, + "regions": 5, + "simulations": 0, + "reports": 0, + "report_runs": 0, + } + ) + assert _identifiers(publication_engine) == first_ids + + with publication_engine.connect() as connection: + datasets = connection.execute( + text( + "SELECT name, storage_path, is_output_dataset " + "FROM datasets ORDER BY name" + ) + ).all() + region_defaults = connection.execute( + text( + """ + SELECT region.code, dataset.name + FROM regions AS region + JOIN datasets AS dataset ON dataset.id = region.default_dataset_id + ORDER BY region.code + """ + ) + ).all() + assert all( + storage_path is None and not output for _, storage_path, output in datasets + ) + assert dict(region_defaults) == { + "country/england": "enhanced_frs_2024_25", + "place/CA-44000": "populace_us_2024", + "state/ca": "populace_us_ca_2024", + "uk": "enhanced_frs_2024_25", + "us": "populace_us_2024", + } + + +@pytest.mark.parametrize( + "failure_point", + ["during_copy", "after_reconciliation", "after_validation"], +) +def test_failure_during_publication_rolls_back_every_catalog_write( + publication_engine: Engine, + failure_point: str, +) -> None: + def fail(point: str, _connection) -> None: + if point == failure_point: + raise RuntimeError("injected failure") + + with pytest.raises(RuntimeError, match="injected failure"): + publish_catalog( + publication_engine, + normalized_catalog(), + checkpoint=fail, + ) + + assert all(count == 0 for count in _counts(publication_engine).values()) + + +def test_same_version_with_changed_content_fails_without_mutation( + publication_engine: Engine, +) -> None: + catalog = normalized_catalog() + publish_catalog(publication_engine, catalog) + before_ids = _identifiers(publication_engine) + before_counts = _counts(publication_engine) + us = catalog.country("us") + changed_variable = replace(us.variables[0], label="Changed label") + changed_us = replace(us, variables=(changed_variable,)) + changed = replace( + catalog, + countries=(changed_us, catalog.country("uk")), + ) + + with pytest.raises(CatalogPublicationError, match="differs"): + publish_catalog(publication_engine, changed) + + assert _identifiers(publication_engine) == before_ids + assert _counts(publication_engine) == before_counts + + +def test_republish_preserves_existing_run_and_dataset_version_records( + publication_engine: Engine, +) -> None: + catalog = normalized_catalog() + publish_catalog(publication_engine, catalog) + + with Session(publication_engine) as session: + model = session.exec( + select(TaxBenefitModel).where(TaxBenefitModel.name == "policyengine-us") + ).one() + model_version = session.exec( + select(TaxBenefitModelVersion).where( + TaxBenefitModelVersion.model_id == model.id + ) + ).one() + input_dataset = session.exec( + select(Dataset).where( + Dataset.tax_benefit_model_version_id == model_version.id, + Dataset.name == "populace_us_2024", + ) + ).one() + region = session.exec( + select(Region).where( + Region.tax_benefit_model_version_id == model_version.id, + Region.code == "us", + ) + ).one() + output_dataset = Dataset( + name="existing-run-output", + description="Existing generated output", + storage_path="private-output-reference", + year=2025, + is_output_dataset=True, + tax_benefit_model_version_id=model_version.id, + ) + session.add(output_dataset) + session.flush() + dataset_version = DatasetVersion( + name="existing-user-version", + description="Existing independently versioned data", + dataset_id=output_dataset.id, + tax_benefit_model_id=model.id, + ) + simulation = Simulation( + dataset_id=input_dataset.id, + output_dataset_id=output_dataset.id, + tax_benefit_model_version_id=model_version.id, + region_id=region.id, + ) + report = Report( + label="Existing report", + country="us", + tax_benefit_model_id=model.id, + dataset_id=input_dataset.id, + region_id=region.id, + ) + session.add_all((dataset_version, simulation, report)) + session.flush() + report_run = ReportRun( + report_id=report.id, + country_package_version="existing-country-version", + policyengine_version="existing-policyengine-version", + ) + session.add(report_run) + session.commit() + protected_ids = ( + dataset_version.id, + simulation.id, + report.id, + report_run.id, + simulation.dataset_id, + simulation.output_dataset_id, + report.dataset_id, + ) + + publish_catalog(publication_engine, catalog) + + with Session(publication_engine) as session: + persisted_dataset_version = session.get(DatasetVersion, protected_ids[0]) + persisted_simulation = session.get(Simulation, protected_ids[1]) + persisted_report = session.get(Report, protected_ids[2]) + persisted_report_run = session.get(ReportRun, protected_ids[3]) + assert persisted_dataset_version is not None + assert persisted_simulation is not None + assert persisted_report is not None + assert persisted_report_run is not None + assert persisted_simulation.dataset_id == protected_ids[4] + assert persisted_simulation.output_dataset_id == protected_ids[5] + assert persisted_report.dataset_id == protected_ids[6] + + +def test_new_version_is_additive_and_concurrent_retries_serialize( + publication_engine: Engine, +) -> None: + first = normalized_catalog() + newer = normalized_catalog(policyengine_version="4.21.0") + publish_catalog(publication_engine, first) + + with ThreadPoolExecutor(max_workers=2) as executor: + results = list( + executor.map( + lambda _index: publish_catalog(publication_engine, newer), + range(2), + ) + ) + + assert [result.policyengine_version for result in results] == [ + "4.21.0", + "4.21.0", + ] + assert _counts(publication_engine) == { + "tax_benefit_models": 2, + "tax_benefit_model_versions": 4, + "variables": 4, + "parameter_nodes": 4, + "parameters": 4, + "parameter_values": 8, + "datasets": 6, + "dataset_versions": 0, + "regions": 10, + "simulations": 0, + "reports": 0, + "report_runs": 0, + } diff --git a/tests/integration/test_v2_catalog_publication_qualification.py b/tests/integration/test_v2_catalog_publication_qualification.py new file mode 100644 index 000000000..2a0fe1bde --- /dev/null +++ b/tests/integration/test_v2_catalog_publication_qualification.py @@ -0,0 +1,123 @@ +"""Production-scale qualification for the installed PolicyEngine.py catalog.""" + +from __future__ import annotations + +import gc +import json +import os +from pathlib import Path +import tracemalloc + +from alembic import command +from alembic.config import Config +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.engine import make_url +from sqlalchemy.pool import NullPool + +from policyengine_api.data.v2.catalog.extraction import extract_installed_catalog +from policyengine_api.data.v2.catalog.publication import publish_catalog +from policyengine_api.data.v2.settings import V2_MIGRATION_DATABASE_URL + + +REPO = Path(__file__).parents[2] +DISPOSABLE_DATABASE = "policyengine_v2_alembic_test" +MAX_ADDITIONAL_PUBLISHER_MEMORY_BYTES = 256 * 1024 * 1024 + +pytestmark = pytest.mark.skipif( + os.environ.get("RUN_V2_CATALOG_PUBLICATION_QUALIFICATION") != "1", + reason=( + "publishes the complete catalog; set RUN_V2_CATALOG_PUBLICATION_QUALIFICATION=1" + ), +) + + +def _disposable_url() -> str: + database_url = os.environ.get(V2_MIGRATION_DATABASE_URL, "") + if not database_url: + pytest.skip(f"{V2_MIGRATION_DATABASE_URL} is not set") + url = make_url(database_url) + if url.database != DISPOSABLE_DATABASE or url.host not in { + "localhost", + "127.0.0.1", + "::1", + "postgres", + }: + pytest.fail("catalog qualification requires disposable local Postgres") + return database_url + + +def test_complete_installed_catalog_bulk_publication() -> None: + database_url = _disposable_url() + command.upgrade(Config(str(REPO / "alembic-v2.ini")), "head") + engine = create_engine(database_url, poolclass=NullPool) + with engine.begin() as connection: + connection.execute(text("TRUNCATE tax_benefit_models CASCADE")) + + try: + catalog = extract_installed_catalog() + gc.collect() + tracemalloc.start() + evidence = publish_catalog(engine, catalog) + _, peak_bytes = tracemalloc.get_traced_memory() + tracemalloc.stop() + + assert peak_bytes < MAX_ADDITIONAL_PUBLISHER_MEMORY_BYTES + assert evidence.entity_counts == catalog.entity_counts() + with engine.connect() as connection: + persisted_counts = connection.execute( + text( + """ + SELECT + (SELECT count(*) FROM tax_benefit_models), + (SELECT count(*) FROM tax_benefit_model_versions), + (SELECT count(*) FROM variables), + (SELECT count(*) FROM parameter_nodes), + (SELECT count(*) FROM parameters), + (SELECT count(*) FROM parameter_values + WHERE policy_id IS NULL AND dynamic_id IS NULL), + (SELECT count(*) FROM datasets + WHERE NOT is_output_dataset AND storage_path IS NULL), + (SELECT count(*) FROM regions) + """ + ) + ).one() + representative = connection.execute( + text( + """ + SELECT + EXISTS (SELECT 1 FROM variables + WHERE name = 'employment_income'), + EXISTS (SELECT 1 FROM parameters + WHERE name = 'gov.benefit_uprating_cpi'), + EXISTS (SELECT 1 FROM regions WHERE code = 'us'), + EXISTS (SELECT 1 FROM regions WHERE code = 'uk') + """ + ) + ).one() + assert tuple(persisted_counts) == ( + evidence.entity_counts["models"], + evidence.entity_counts["model_versions"], + evidence.entity_counts["variables"], + evidence.entity_counts["parameter_nodes"], + evidence.entity_counts["parameters"], + evidence.entity_counts["parameter_values"], + evidence.entity_counts["datasets"], + evidence.entity_counts["regions"], + ) + assert all(representative) + print( + json.dumps( + { + **evidence.as_dict(), + "peak_additional_publisher_memory_bytes": peak_bytes, + }, + sort_keys=True, + ) + ) + finally: + if tracemalloc.is_tracing(): + tracemalloc.stop() + with engine.begin() as connection: + connection.execute(text("TRUNCATE tax_benefit_models CASCADE")) + engine.dispose() diff --git a/tests/integration/test_v2_metadata_routes.py b/tests/integration/test_v2_metadata_routes.py new file mode 100644 index 000000000..40a38b4ba --- /dev/null +++ b/tests/integration/test_v2_metadata_routes.py @@ -0,0 +1,246 @@ +"""Postgres-backed integration coverage for v2 metadata preview reads.""" + +from __future__ import annotations + +import os +from pathlib import Path +from uuid import uuid4 + +from alembic import command +from alembic.config import Config +from fastapi.testclient import TestClient +from flask import Flask, jsonify +import pytest +from sqlalchemy import create_engine, text +from sqlalchemy.engine import Engine, make_url +from sqlalchemy.pool import NullPool +from sqlmodel import Session, select + +from policyengine_api.asgi_factory import create_asgi_app +from policyengine_api.data.v2.catalog.publication import publish_catalog +from policyengine_api.data.v2.catalog.query import V2MetadataQueryService +from policyengine_api.data.v2.models import ( + Dataset, + TaxBenefitModel, + TaxBenefitModelVersion, +) +from policyengine_api.data.v2.settings import V2_MIGRATION_DATABASE_URL +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies +from policyengine_api.migration_flags import ( + RouteImplementation, + RouteImplementationSettings, +) +from tests.fixtures.v2_catalog import POLICYENGINE_VERSION, normalized_catalog + + +REPO = Path(__file__).parents[2] +DISPOSABLE_DATABASE = "policyengine_v2_alembic_test" + + +def _disposable_url() -> str: + database_url = os.environ.get(V2_MIGRATION_DATABASE_URL, "") + if not database_url: + pytest.skip(f"{V2_MIGRATION_DATABASE_URL} is not set") + url = make_url(database_url) + if url.database != DISPOSABLE_DATABASE or url.host not in { + "localhost", + "127.0.0.1", + "::1", + "postgres", + }: + pytest.fail("v2 preview route tests require disposable local Postgres") + return database_url + + +@pytest.fixture +def published_engine() -> Engine: + database_url = _disposable_url() + command.upgrade(Config(str(REPO / "alembic-v2.ini")), "head") + engine = create_engine(database_url, poolclass=NullPool) + with engine.begin() as connection: + connection.execute(text("TRUNCATE tax_benefit_models CASCADE")) + publish_catalog(engine, normalized_catalog()) + yield engine + with engine.begin() as connection: + connection.execute(text("TRUNCATE tax_benefit_models CASCADE")) + engine.dispose() + + +def _client(engine: Engine) -> TestClient: + flask_app = Flask(__name__) + + @flask_app.get("//metadata") + def v1_metadata(country_id: str): + return jsonify({"status": "ok", "result": {"country": country_id}}) + + dependencies = NativeRouteDependencies( + readiness_probe=lambda: True, + gateway_client_factory=lambda: None, + metadata_reader_factory=lambda: None, + specification_provider=lambda: {}, + v2_metadata_reader_factory=lambda: V2MetadataQueryService( + Session(engine), + running_policyengine_version=POLICYENGINE_VERSION, + ), + ) + settings = RouteImplementationSettings( + health=RouteImplementation.FLASK_FALLBACK, + specification=RouteImplementation.FLASK_FALLBACK, + metadata=RouteImplementation.FLASK_FALLBACK, + ) + return TestClient( + create_asgi_app( + flask_app, + dependencies=dependencies, + route_settings=settings, + ) + ) + + +def test_postgres_preview_returns_complete_us_and_uk_catalogs_without_writes( + published_engine: Engine, +) -> None: + client = _client(published_engine) + with published_engine.connect() as connection: + before = connection.execute( + text( + """ + SELECT + (SELECT count(*) FROM variables), + (SELECT count(*) FROM parameters), + (SELECT count(*) FROM parameter_values), + (SELECT count(*) FROM datasets), + (SELECT count(*) FROM regions) + """ + ) + ).one() + + us = client.get("/v2/us/metadata") + uk = client.get("/v2/uk/metadata") + + assert us.status_code == uk.status_code == 200 + assert us.json()["result"]["current_law_id"] == 2 + assert uk.json()["result"]["current_law_id"] == 1 + for country_id, response in (("us", us), ("uk", uk)): + result = response.json()["result"] + assert result["model"]["name"] == f"policyengine-{country_id}" + assert result["model_version"]["version"] == POLICYENGINE_VERSION + assert result["variables"] + assert result["parameter_nodes"] + assert result["parameters"] + assert result["parameters"][0]["values"] + assert result["datasets"] + assert result["regions"] + assert all( + not dataset["is_output_dataset"] and dataset["storage_path"] is None + for dataset in result["datasets"] + ) + assert all( + isinstance(period["name"], int) and isinstance(period["label"], str) + for period in result["economy_options"]["time_period"] + ) + + with published_engine.connect() as connection: + after = connection.execute( + text( + """ + SELECT + (SELECT count(*) FROM variables), + (SELECT count(*) FROM parameters), + (SELECT count(*) FROM parameter_values), + (SELECT count(*) FROM datasets), + (SELECT count(*) FROM regions) + """ + ) + ).one() + assert after == before + + +def test_postgres_preview_excludes_an_existing_output_dataset( + published_engine: Engine, +) -> None: + with Session(published_engine) as session: + model = session.exec( + select(TaxBenefitModel).where(TaxBenefitModel.name == "policyengine-us") + ).one() + session.add( + Dataset( + id=uuid4(), + tax_benefit_model_version_id=session.exec( + select(TaxBenefitModelVersion).where( + TaxBenefitModelVersion.model_id == model.id, + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ) + .one() + .id, + name="existing-output", + description="Generated output", + storage_path="private-output-reference", + year=2026, + is_output_dataset=True, + ) + ) + session.commit() + + response = _client(published_engine).get("/v2/us/metadata") + + assert response.status_code == 200 + assert "existing-output" not in { + dataset["name"] for dataset in response.json()["result"]["datasets"] + } + + +@pytest.mark.parametrize("country_id", ["us", "uk"]) +def test_postgres_preview_defaults_to_running_version_and_accepts_exact_version( + published_engine: Engine, + country_id: str, +) -> None: + publish_catalog( + published_engine, + normalized_catalog(policyengine_version="5.0.5"), + ) + client = _client(published_engine) + + default_response = client.get(f"/v2/{country_id}/metadata") + selected_response = client.get( + f"/v2/{country_id}/metadata", + params={"policyengine_version": "5.0.5"}, + ) + + assert default_response.status_code == selected_response.status_code == 200 + default_result = default_response.json()["result"] + selected_result = selected_response.json()["result"] + assert default_result["model_version"]["version"] == POLICYENGINE_VERSION + assert selected_result["model_version"]["version"] == "5.0.5" + assert ( + default_result["model_version"]["id"] != selected_result["model_version"]["id"] + ) + assert {dataset["id"] for dataset in default_result["datasets"]}.isdisjoint( + dataset["id"] for dataset in selected_result["datasets"] + ) + assert {region["id"] for region in default_result["regions"]}.isdisjoint( + region["id"] for region in selected_result["regions"] + ) + + +def test_postgres_preview_distinguishes_invalid_and_absent_versions( + published_engine: Engine, +) -> None: + client = _client(published_engine) + + invalid = client.get( + "/v2/us/metadata", + params={"policyengine_version": "not a version"}, + ) + absent = client.get( + "/v2/us/metadata", + params={"policyengine_version": "4.99.0"}, + ) + + assert invalid.status_code == 400 + assert invalid.json()["status"] == "error" + assert invalid.json()["message"] + assert absent.status_code == 404 + assert absent.json()["status"] == "error" + assert absent.json()["message"] diff --git a/tests/unit/routes/test_migration_context_logging.py b/tests/unit/routes/test_migration_context_logging.py index 38e107fb0..dde356f26 100644 --- a/tests/unit/routes/test_migration_context_logging.py +++ b/tests/unit/routes/test_migration_context_logging.py @@ -10,6 +10,7 @@ RouteImplementationSettings, ) from policyengine_api.migration_logging import register_migration_request_logging +from policyengine_api.migration_logging import log_migration_request from policyengine_api.request_context import ( REQUEST_ID_HEADER, current_request_id, @@ -253,6 +254,28 @@ def test_native_metadata_logs_country_and_actual_implementation(): assert log_payload["migration"]["route_impl"] == "fastapi_native" +def test_v2_metadata_preview_logs_its_actual_supabase_read_source(monkeypatch): + monkeypatch.setenv("DB_READ_METADATA", "invalid-unprefixed-setting") + monkeypatch.setenv("DB_WRITE_METADATA", "invalid-unprefixed-setting") + + with patch("policyengine_api.migration_logging.logger") as mock_logger: + log_migration_request( + request_id="request-123", + method="GET", + path="/v2/us/metadata", + status_code=200, + started_at=None, + country_id="us", + route_impl=RouteImplementation.FASTAPI_NATIVE, + ) + + migration_context = mock_logger.log_struct.call_args.args[0]["migration"] + assert migration_context["route_group"] == "metadata" + assert migration_context["route_impl"] == "fastapi_native" + assert migration_context["db_write"] is None + assert migration_context["db_read"] == "supabase" + + def test_native_route_failure_logs_country_and_actual_implementation(): dependencies = NativeRouteDependencies( readiness_probe=lambda: True, diff --git a/tests/unit/test_alembic_workflows.py b/tests/unit/test_alembic_workflows.py index d077934e2..8b84f4a1a 100644 --- a/tests/unit/test_alembic_workflows.py +++ b/tests/unit/test_alembic_workflows.py @@ -131,6 +131,12 @@ def test_reusable_v2_check_uses_disposable_postgres_and_real_redis(): assert "bash .github/scripts/test_alembic_v2_lifecycle.sh" in workflow assert "test_alembic_v2.py" in lifecycle_script assert "test_alembic_v2_lifecycle.py" in lifecycle_script + assert "test_v2_catalog_installed.py" in workflow + assert "RUN_V2_CATALOG_COMPATIBILITY" in workflow + assert "test_v2_catalog_publication.py" in workflow + assert "test_v2_metadata_routes.py" in workflow + assert "test_v2_catalog_publication_qualification.py" in workflow + assert "RUN_V2_CATALOG_PUBLICATION_QUALIFICATION" in workflow assert "test_runtime_cache_redis.py" in workflow assert "uv sync --frozen" in workflow diff --git a/tests/unit/test_cloud_run_deploy_scripts.py b/tests/unit/test_cloud_run_deploy_scripts.py index 6257f0f3f..f2c7d6acd 100644 --- a/tests/unit/test_cloud_run_deploy_scripts.py +++ b/tests/unit/test_cloud_run_deploy_scripts.py @@ -17,6 +17,9 @@ STAGING_CLOUD_RUN_SERVICE = "policyengine-api-staging" TEST_V2_PROJECT_REF = "abcdefghijklmnopqrst" TEST_V2_ENVIRONMENT = "test-foundation" +TEST_V2_RUNTIME_SECRET_RESOURCE = ( + "projects/test-project/secrets/v2-runtime-database-url/versions/latest" +) CLOUD_RUN_SERVICE_SCRIPTS = ( "scripts/deploy_cloud_run_candidate.sh", "scripts/capture_cloud_run_service_state.sh", @@ -106,6 +109,7 @@ def _v2_target_env() -> dict[str, str]: return { "V2_SUPABASE_PROJECT_REF": TEST_V2_PROJECT_REF, "V2_SUPABASE_ENVIRONMENT": TEST_V2_ENVIRONMENT, + "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE": (TEST_V2_RUNTIME_SECRET_RESOURCE), } @@ -744,7 +748,10 @@ def test_validate_app_engine_deploy_env_accepts_direct_mode_from_environment(): ) @pytest.mark.parametrize( "missing_name", - ["V2_SUPABASE_PROJECT_REF", "V2_SUPABASE_ENVIRONMENT"], + [ + "V2_SUPABASE_PROJECT_REF", + "V2_SUPABASE_ENVIRONMENT", + ], ) def test_deployment_validation_requires_supabase_target_variables( validation_script, @@ -759,6 +766,24 @@ def test_deployment_validation_requires_supabase_target_variables( assert missing_name in result.stderr +def test_only_cloud_run_requires_v2_runtime_database_configuration(): + env = _script_env(**_required_runtime_env()) + env.pop("V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE") + + app_engine = _run_script( + ".github/scripts/validate_app_engine_deploy_env.sh", + env, + ) + cloud_run = _run_script( + ".github/scripts/validate_cloud_run_deploy_env.sh", + env, + ) + + assert app_engine.returncode == 0, app_engine.stderr + assert cloud_run.returncode == 1 + assert "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE" in cloud_run.stderr + + @pytest.mark.parametrize( ("entrypoint", "selected_url_env", "selected_url"), [ @@ -840,8 +865,10 @@ def test_app_engine_bundle_contains_runtime_environment_placeholders(): assert 'RUNTIME_CACHE_MODE: "deployed"' in app_config assert 'V2_SUPABASE_PROJECT_REF: ".v2_supabase_project_ref"' in app_config assert 'V2_SUPABASE_ENVIRONMENT: ".v2_supabase_environment"' in app_config + assert "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE" not in app_config assert '"V2_SUPABASE_PROJECT_REF"' in export_script assert '"V2_SUPABASE_ENVIRONMENT"' in export_script + assert '"V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE"' not in export_script assert ( 'RUNTIME_CACHE_URL_SECRET_RESOURCE: ".runtime_cache_url_secret_resource"' in app_config @@ -867,6 +894,16 @@ def test_deployment_jobs_read_supabase_identity_from_github_environment_variable assert "V2_SUPABASE_PROJECT_REF: ${{ vars.V2_SUPABASE_PROJECT_REF }}" in job assert "V2_SUPABASE_ENVIRONMENT: ${{ vars.V2_SUPABASE_ENVIRONMENT }}" in job + for job_name in ("deploy-cloud-run-staging", "deploy-cloud-run-candidate"): + job = _workflow_job_block(workflow, job_name) + assert ( + "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: " + "${{ secrets.V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE }}" in job + ) + for job_name in ("deploy-staging", "deploy-production-candidate"): + job = _workflow_job_block(workflow, job_name) + assert "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE" not in job + assert TEST_V2_PROJECT_REF not in workflow assert TEST_V2_ENVIRONMENT not in workflow @@ -1022,6 +1059,7 @@ def test_app_engine_export_requires_only_selected_url( "projects/policyengine-api/secrets/" "policyengine-api-prod-runtime-cache-url/versions/latest" in rendered_app_config ) + assert TEST_V2_RUNTIME_SECRET_RESOURCE not in rendered_app_config for resource in APP_ENGINE_SECRET_RESOURCES.values(): assert resource in rendered_app_config assert not (tmp_path / ".dbpw").exists() @@ -1097,6 +1135,10 @@ def test_deploy_cloud_run_candidate_dry_run_never_shifts_traffic(): ) assert f"V2_SUPABASE_PROJECT_REF={TEST_V2_PROJECT_REF}" in result.stdout assert f"V2_SUPABASE_ENVIRONMENT={TEST_V2_ENVIRONMENT}" in result.stdout + assert ( + f"V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE=" + f"{TEST_V2_RUNTIME_SECRET_RESOURCE}" in result.stdout + ) assert "V2_DATABASE_URL" not in result.stdout assert "V2_STORAGE_ADMIN_KEY" not in result.stdout for env_name, secret_ref in CLOUD_RUN_SECRET_MAPPINGS.items(): @@ -1743,10 +1785,8 @@ def test_push_workflow_staging_fully_gates_all_production_deployments(): docker_publish = _workflow_job_block(workflow, "docker") cloud_run_production = _workflow_job_block(workflow, "deploy-cloud-run-candidate") - production_gate_dependency = ( - "needs: ensure-production-model-version-aligns-with-sim-api" - ) - assert production_gate_dependency in app_engine_candidate + production_initialization_dependency = "needs: initialize-v2-production" + assert production_initialization_dependency in app_engine_candidate assert 'APP_ENGINE_PROMOTE: "0"' in app_engine_candidate assert ( "bash .github/scripts/promote_app_engine_version.sh" not in app_engine_candidate @@ -1761,7 +1801,7 @@ def test_push_workflow_staging_fully_gates_all_production_deployments(): "${{ needs.deploy-production-candidate.outputs.version }}" in app_engine_promotion ) - assert production_gate_dependency in cloud_run_production + assert production_initialization_dependency in cloud_run_production assert "needs: promote-production" in docker_publish assert "stage3-prod-" in cloud_run_production assert "Build and push Cloud Run image" not in cloud_run_production diff --git a/tests/unit/test_migration_contract_artifacts.py b/tests/unit/test_migration_contract_artifacts.py index 014e3f8c1..93f39f193 100644 --- a/tests/unit/test_migration_contract_artifacts.py +++ b/tests/unit/test_migration_contract_artifacts.py @@ -10,8 +10,8 @@ def test_migration_contract_payload_summarizes_route_contracts(): assert payload["version"] == 1 assert payload["metadata"] == { "route_group_count": 9, - "workflow_count": 7, - "request_count": 14, + "workflow_count": 8, + "request_count": 16, "db_entity_count": 6, "sim_flow_count": 3, } @@ -20,6 +20,7 @@ def test_migration_contract_payload_summarizes_route_contracts(): "household_save_edit_read", "household_calculate", "region_selection", + "region_selection_v2_preview", "simulation_submit_poll", "report_create_poll", "budget_window_submit_poll", diff --git a/tests/unit/test_migration_flags.py b/tests/unit/test_migration_flags.py index 8368b0462..55dd39f78 100644 --- a/tests/unit/test_migration_flags.py +++ b/tests/unit/test_migration_flags.py @@ -115,7 +115,9 @@ def test_invalid_migration_flag_raises(monkeypatch): ("/health", "health"), ("/simulation-gateway-check", "health"), ("/readiness-check", "health"), + ("/v2/openapi.json", "specification"), ("/us/metadata", "metadata"), + ("/v2/us/metadata", "metadata"), ("/us/policy/1", "policy"), ("/us/policies", "policy"), ("/us/household/1", "household"), diff --git a/tests/unit/v2/test_alembic_v2.py b/tests/unit/v2/test_alembic_v2.py index dfb0f0dc9..815acf476 100644 --- a/tests/unit/v2/test_alembic_v2.py +++ b/tests/unit/v2/test_alembic_v2.py @@ -217,11 +217,12 @@ def test_v2_files_are_mechanically_separate_from_v1() -> None: assert all("migrations/v1" not in str(path) for path in v2_files) -def test_v2_revision_chain_is_one_generated_correction_bounded_baseline() -> None: +def test_v2_revision_chain_has_baseline_and_generated_stage_9_revision() -> None: config = Config(str(REPO / "alembic-v2.ini")) script = ScriptDirectory.from_config(config) - assert script.get_heads() == ["f5ef4347cb2a"] + assert script.get_heads() == ["68b4a5ae5dc5"] assert [revision.revision for revision in script.walk_revisions()] == [ + "68b4a5ae5dc5", "f5ef4347cb2a", ] @@ -268,6 +269,24 @@ def test_v2_revision_chain_is_one_generated_correction_bounded_baseline() -> Non "v2_simulation_type", } + stage_9_revision = ( + REPO + / "migrations/v2/versions/68b4a5ae5dc5_version_metadata_catalog_snapshots.py" + ).read_text(encoding="utf-8") + assert ( + "Generation: uv run alembic -c alembic-v2.ini revision --autogenerate" + in stage_9_revision + ) + assert 'down_revision: Union[str, None] = "f5ef4347cb2a"' in stage_9_revision + assert "uq_parameter_values_canonical_parameter_start_date" in stage_9_revision + assert "op.create_index(" in stage_9_revision + assert "op.drop_index(" in stage_9_revision + assert "tax_benefit_model_version_id" in stage_9_revision + assert "metadata_time_periods" in stage_9_revision + assert "current_law_id" in stage_9_revision + assert "op.execute(" not in stage_9_revision + assert "op.bulk_insert(" not in stage_9_revision + def test_alembic_rejects_unknown_missing_and_divergent_history(tmp_path: Path) -> None: original = REPO / "migrations/v2" @@ -276,9 +295,10 @@ def test_alembic_rejects_unknown_missing_and_divergent_history(tmp_path: Path) - (missing / "versions/f5ef4347cb2a_establish_v2_platform_baseline.py").unlink() missing_config = Config() missing_config.set_main_option("script_location", str(missing)) - missing_script = ScriptDirectory.from_config(missing_config) - with pytest.raises((CommandError, ResolutionError)): - missing_script.get_revision("f5ef4347cb2a") + with pytest.warns(UserWarning, match=r"Revision f5ef4347cb2a .* is not present"): + missing_script = ScriptDirectory.from_config(missing_config) + with pytest.raises((CommandError, KeyError, ResolutionError)): + missing_script.get_revision("f5ef4347cb2a") divergent = tmp_path / "divergent" shutil.copytree(original, divergent) diff --git a/tests/unit/v2/test_catalog_extraction.py b/tests/unit/v2/test_catalog_extraction.py new file mode 100644 index 000000000..baa59d528 --- /dev/null +++ b/tests/unit/v2/test_catalog_extraction.py @@ -0,0 +1,344 @@ +"""Focused extraction tests for the PolicyEngine.py v2 catalog.""" + +from __future__ import annotations + +import ast +from datetime import datetime, timezone +from importlib.metadata import PackageNotFoundError +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from policyengine_api.data.v2.catalog.extraction import ( + CatalogExtractionError, + extract_catalog, + normalize_json_value, +) +from tests.fixtures.v2_catalog import ( + DEPENDENCY_VERSIONS, + POLICYENGINE_VERSION, + bundle, + installed_version, + normalized_catalog, + source_models, +) + + +def test_extracts_canonical_policyengine_catalog_and_dataset_defaults() -> None: + catalog = normalized_catalog() + us = catalog.country("us") + uk = catalog.country("uk") + + assert catalog.policyengine_version == POLICYENGINE_VERSION + assert dict(catalog.dependency_versions) == DEPENDENCY_VERSIONS + assert us.model_version.version == POLICYENGINE_VERSION + assert us.model_version.version != DEPENDENCY_VERSIONS["policyengine-us"] + assert uk.model_version.version != DEPENDENCY_VERSIONS["policyengine-uk"] + + assert {dataset.name for dataset in us.datasets} == { + "populace_us_2024", + "populace_us_ca_2024", + } + assert {dataset.name for dataset in uk.datasets} == {"enhanced_frs_2024_25"} + assert all( + not dataset.is_output_dataset for dataset in (*us.datasets, *uk.datasets) + ) + assert all(dataset.storage_path is None for dataset in (*us.datasets, *uk.datasets)) + + us_datasets = {dataset.id: dataset.name for dataset in us.datasets} + us_defaults = { + region.code: us_datasets[region.default_dataset_id] for region in us.regions + } + assert us_defaults == { + "place/CA-44000": "populace_us_2024", + "state/ca": "populace_us_ca_2024", + "us": "populace_us_2024", + } + assert [ + (summary.region_type, summary.count) for summary in us.fallback_summaries + ] == [("place", 1)] + assert all(region.default_dataset_id == uk.datasets[0].id for region in uk.regions) + + +def test_durable_ids_are_deterministic_and_stable_at_the_intended_scope() -> None: + first = normalized_catalog() + repeated = normalized_catalog() + newer = normalized_catalog(policyengine_version="5.0.5") + + assert first == repeated + for country_id in ("us", "uk"): + old = first.country(country_id) + new = newer.country(country_id) + assert old.model.id == new.model.id + assert old.model_version.id != new.model_version.id + assert {dataset.id for dataset in old.datasets}.isdisjoint( + dataset.id for dataset in new.datasets + ) + assert {region.id for region in old.regions}.isdisjoint( + region.id for region in new.regions + ) + assert all( + dataset.model_version_id == old.model_version.id for dataset in old.datasets + ) + assert all( + region.model_version_id == old.model_version.id for region in old.regions + ) + + +def test_parameter_values_are_nested_and_iterated_in_bounded_batches() -> None: + us = normalized_catalog().country("us") + + assert len(us.parameters) == 1 + assert len(us.parameters[0].values) == 2 + assert [len(batch) for batch in us.parameter_value_batches(batch_size=1)] == [ + 1, + 1, + ] + assert [len(batch) for batch in us.parameter_value_batches(batch_size=100)] == [2] + with pytest.raises(ValueError, match="positive"): + tuple(us.parameter_value_batches(batch_size=0)) + + +def test_normalizes_supported_json_scalars_dates_and_containers() -> None: + class Scalar: + def item(self): + return 7 + + assert normalize_json_value( + { + "date": datetime(2026, 1, 1, tzinfo=timezone.utc), + "values": (Scalar(), True, 1.5), + } + ) == { + "date": "2026-01-01T00:00:00+00:00", + "values": [7, True, 1.5], + } + assert normalize_json_value(float("inf")) == "Infinity" + assert normalize_json_value(float("-inf")) == "-Infinity" + + +@pytest.mark.parametrize("value", [float("nan"), object(), {1: "x"}]) +def test_rejects_non_json_values(value: object) -> None: + with pytest.raises(CatalogExtractionError): + normalize_json_value(value) + + +def test_rejects_missing_or_mismatched_manifest_dependencies() -> None: + missing = bundle() + del missing["packages"]["policyengine-core"] + with pytest.raises(CatalogExtractionError, match="omits policyengine-core"): + extract_catalog( + bundle=missing, + policyengine_version=POLICYENGINE_VERSION, + models=source_models(), + installed_version=installed_version, + ) + + with pytest.raises(CatalogExtractionError, match="does not match"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=source_models(), + installed_version=lambda name: ( + "wrong" if name == "policyengine-us" else installed_version(name) + ), + ) + + def absent(_name: str) -> str: + raise PackageNotFoundError + + with pytest.raises(CatalogExtractionError, match="is absent"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=source_models(), + installed_version=absent, + ) + + +@pytest.mark.parametrize("version", ["", "0.0.0", "unknown"]) +def test_rejects_placeholder_policyengine_versions(version: str) -> None: + with pytest.raises(CatalogExtractionError, match="placeholder"): + extract_catalog( + bundle=bundle(policyengine_version=version), + policyengine_version=version, + models=source_models(policyengine_version=version), + installed_version=installed_version, + ) + + +def test_rejects_incomplete_public_model_catalogs() -> None: + models = source_models() + del models["uk"] + with pytest.raises(CatalogExtractionError, match="absent for: uk"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + models = source_models() + models["us"].release_manifest.default_dataset = "unreviewed_us_2025" + with pytest.raises(CatalogExtractionError, match="reviewed selection"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + +def test_ignores_only_the_unnamed_structural_parameter_root() -> None: + models = source_models() + models["uk"].parameter_nodes_by_name[""] = SimpleNamespace( + name="", label=None, description=None + ) + + catalog = extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + assert [node.name for node in catalog.country("uk").parameter_nodes] == [ + "gov.example" + ] + + +def test_uses_and_validates_public_name_indexed_mappings() -> None: + models = source_models() + variable = models["uk"].variables_by_name.pop("employment_income") + models["uk"].variables_by_name["wrong_name"] = variable + with pytest.raises(CatalogExtractionError, match="does not match record name"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + +def test_rejects_duplicate_keys_invalid_intervals_and_unknown_region_types() -> None: + duplicate = source_models() + duplicate["us"].parameters_by_name = [] + with pytest.raises(CatalogExtractionError, match="name-indexed mapping"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=duplicate, + installed_version=installed_version, + ) + + unknown_region = source_models() + unknown_region["us"].region_registry[1].region_type = "province" + with pytest.raises(CatalogExtractionError, match="unsupported type"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=unknown_region, + installed_version=installed_version, + ) + + invalid_interval = source_models() + invalid_interval["us"].parameters_by_name["gov.example.rate"].parameter_values[ + 0 + ].end_date = datetime(2026, 1, 1, tzinfo=timezone.utc) + with pytest.raises(CatalogExtractionError, match="canonical inclusive intervals"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=invalid_interval, + installed_version=installed_version, + ) + + +def test_rejects_parameter_values_that_are_not_oldest_to_newest() -> None: + models = source_models() + parameter = models["us"].parameters_by_name["gov.example.rate"] + parameter.parameter_values.reverse() + with pytest.raises(CatalogExtractionError, match="oldest to newest"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + +def test_rejects_duplicate_parameter_value_start_dates() -> None: + models = source_models() + parameter = models["us"].parameters_by_name["gov.example.rate"] + parameter.parameter_values.append(parameter.parameter_values[1]) + + with pytest.raises(CatalogExtractionError, match="duplicate value start date"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + +def test_preserves_consecutive_day_parameter_value_intervals() -> None: + models = source_models() + parameter = models["us"].parameters_by_name["gov.example.rate"] + parameter.parameter_values = [ + SimpleNamespace( + value=0.1, + start_date=datetime(2026, 1, 1, tzinfo=timezone.utc), + end_date=datetime(2026, 1, 1, tzinfo=timezone.utc), + ), + SimpleNamespace( + value=0.2, + start_date=datetime(2026, 1, 2, tzinfo=timezone.utc), + end_date=None, + ), + ] + + catalog = extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + values = catalog.country("us").parameters[0].values + assert values[0].start_date == values[0].end_date + assert values[1].start_date == datetime(2026, 1, 2, tzinfo=timezone.utc) + assert values[1].end_date is None + + +def test_extractor_has_no_direct_core_country_or_v1_metadata_imports() -> None: + source_path = ( + Path(__file__).parents[3] + / "policyengine_api" + / "data" + / "v2" + / "catalog" + / "extraction.py" + ) + tree = ast.parse(source_path.read_text(encoding="utf-8")) + imported_modules = { + alias.name + for node in ast.walk(tree) + if isinstance(node, ast.Import) + for alias in node.names + } | { + node.module or "" for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) + } + + assert not any( + module.startswith( + ( + "policyengine_core", + "policyengine_us", + "policyengine_uk", + "policyengine_api.country", + "policyengine_api.services.metadata_service", + ) + ) + for module in imported_modules + ) diff --git a/tests/unit/v2/test_catalog_initialization.py b/tests/unit/v2/test_catalog_initialization.py new file mode 100644 index 000000000..3dbdc85c3 --- /dev/null +++ b/tests/unit/v2/test_catalog_initialization.py @@ -0,0 +1,143 @@ +"""Command-boundary tests for explicit Stage 9 initialization.""" + +from __future__ import annotations + +import ast +from pathlib import Path + +import pytest + +from policyengine_api.data.v2.catalog import initialization +from policyengine_api.data.v2.catalog.publication import PublicationEvidence +from policyengine_api.data.v2.settings import ( + V2_DATA_WRITE_DATABASE_URL, + V2_SUPABASE_ENVIRONMENT, + V2_SUPABASE_PROJECT_REF, +) +from tests.fixtures.v2_catalog import ( + DEPENDENCY_VERSIONS, + POLICYENGINE_VERSION, + normalized_catalog, +) + + +ENVIRONMENT = { + V2_DATA_WRITE_DATABASE_URL: ( + "postgresql+psycopg://data-writer:test-password@db.example.com/" + "postgres?sslmode=require" + ), + V2_SUPABASE_PROJECT_REF: "abcdefghijklmnopqrst", + V2_SUPABASE_ENVIRONMENT: "test-foundation", + "V2_RUNTIME_DATABASE_URL": "invalid-runtime-value", + "V2_MIGRATION_DATABASE_URL": "invalid-migration-value", +} + + +class FakeEngine: + def __init__(self) -> None: + self.disposed = False + + def dispose(self) -> None: + self.disposed = True + + +def _evidence() -> PublicationEvidence: + return PublicationEvidence( + policyengine_version=POLICYENGINE_VERSION, + dependency_versions=( + ("policyengine-core", DEPENDENCY_VERSIONS["policyengine-core"]), + ), + entity_counts={"models": 2}, + fallback_summaries=(("us", "state", 1),), + elapsed_seconds=1.25, + ) + + +def test_initialization_uses_only_data_write_settings_and_disposes_engine() -> None: + engine = FakeEngine() + observed = {} + + def build(settings): + observed["username"] = settings.connection.url.username + return engine + + def publish(selected_engine, catalog): + observed["engine"] = selected_engine + observed["version"] = catalog.policyengine_version + return _evidence() + + result = initialization.initialize_catalog( + ENVIRONMENT, + extractor=normalized_catalog, + engine_builder=build, + publisher=publish, + ) + + assert result == _evidence() + assert observed == { + "username": "data-writer", + "engine": engine, + "version": POLICYENGINE_VERSION, + } + assert engine.disposed + + +def test_extraction_failure_opens_no_engine() -> None: + def fail_extraction(): + raise RuntimeError("extraction failed") + + def reject_engine(_settings): + raise AssertionError("engine must not be built") + + with pytest.raises(RuntimeError, match="extraction failed"): + initialization.initialize_catalog( + ENVIRONMENT, + extractor=fail_extraction, + engine_builder=reject_engine, + ) + + +def test_main_emits_typed_success_and_redacted_unexpected_failure( + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + monkeypatch.setattr(initialization, "initialize_catalog", _evidence) + assert initialization.main() == 0 + success = capsys.readouterr() + assert '"outcome": "ok"' in success.out + assert success.err == "" + + def fail(): + raise RuntimeError("postgresql://user:do-not-print@secret-host/database") + + monkeypatch.setattr(initialization, "initialize_catalog", fail) + assert initialization.main() == 1 + failure = capsys.readouterr() + assert '"outcome": "error"' in failure.err + assert "do-not-print" not in failure.err + assert "secret-host" not in failure.err + + +def test_initializer_is_not_imported_by_application_startup_modules() -> None: + repo = Path(__file__).parents[3] + for relative_path in ( + "policyengine_api/api.py", + "policyengine_api/asgi.py", + "policyengine_api/asgi_factory.py", + "policyengine_api/app_engine_runtime.py", + ): + tree = ast.parse((repo / relative_path).read_text(encoding="utf-8")) + imported = { + node.module or "" + for node in ast.walk(tree) + if isinstance(node, ast.ImportFrom) + } | { + alias.name + for node in ast.walk(tree) + if isinstance(node, ast.Import) + for alias in node.names + } + assert not any( + module.startswith("policyengine_api.data.v2.catalog.initialization") + for module in imported + ) diff --git a/tests/unit/v2/test_catalog_publication.py b/tests/unit/v2/test_catalog_publication.py new file mode 100644 index 000000000..b7e72ac6a --- /dev/null +++ b/tests/unit/v2/test_catalog_publication.py @@ -0,0 +1,221 @@ +"""Focused structural tests for PostgreSQL catalog publication.""" + +from __future__ import annotations + +from io import StringIO +from pathlib import Path +from unittest.mock import patch + +from alembic.config import Config +from alembic.script import ScriptDirectory +import pytest + +from policyengine_api.data.v2.catalog import publication + + +REPO = Path(__file__).parents[3] + + +class FakeCopy: + def __init__(self) -> None: + self.rows = [] + + def __enter__(self): + return self + + def __exit__(self, *_args) -> None: + return None + + def write_row(self, row) -> None: + self.rows.append(row) + + +class FakeCursor: + def __init__(self) -> None: + self.statement = None + self.copy_operation = FakeCopy() + + def __enter__(self): + return self + + def __exit__(self, *_args) -> None: + return None + + def copy(self, statement: str) -> FakeCopy: + self.statement = statement + return self.copy_operation + + +class FakeDriverConnection: + def __init__(self) -> None: + self.selected_cursor = FakeCursor() + + def cursor(self) -> FakeCursor: + return self.selected_cursor + + +class FakeConnection: + def __init__(self) -> None: + self.connection = type( + "ConnectionProxy", + (), + {"driver_connection": FakeDriverConnection()}, + )() + + +def test_expected_publication_revision_is_the_alembic_head() -> None: + config = Config(str(REPO / "alembic-v2.ini"), output_buffer=StringIO()) + script = ScriptDirectory.from_config(config) + + assert script.get_current_head() == publication.EXPECTED_ALEMBIC_REVISION + + +class _ScalarResult: + def __init__(self, value=None, values=()): + self.value = value + self.values = values + + def scalar_one(self): + return self.value + + def scalars(self): + return iter(self.values) + + +class _RevisionConnection: + def __init__(self, *, dialect: str, version_table=None, revisions=()): + self.dialect = type("Dialect", (), {"name": dialect})() + self.results = iter( + ( + _ScalarResult(value=version_table), + _ScalarResult(values=revisions), + ) + ) + + def execute(self, _statement): + return next(self.results) + + +def test_revision_check_rejects_non_postgres_missing_and_wrong_revisions() -> None: + with pytest.raises(publication.CatalogPublicationError, match="PostgreSQL"): + publication._verify_expected_revision(_RevisionConnection(dialect="sqlite")) + + with pytest.raises(publication.CatalogPublicationError, match="table is absent"): + publication._verify_expected_revision( + _RevisionConnection(dialect="postgresql", version_table=None) + ) + + with pytest.raises(publication.CatalogPublicationError, match="expected"): + publication._verify_expected_revision( + _RevisionConnection( + dialect="postgresql", + version_table="alembic_version", + revisions=("wrong-revision",), + ) + ) + + +def test_copy_streams_rows_through_psycopg_without_an_orm_write() -> None: + connection = FakeConnection() + rows = ((index, f"row-{index}") for index in range(3)) + + count = publication._copy_rows( + connection, + table_name="stage_catalog_models", + columns=("id", "name"), + rows=rows, + ) + + cursor = connection.connection.driver_connection.selected_cursor + assert count == 3 + assert cursor.statement == ("COPY stage_catalog_models (id, name) FROM STDIN") + assert cursor.copy_operation.rows == [ + (0, "row-0"), + (1, "row-1"), + (2, "row-2"), + ] + + +def test_publication_sql_uses_private_staging_and_set_based_inserts() -> None: + assert all( + "CREATE TEMP TABLE" in statement and "ON COMMIT DROP" in statement + for statement in publication.TEMP_TABLE_STATEMENTS + ) + assert all( + "INSERT INTO" in statement for statement in publication.SET_BASED_INSERT_SQL + ) + assert all("SELECT" in statement for statement in publication.SET_BASED_INSERT_SQL) + assert all( + "VALUES" not in statement for statement in publication.SET_BASED_INSERT_SQL + ) + + class Result: + def scalar_one(self): + return None + + class Connection: + statement = "" + parameters = {} + + def execute(self, statement, parameters): + self.statement = str(statement) + self.parameters = parameters + return Result() + + connection = Connection() + publication._acquire_publication_lock(connection) + assert "pg_advisory_xact_lock" in connection.statement + assert connection.parameters == { + "lock_key": publication.PUBLICATION_ADVISORY_LOCK_KEY + } + + +def test_completion_evidence_contains_only_reviewed_non_secret_fields() -> None: + evidence = publication.PublicationEvidence( + policyengine_version="4.20.3", + dependency_versions=(("policyengine-core", "3.30.0"),), + entity_counts={"parameters": 12}, + fallback_summaries=(("us", "state", 3),), + elapsed_seconds=2.3456, + ).as_dict() + + assert evidence == { + "outcome": "ok", + "policyengine_version": "4.20.3", + "dependency_versions": {"policyengine-core": "3.30.0"}, + "entity_counts": {"parameters": 12}, + "fallback_summaries": [ + {"country_id": "us", "region_type": "state", "count": 3} + ], + "elapsed_seconds": 2.346, + } + assert ( + not { + "database_url", + "credentials", + "parameter_values", + "artifact_location", + "dataset_release", + } + & evidence.keys() + ) + + +def test_fallback_summary_emits_one_non_secret_warning() -> None: + fallback_summaries = ( + ("us", "congressional_district", 436), + ("us", "place", 333), + ("us", "state", 51), + ) + + with patch.object(publication.LOGGER, "warning") as warning: + publication._log_fallback_warning(fallback_summaries) + + warning.assert_called_once_with( + "PolicyEngine.py regional dataset fallback summary: %s", + fallback_summaries, + ) + message_template, logged_summaries = warning.call_args.args + rendered_message = message_template % (logged_summaries,) + assert "regional dataset fallback summary" in rendered_message + assert "postgresql://" not in rendered_message diff --git a/tests/unit/v2/test_import_side_effects.py b/tests/unit/v2/test_import_side_effects.py index 3ee987b2c..eeb6be590 100644 --- a/tests/unit/v2/test_import_side_effects.py +++ b/tests/unit/v2/test_import_side_effects.py @@ -18,7 +18,9 @@ V2_ENVIRONMENT_NAMES = ( "V2_RUNTIME_DATABASE_URL", + "V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE", "V2_MIGRATION_DATABASE_URL", + "V2_DATA_WRITE_DATABASE_URL", "V2_SUPABASE_PROJECT_REF", "V2_SUPABASE_ENVIRONMENT", ) @@ -60,9 +62,11 @@ def reject_connect(*args, **kwargs): import policyengine_api.data.v2.settings import policyengine_api.data.v2.database from policyengine_api.data.v2.models import V2_METADATA +import sys after = set(pathlib.Path.cwd().iterdir()) assert before == after assert len(V2_METADATA.tables) == 32 +assert "policyengine_api.data.v2.catalog.initialization" not in sys.modules """ result = subprocess.run( @@ -104,10 +108,12 @@ def test_import_startup_and_request_never_create_runtime_sqlite( environment["FLASK_DEBUG"] = debug script = """ from pathlib import Path +import sys from policyengine_api.api import app response = app.test_client().get('/liveness-check') assert response.status_code == 200 +assert "policyengine_api.data.v2.catalog.initialization" not in sys.modules assert not Path('policyengine.db').exists() assert not list(Path.cwd().glob('*.db')) assert not list(Path.cwd().glob('*.init.lock')) diff --git a/tests/unit/v2/test_metadata_deployment.py b/tests/unit/v2/test_metadata_deployment.py new file mode 100644 index 000000000..cb6a7e806 --- /dev/null +++ b/tests/unit/v2/test_metadata_deployment.py @@ -0,0 +1,96 @@ +"""Deployment ordering and credential-isolation tests for Stage 9.""" + +from __future__ import annotations + +from pathlib import Path +import re +import subprocess + + +REPO = Path(__file__).resolve().parents[3] + + +def _read(relative_path: str) -> str: + return (REPO / relative_path).read_text(encoding="utf-8") + + +def _job(workflow: str, name: str) -> str: + match = re.search( + rf"^ {re.escape(name)}:\n(?P.*?)(?=^ [\w-]+:|\Z)", + workflow, + flags=re.MULTILINE | re.DOTALL, + ) + assert match is not None + return match.group("body") + + +def _step(workflow: str, name: str, next_name: str | None = None) -> str: + start = workflow.index(f" - name: {name}") + end = len(workflow) + if next_name is not None: + end = workflow.index(f" - name: {next_name}", start) + return workflow[start:end] + + +def test_reusable_initialization_workflow_separates_database_credentials() -> None: + workflow = _read(".github/workflows/initialize-v2-metadata.yml") + migration = _step( + workflow, + "Upgrade and verify the v2 schema", + "Publish and validate the v2 metadata catalog", + ) + publication = _step( + workflow, + "Publish and validate the v2 metadata catalog", + ) + + assert "workflow_call:" in workflow + assert "workflow_dispatch:" in workflow + assert "environment: ${{ inputs.deployment_environment }}" in workflow + assert "V2_MIGRATION_DATABASE_URL" in migration + assert "V2_DATA_WRITE_DATABASE_URL" not in migration + assert "V2_DATA_WRITE_DATABASE_URL" in publication + assert "V2_MIGRATION_DATABASE_URL" not in publication + assert "V2_RUNTIME_DATABASE_URL" not in workflow + assert "scripts/initialize_v2_metadata.py" in publication + + +def test_schema_upgrade_precedes_atomic_catalog_publication() -> None: + script = _read(".github/scripts/migrate_v2_metadata_schema.sh") + upgrade = script.index("upgrade head") + current = script.index("current --check-heads") + drift = script.index("alembic -c alembic-v2.ini check") + + assert "set -euo pipefail" in script + assert upgrade < current < drift + + +def test_initialization_success_is_required_before_candidate_creation() -> None: + workflow = _read(".github/workflows/push.yml") + staging_initialization = _job(workflow, "initialize-v2-staging") + production_initialization = _job(workflow, "initialize-v2-production") + + assert "deployment_environment: staging" in staging_initialization + assert "migrate-v1-cloud-sql" in staging_initialization + for job_name in ("deploy-staging", "deploy-cloud-run-staging"): + assert "initialize-v2-staging" in _job(workflow, job_name) + + assert ( + "needs: ensure-production-model-version-aligns-with-sim-api" + in production_initialization + ) + assert "deployment_environment: production" in production_initialization + for job_name in ("deploy-production-candidate", "deploy-cloud-run-candidate"): + assert "needs: initialize-v2-production" in _job(workflow, job_name) + + +def test_stage_9_deployment_shell_script_is_syntax_valid() -> None: + result = subprocess.run( + ["bash", "-n", ".github/scripts/migrate_v2_metadata_schema.sh"], + cwd=REPO, + capture_output=True, + text=True, + check=False, + ) + + assert result.returncode == 0, result.stderr diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py new file mode 100644 index 000000000..dc0f4b194 --- /dev/null +++ b/tests/unit/v2/test_metadata_query.py @@ -0,0 +1,488 @@ +"""Typed read-only query coverage for the v2 metadata preview.""" + +from __future__ import annotations + +import ast +from pathlib import Path +from uuid import uuid4 + +from pydantic import TypeAdapter, ValidationError +import pytest +from sqlalchemy.pool import StaticPool +from sqlmodel import Session, create_engine, select + +from policyengine_api.data.v2.catalog.query import ( + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, + UnsupportedPreviewCountryError, + V2MetadataQueryService, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataErrorResponse, + MetadataPreviewResponse, + MetadataSuccessResponse, +) +from policyengine_api.data.v2.models import ( + Dataset, + Parameter, + ParameterNode, + ParameterValue, + Region, + RegionType, + TaxBenefitModel, + TaxBenefitModelVersion, + V2_METADATA, + Variable, +) +from tests.fixtures.v2_catalog import POLICYENGINE_VERSION, normalized_catalog + + +@pytest.fixture +def catalog_session() -> Session: + engine = create_engine( + "sqlite://", + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) + V2_METADATA.create_all(engine) + catalog = normalized_catalog() + with Session(engine) as session: + for country in catalog.countries: + session.add( + TaxBenefitModel( + id=country.model.id, + name=country.model.name, + description=country.model.description, + ) + ) + session.add( + TaxBenefitModelVersion( + id=country.model_version.id, + model_id=country.model.id, + version=country.model_version.version, + description=country.model_version.description, + current_law_id=country.model_version.current_law_id, + metadata_time_periods=list( + country.model_version.metadata_time_periods + ), + ) + ) + session.add_all( + Variable( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + label=record.label, + entity=record.entity, + description=record.description, + data_type=record.data_type, + possible_values=record.possible_values, + default_value=record.default_value, + adds=record.adds, + subtracts=record.subtracts, + ) + for record in country.variables + ) + session.add_all( + ParameterNode( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + label=record.label, + description=record.description, + ) + for record in country.parameter_nodes + ) + session.add_all( + Parameter( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + label=record.label, + description=record.description, + data_type=record.data_type, + unit=record.unit, + ) + for record in country.parameters + ) + session.add_all( + ParameterValue( + id=value.id, + parameter_id=value.parameter_id, + value_json=value.value_json, + start_date=value.start_date, + end_date=value.end_date, + ) + for parameter in country.parameters + for value in parameter.values + ) + session.add_all( + Dataset( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + description=record.description, + year=record.year, + ) + for record in country.datasets + ) + session.add_all( + Region( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + default_dataset_id=record.default_dataset_id, + code=record.code, + label=record.label, + region_type=RegionType(record.region_type), + requires_filter=record.requires_filter, + filter_field=record.filter_field, + filter_value=record.filter_value, + filter_strategy=record.filter_strategy, + parent_code=record.parent_code, + state_code=record.state_code, + state_name=record.state_name, + ) + for record in country.regions + ) + session.commit() + + session = Session(engine) + try: + yield session + finally: + session.close() + engine.dispose() + + +def _add_country_version( + session: Session, + *, + policyengine_version: str, + current_law_id: int, + time_periods: list[int], +) -> None: + country = normalized_catalog(policyengine_version=policyengine_version).country( + "us" + ) + session.add( + TaxBenefitModelVersion( + id=country.model_version.id, + model_id=country.model.id, + version=policyengine_version, + description=f"US model for {policyengine_version}", + current_law_id=current_law_id, + metadata_time_periods=time_periods, + ) + ) + session.add_all( + Variable( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + label=record.label, + entity=record.entity, + description=record.description, + data_type=record.data_type, + possible_values=record.possible_values, + default_value=record.default_value, + adds=record.adds, + subtracts=record.subtracts, + ) + for record in country.variables + ) + session.add_all( + ParameterNode( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + label=record.label, + description=record.description, + ) + for record in country.parameter_nodes + ) + session.add_all( + Parameter( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + label=record.label, + description=record.description, + data_type=record.data_type, + unit=record.unit, + ) + for record in country.parameters + ) + session.add_all( + ParameterValue( + id=value.id, + parameter_id=value.parameter_id, + value_json=value.value_json, + start_date=value.start_date, + end_date=value.end_date, + ) + for parameter in country.parameters + for value in parameter.values + ) + session.add_all( + Dataset( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + name=record.name, + description=f"{record.description} for {policyengine_version}", + year=record.year, + ) + for record in country.datasets + ) + session.add_all( + Region( + id=record.id, + tax_benefit_model_version_id=country.model_version.id, + default_dataset_id=record.default_dataset_id, + code=record.code, + label=f"{record.label} {policyengine_version}", + region_type=RegionType(record.region_type), + requires_filter=record.requires_filter, + filter_field=record.filter_field, + filter_value=record.filter_value, + filter_strategy=record.filter_strategy, + parent_code=record.parent_code, + state_code=record.state_code, + state_name=record.state_name, + ) + for record in country.regions + ) + session.commit() + + +def test_query_serializes_complete_typed_metadata_without_writes( + catalog_session: Session, +) -> None: + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + model_classes = ( + TaxBenefitModel, + TaxBenefitModelVersion, + Variable, + ParameterNode, + Parameter, + ParameterValue, + Dataset, + Region, + ) + before = { + model_class.__tablename__: len(catalog_session.exec(select(model_class)).all()) + for model_class in model_classes + } + + result = service.get_metadata("us") + + assert result.current_law_id == 2 + assert result.model.name == "policyengine-us" + assert result.model_version.version == POLICYENGINE_VERSION + assert [variable.name for variable in result.variables] == ["employment_income"] + assert [parameter.name for parameter in result.parameters] == ["gov.example.rate"] + assert [value.value for value in result.parameters[0].values] == [0.1, 0.2] + assert {dataset.name for dataset in result.datasets} == { + "populace_us_2024", + "populace_us_ca_2024", + } + assert all(not dataset.is_output_dataset for dataset in result.datasets) + assert all(dataset.storage_path is None for dataset in result.datasets) + assert result.economy_options.region[0].name == "place/CA-44000" + assert result.economy_options.time_period[0].name == 2035 + assert result.economy_options.time_period[-1].name == 2022 + assert [option.name for option in result.economy_options.datasets] == [ + "populace_us_2024" + ] + after = { + model_class.__tablename__: len(catalog_session.exec(select(model_class)).all()) + for model_class in model_classes + } + assert after == before + + +def test_query_excludes_output_datasets(catalog_session: Session) -> None: + model = catalog_session.exec( + select(TaxBenefitModel).where(TaxBenefitModel.name == "policyengine-us") + ).one() + catalog_session.add( + Dataset( + id=uuid4(), + tax_benefit_model_version_id=catalog_session.exec( + select(TaxBenefitModelVersion).where( + TaxBenefitModelVersion.model_id == model.id, + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ) + .one() + .id, + name="simulation-output", + description="Generated result", + storage_path="private-output-reference", + year=2026, + is_output_dataset=True, + ) + ) + catalog_session.commit() + + result = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ).get_metadata("us") + + assert "simulation-output" not in {dataset.name for dataset in result.datasets} + + +def test_query_defaults_to_running_version_and_allows_exact_override( + catalog_session: Session, +) -> None: + _add_country_version( + catalog_session, + policyengine_version="5.0.5", + current_law_id=22, + time_periods=[2041, 2040], + ) + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + default = service.get_metadata("us") + selected = service.get_metadata("us", "5.0.5") + + assert default.model_version.version == POLICYENGINE_VERSION + assert default.current_law_id == 2 + assert default.economy_options.time_period[0].name == 2035 + assert selected.model_version.version == "5.0.5" + assert selected.model.description == "US model for 5.0.5" + assert selected.current_law_id == 22 + assert [option.name for option in selected.economy_options.time_period] == [ + 2041, + 2040, + ] + assert {dataset.id for dataset in default.datasets}.isdisjoint( + dataset.id for dataset in selected.datasets + ) + assert {region.id for region in default.regions}.isdisjoint( + region.id for region in selected.regions + ) + + +def test_query_rejects_invalid_or_absent_selected_versions( + catalog_session: Session, +) -> None: + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + for invalid in ("", f" {POLICYENGINE_VERSION}", "not a version", "0.0.0"): + with pytest.raises(InvalidPolicyEngineVersionError): + service.get_metadata("us", invalid) + with pytest.raises(MetadataCatalogVersionNotFoundError): + service.get_metadata("us", "4.99.0") + with pytest.raises(MetadataCatalogUnavailableError): + V2MetadataQueryService( + catalog_session, + running_policyengine_version="4.99.0", + ).get_metadata("us") + + +def test_query_rejects_unsupported_and_incomplete_catalogs( + catalog_session: Session, +) -> None: + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + with pytest.raises(UnsupportedPreviewCountryError): + service.get_metadata("ca") + + uk_model = catalog_session.exec( + select(TaxBenefitModel).where(TaxBenefitModel.name == "policyengine-uk") + ).one() + uk_model_version = catalog_session.exec( + select(TaxBenefitModelVersion).where( + TaxBenefitModelVersion.model_id == uk_model.id, + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ).one() + for region in catalog_session.exec( + select(Region).where(Region.tax_benefit_model_version_id == uk_model_version.id) + ).all(): + catalog_session.delete(region) + catalog_session.commit() + with pytest.raises(MetadataCatalogUnavailableError, match="incomplete"): + service.get_metadata("uk") + + +def test_response_outcomes_are_discriminated_and_strict( + catalog_session: Session, +) -> None: + result = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ).get_metadata("uk") + adapter = TypeAdapter(MetadataPreviewResponse) + + success = adapter.validate_python( + MetadataSuccessResponse(result=result).model_dump() + ) + error = adapter.validate_python( + MetadataErrorResponse(message="Catalog unavailable").model_dump() + ) + assert success.status == "ok" + assert error.status == "error" + + with pytest.raises(ValidationError): + adapter.validate_python({"status": "ok", "message": None}) + with pytest.raises(ValidationError): + adapter.validate_python({"status": "error", "message": ""}) + with pytest.raises(ValidationError): + adapter.validate_python({"status": "error", "message": " "}) + with pytest.raises(ValidationError): + adapter.validate_python( + { + "status": "error", + "message": "Failure", + "result": result.model_dump(), + } + ) + with pytest.raises(ValidationError): + adapter.validate_python({"status": "pending", "message": "Wait"}) + + +def test_query_module_imports_no_policyengine_or_v1_metadata_source() -> None: + source_path = ( + Path(__file__).parents[3] + / "policyengine_api" + / "data" + / "v2" + / "catalog" + / "query.py" + ) + tree = ast.parse(source_path.read_text(encoding="utf-8")) + imported = { + alias.name + for node in ast.walk(tree) + if isinstance(node, ast.Import) + for alias in node.names + } | { + node.module or "" for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) + } + assert not any( + module.startswith( + ( + "policyengine_core", + "policyengine_us", + "policyengine_uk", + "policyengine_api.country", + "policyengine_api.services.metadata_service", + "policyengine_api.data.v2.catalog.extraction", + ) + ) + for module in imported + ) diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py new file mode 100644 index 000000000..e5cac18e5 --- /dev/null +++ b/tests/unit/v2/test_metadata_routes.py @@ -0,0 +1,371 @@ +"""Typed route coverage for the dormant v2 metadata preview.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from uuid import uuid4 + +from fastapi.testclient import TestClient +from flask import Flask, jsonify +import pytest + +from policyengine_api.asgi_factory import create_asgi_app +from policyengine_api.data.v2.catalog.query import ( + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDataset, + MetadataEconomyOptions, + MetadataModel, + MetadataModelVersion, + MetadataParameter, + MetadataParameterNode, + MetadataParameterValue, + MetadataRegion, + MetadataRegionOption, + MetadataResult, + MetadataTimePeriodOption, + MetadataVariable, +) +from policyengine_api.data.v2.settings import V2ConfigurationError +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies +from policyengine_api.migration_flags import ( + RouteImplementation, + RouteImplementationSettings, +) + + +class Reader: + def __init__(self, result: MetadataResult, error: Exception | None = None): + self.result = result + self.error = error + self.calls = [] + self.closed = False + + def get_metadata( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataResult: + self.calls.append((country_id, policyengine_version)) + if self.error is not None: + raise self.error + return self.result + + def close(self) -> None: + self.closed = True + + +def _result(country_id: str = "us") -> MetadataResult: + model_id = uuid4() + version_id = uuid4() + dataset_id = uuid4() + region_id = uuid4() + parameter_id = uuid4() + return MetadataResult( + current_law_id=2 if country_id == "us" else 1, + model=MetadataModel( + id=model_id, + name=f"policyengine-{country_id}", + description="Model", + ), + model_version=MetadataModelVersion( + id=version_id, + model_id=model_id, + version="4.20.3", + description="PolicyEngine.py catalog", + ), + variables=[ + MetadataVariable( + id=uuid4(), + name="employment_income", + label="Employment income", + entity="person", + description=None, + data_type="float", + possible_values=None, + default_value=0, + adds=None, + subtracts=None, + ) + ], + parameter_nodes=[ + MetadataParameterNode( + id=uuid4(), + name="gov.example", + label="Example", + description=None, + ) + ], + parameters=[ + MetadataParameter( + id=parameter_id, + name="gov.example.rate", + label="Rate", + description=None, + data_type="float", + unit="/1", + values=[ + MetadataParameterValue( + id=uuid4(), + value=0.1, + start_date=datetime(2025, 1, 1, tzinfo=timezone.utc), + end_date=None, + ) + ], + ) + ], + datasets=[ + MetadataDataset( + id=dataset_id, + name=f"populace_{country_id}_2024", + description="Populace", + year=2024, + ) + ], + regions=[ + MetadataRegion( + id=region_id, + code=country_id, + label=country_id.upper(), + region_type="national", + requires_filter=False, + filter_field=None, + filter_value=None, + filter_strategy=None, + parent_code=None, + state_code=None, + state_name=None, + default_dataset_id=dataset_id, + ) + ], + economy_options=MetadataEconomyOptions( + region=[ + MetadataRegionOption( + name=country_id, + label=country_id.upper(), + type="national", + ) + ], + time_period=[MetadataTimePeriodOption(name=2026, label="2026")], + datasets=[], + ), + ) + + +def _client(factory) -> TestClient: + flask_app = Flask(__name__) + + @flask_app.get("//metadata") + def v1_metadata(country_id: str): + return jsonify( + { + "status": "ok", + "message": None, + "result": {"source": "v1", "country_id": country_id}, + } + ) + + dependencies = NativeRouteDependencies( + readiness_probe=lambda: True, + gateway_client_factory=lambda: None, + metadata_reader_factory=lambda: None, + specification_provider=lambda: {}, + v2_metadata_reader_factory=factory, + ) + settings = RouteImplementationSettings( + health=RouteImplementation.FLASK_FALLBACK, + specification=RouteImplementation.FLASK_FALLBACK, + metadata=RouteImplementation.FLASK_FALLBACK, + ) + return TestClient( + create_asgi_app( + flask_app, + dependencies=dependencies, + route_settings=settings, + ), + raise_server_exceptions=False, + ) + + +@pytest.mark.parametrize("country_id", ["us", "uk"]) +def test_preview_get_returns_typed_catalog_response(country_id: str) -> None: + readers = [] + + def factory(): + reader = Reader(_result(country_id)) + readers.append(reader) + return reader + + response = _client(factory).get(f"/v2/{country_id}/metadata") + + assert response.status_code == 200 + assert response.headers["content-type"].startswith("application/json") + payload = response.json() + assert payload["status"] == "ok" + assert payload["message"] is None + assert payload["result"]["current_law_id"] == (2 if country_id == "us" else 1) + assert payload["result"]["model_version"]["version"] == "4.20.3" + assert payload["result"]["economy_options"]["region"][0]["name"] == country_id + assert isinstance( + payload["result"]["economy_options"]["time_period"][0]["name"], + int, + ) + assert readers[0].calls == [(country_id, None)] + assert readers[0].closed + + +def test_unsupported_country_and_methods_return_typed_client_errors() -> None: + calls = [] + + def factory(): + calls.append("called") + return Reader(_result()) + + client = _client(factory) + country_response = client.get("/v2/ca/metadata") + assert country_response.status_code == 404 + assert country_response.json()["status"] == "error" + assert country_response.json()["message"] + + for method in ("POST", "PUT", "PATCH", "DELETE", "OPTIONS"): + response = client.request(method, "/v2/us/metadata") + assert response.status_code == 405 + assert response.json()["status"] == "error" + assert response.json()["message"] + assert calls == [] + + +@pytest.mark.parametrize( + ("error", "expected_status"), + [ + (MetadataCatalogUnavailableError("missing"), 503), + (V2ConfigurationError("missing URL"), 503), + (RuntimeError("private database detail"), 500), + ], +) +def test_preview_failures_are_typed_and_hide_internal_details( + error: Exception, + expected_status: int, +) -> None: + reader = Reader(_result(), error=error) + response = _client(lambda: reader).get("/v2/us/metadata") + + assert response.status_code == expected_status + assert response.json()["status"] == "error" + assert response.json()["message"] + assert "private database detail" not in response.text + assert reader.closed + + +@pytest.mark.parametrize( + ("version", "error", "expected_status"), + [ + ( + "not a version", + InvalidPolicyEngineVersionError("invalid PolicyEngine.py version"), + 400, + ), + ( + "4.99.0", + MetadataCatalogVersionNotFoundError( + "PolicyEngine.py 4.99.0 is not published for us" + ), + 404, + ), + ], +) +def test_preview_version_selector_returns_typed_client_errors( + version: str, + error: Exception, + expected_status: int, +) -> None: + reader = Reader(_result(), error=error) + + response = _client(lambda: reader).get( + "/v2/us/metadata", + params={"policyengine_version": version}, + ) + + assert response.status_code == expected_status + assert response.json()["status"] == "error" + assert response.json()["message"] + assert reader.calls == [("us", version)] + assert reader.closed + + +def test_preview_passes_explicit_version_to_reader() -> None: + result = _result() + result.model_version.version = "4.19.0" + reader = Reader(result) + + response = _client(lambda: reader).get( + "/v2/us/metadata", + params={"policyengine_version": "4.19.0"}, + ) + + assert response.status_code == 200 + assert response.json()["result"]["model_version"]["version"] == "4.19.0" + assert reader.calls == [("us", "4.19.0")] + + +def test_preview_reads_repeat_without_changing_v1_routing() -> None: + reader = Reader(_result()) + client = _client(lambda: reader) + + first = client.get("/v2/us/metadata") + second = client.get("/v2/us/metadata") + v1 = client.get("/us/metadata") + + assert first.json() == second.json() + assert reader.calls == [("us", None), ("us", None)] + assert v1.json()["result"] == {"source": "v1", "country_id": "us"} + + +def test_openapi_references_explicit_preview_response_schemas() -> None: + client = _client(lambda: Reader(_result())) + response = client.get("/v2/openapi.json") + + assert response.status_code == 200 + assert response.headers["content-type"].startswith("application/json") + schema = response.json() + assert set(schema["paths"]) == { + "/v2/us/metadata", + "/v2/uk/metadata", + "/v2/{country_id}/metadata", + } + + for path in ("/v2/us/metadata", "/v2/uk/metadata"): + operation = schema["paths"][path]["get"] + assert set(operation["responses"]) >= { + "200", + "400", + "404", + "405", + "500", + "503", + } + assert ( + operation["responses"]["200"]["content"]["application/json"]["schema"][ + "$ref" + ] + == "#/components/schemas/MetadataSuccessResponse" + ) + for status in ("400", "404", "405", "500", "503"): + assert ( + operation["responses"][status]["content"]["application/json"]["schema"][ + "$ref" + ] + == "#/components/schemas/MetadataErrorResponse" + ) + + unsupported = schema["paths"]["/v2/{country_id}/metadata"] + assert set(unsupported) == {"get"} + assert ( + unsupported["get"]["responses"]["404"]["content"]["application/json"]["schema"][ + "$ref" + ] + == "#/components/schemas/MetadataErrorResponse" + ) diff --git a/tests/unit/v2/test_model_persistence.py b/tests/unit/v2/test_model_persistence.py index b8879b1e7..efbaa5c05 100644 --- a/tests/unit/v2/test_model_persistence.py +++ b/tests/unit/v2/test_model_persistence.py @@ -1,5 +1,6 @@ """Canonical SQLModel persistence and bounded SQLAlchemy escape-hatch tests.""" +from datetime import datetime, timezone from pathlib import Path import pytest @@ -7,6 +8,8 @@ from sqlmodel import Session, create_engine, select from policyengine_api.data.v2.models import ( + Parameter, + ParameterValue, TaxBenefitModel, TaxBenefitModelVersion, User, @@ -21,7 +24,12 @@ def test_ordinary_persistence_uses_sqlmodel_session_select_and_exec() -> None: engine = create_engine("sqlite://") V2_METADATA.create_all(engine) model = TaxBenefitModel(name="test-country", description="Test model") - version = TaxBenefitModelVersion(model=model, version="1.2.3") + version = TaxBenefitModelVersion( + model=model, + version="1.2.3", + current_law_id=1, + metadata_time_periods=[2026], + ) with Session(engine) as session: session.add(version) @@ -76,6 +84,43 @@ def test_user_primary_country_can_change_between_us_and_uk_only() -> None: engine.dispose() +def test_canonical_parameter_values_are_unique_by_parameter_and_start_date() -> None: + engine = create_engine("sqlite://") + V2_METADATA.create_all(engine) + model = TaxBenefitModel(name="canonical-values") + version = TaxBenefitModelVersion( + model=model, + version="4.20.3", + current_law_id=1, + metadata_time_periods=[2026], + ) + parameter = Parameter( + name="gov.example.rate", + tax_benefit_model_version=version, + ) + start_date = datetime(2026, 1, 1, tzinfo=timezone.utc) + + with Session(engine) as session: + session.add_all( + [ + ParameterValue( + parameter=parameter, + value_json=0.1, + start_date=start_date, + ), + ParameterValue( + parameter=parameter, + value_json=0.2, + start_date=start_date, + ), + ] + ) + with pytest.raises(IntegrityError): + session.commit() + + engine.dispose() + + def test_v2_models_do_not_create_a_parallel_sqlalchemy_orm_layer() -> None: models_directory = ( Path(__file__).parents[3] / "policyengine_api" / "data" / "v2" / "models" diff --git a/tests/unit/v2/test_models.py b/tests/unit/v2/test_models.py index 3fd8cd1f5..dbdab10ee 100644 --- a/tests/unit/v2/test_models.py +++ b/tests/unit/v2/test_models.py @@ -7,11 +7,13 @@ from policyengine_api.data.v1_models import V1Base from policyengine_api.data.v2.models import ( + DatasetVersion, Dynamic, Household, HouseholdJob, Policy, Simulation, + TaxBenefitModelVersion, User, UserHouseholdAssociation, UserPolicy, @@ -197,7 +199,7 @@ def test_user_primary_country_is_required_and_limited_to_supported_values() -> N } -def test_regions_have_one_same_model_default_logical_dataset() -> None: +def test_regions_have_one_same_model_version_default_logical_dataset() -> None: regions = V2_METADATA.tables["regions"] datasets = V2_METADATA.tables["datasets"] @@ -207,15 +209,15 @@ def test_regions_have_one_same_model_default_logical_dataset() -> None: default_constraint = next( constraint for constraint in regions.foreign_key_constraints - if constraint.name == "fk_regions_default_dataset_model_datasets" + if constraint.name == "fk_regions_default_dataset_model_version" ) assert [element.parent.name for element in default_constraint.elements] == [ "default_dataset_id", - "tax_benefit_model_id", + "tax_benefit_model_version_id", ] assert [element.target_fullname for element in default_constraint.elements] == [ "datasets.id", - "datasets.tax_benefit_model_id", + "datasets.tax_benefit_model_version_id", ] assert default_constraint.ondelete == "RESTRICT" @@ -224,8 +226,12 @@ def test_regions_have_one_same_model_default_logical_dataset() -> None: for constraint in datasets.constraints if isinstance(constraint, sa.UniqueConstraint) } - assert ("tax_benefit_model_id", "name") in unique_column_sets - assert ("id", "tax_benefit_model_id") in unique_column_sets + assert ("tax_benefit_model_version_id", "name") in unique_column_sets + assert ("id", "tax_benefit_model_version_id") in unique_column_sets + assert "tax_benefit_model_id" not in datasets.c + assert "tax_benefit_model_id" not in regions.c + assert not datasets.c.tax_benefit_model_version_id.nullable + assert not regions.c.tax_benefit_model_version_id.nullable assert datasets.c.storage_path.nullable assert "ck_datasets_output_storage_path" in { constraint.name for constraint in datasets.constraints @@ -242,6 +248,11 @@ def test_reports_and_simulations_snapshot_selected_datasets() -> None: assert dataset_foreign_key.ondelete in {"RESTRICT", "SET NULL"} +def test_stage9_adds_no_dataset_version_relationship_to_run_tables() -> None: + for table_name in ("simulations", "reports", "report_runs"): + assert "dataset_version_id" not in V2_METADATA.tables[table_name].c + + def test_run_outputs_reference_report_runs_not_base_reports() -> None: for table_name in RUN_OUTPUT_TABLES: table = V2_METADATA.tables[table_name] @@ -295,9 +306,58 @@ def test_named_checks_and_required_indexes_cover_core_invariants() -> None: "ix_users_email", "ix_simulations_status_created_at", "ix_report_runs_current_output", + "uq_parameter_values_canonical_parameter_start_date", }.issubset(index_names) +def test_stage9_uses_policyengine_version_as_its_only_catalog_release_identity() -> ( + None +): + model_version = V2_METADATA.tables[TaxBenefitModelVersion.__tablename__] + dataset_version = V2_METADATA.tables[DatasetVersion.__tablename__] + forbidden_columns = { + "policyengine_version", + "core_version", + "country_package_version", + "dataset_release", + "dataset_digest", + "catalog_fingerprint", + } + + assert set(model_version.c) >= { + model_version.c.model_id, + model_version.c.version, + model_version.c.current_law_id, + model_version.c.metadata_time_periods, + } + assert forbidden_columns.isdisjoint(model_version.c.keys()) + assert forbidden_columns.isdisjoint(dataset_version.c.keys()) + + unique_column_sets = { + tuple(column.name for column in constraint.columns) + for constraint in model_version.constraints + if isinstance(constraint, sa.UniqueConstraint) + } + assert ("model_id", "version") in unique_column_sets + + +def test_canonical_parameter_values_have_a_postgres_partial_unique_index() -> None: + parameter_values = V2_METADATA.tables["parameter_values"] + index = next( + index + for index in parameter_values.indexes + if index.name == "uq_parameter_values_canonical_parameter_start_date" + ) + + assert index.unique + assert tuple(column.name for column in index.columns) == ( + "parameter_id", + "start_date", + ) + predicate = str(index.dialect_options["postgresql"]["where"]) + assert predicate == "policy_id IS NULL AND dynamic_id IS NULL" + + def test_complete_metadata_compiles_for_postgres_without_mutation() -> None: dialect = postgresql.dialect() diff --git a/tests/unit/v2/test_report_runs.py b/tests/unit/v2/test_report_runs.py index e8deda9c0..a203e3c15 100644 --- a/tests/unit/v2/test_report_runs.py +++ b/tests/unit/v2/test_report_runs.py @@ -244,12 +244,17 @@ def test_outputs_from_repeated_runs_are_preserved(engine) -> None: with Session(engine) as session: report = _create_report(session) model = report.tax_benefit_model - version = TaxBenefitModelVersion(model=model, version="1.2.3") + version = TaxBenefitModelVersion( + model=model, + version="1.2.3", + current_law_id=1, + metadata_time_periods=[2026], + ) dataset = Dataset( name="dataset", storage_path="datasets/test.h5", year=2026, - tax_benefit_model=model, + tax_benefit_model_version=version, ) simulation = Simulation( simulation_type=SimulationType.ECONOMY, diff --git a/tests/unit/v2/test_settings.py b/tests/unit/v2/test_settings.py index ad9a709e7..906a71c0c 100644 --- a/tests/unit/v2/test_settings.py +++ b/tests/unit/v2/test_settings.py @@ -3,11 +3,14 @@ import pytest from policyengine_api.data.v2.settings import ( + V2_DATA_WRITE_DATABASE_URL, V2_MIGRATION_DATABASE_URL, V2_RUNTIME_DATABASE_URL, + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE, V2_SUPABASE_ENVIRONMENT, V2_SUPABASE_PROJECT_REF, V2ConfigurationError, + load_v2_data_write_database_settings, load_v2_migration_database_settings, load_v2_runtime_database_settings, ) @@ -26,21 +29,31 @@ "postgresql+psycopg://migrator:test-migration-password@db.example.com:5432/" "postgres?sslmode=verify-full" ) +DATA_WRITE_URL = ( + "postgresql+psycopg://data-writer:test-data-write-password@db.example.com:5432/" + "postgres?sslmode=verify-ca" +) +RUNTIME_SECRET_RESOURCE = ( + "projects/test-project/secrets/v2-runtime-database-url/versions/latest" +) -def test_runtime_and_migration_urls_are_explicit_and_separate() -> None: +def test_runtime_migration_and_data_write_urls_are_explicit_and_separate() -> None: environment = { **TARGET_ENVIRONMENT, V2_RUNTIME_DATABASE_URL: RUNTIME_URL, V2_MIGRATION_DATABASE_URL: MIGRATION_URL, + V2_DATA_WRITE_DATABASE_URL: DATA_WRITE_URL, } runtime = load_v2_runtime_database_settings(environment) migration = load_v2_migration_database_settings(environment) + data_write = load_v2_data_write_database_settings(environment) assert runtime.connection.url.username == "runtime" assert migration.connection.url.username == "migrator" - assert runtime.target == migration.target + assert data_write.connection.url.username == "data-writer" + assert runtime.target == migration.target == data_write.target def test_postgres_password_is_hidden_from_string_and_repr() -> None: @@ -54,6 +67,79 @@ def test_postgres_password_is_hidden_from_string_and_repr() -> None: assert "***" in rendered +def test_runtime_url_can_be_resolved_lazily_from_secret_manager() -> None: + loaded_resources = [] + + settings = load_v2_runtime_database_settings( + { + **TARGET_ENVIRONMENT, + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: RUNTIME_SECRET_RESOURCE, + }, + secret_loader=lambda resource: loaded_resources.append(resource) or RUNTIME_URL, + ) + + assert loaded_resources == [RUNTIME_SECRET_RESOURCE] + assert settings.connection.url.username == "runtime" + + +def test_direct_runtime_url_does_not_resolve_secret_resource() -> None: + settings = load_v2_runtime_database_settings( + {**TARGET_ENVIRONMENT, V2_RUNTIME_DATABASE_URL: RUNTIME_URL}, + secret_loader=lambda resource: pytest.fail( + f"unexpected secret resolution for {resource}" + ), + ) + + assert settings.connection.url.username == "runtime" + + +@pytest.mark.parametrize( + ("environment", "message"), + [ + ( + { + **TARGET_ENVIRONMENT, + V2_RUNTIME_DATABASE_URL: RUNTIME_URL, + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: RUNTIME_SECRET_RESOURCE, + }, + "set exactly one", + ), + ( + { + **TARGET_ENVIRONMENT, + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: "not-a-resource", + }, + "is invalid", + ), + ], +) +def test_runtime_rejects_ambiguous_or_invalid_secret_sources( + environment: dict[str, str], + message: str, +) -> None: + with pytest.raises(V2ConfigurationError, match=message): + load_v2_runtime_database_settings( + environment, + secret_loader=lambda resource: RUNTIME_URL, + ) + + +def test_runtime_secret_resolution_error_hides_internal_details() -> None: + def fail_to_load(resource: str) -> str: + raise RuntimeError(f"private failure for {resource}") + + with pytest.raises(V2ConfigurationError) as raised: + load_v2_runtime_database_settings( + { + **TARGET_ENVIRONMENT, + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: RUNTIME_SECRET_RESOURCE, + }, + secret_loader=fail_to_load, + ) + + assert "private failure" not in str(raised.value) + + @pytest.mark.parametrize( "url", [ @@ -83,6 +169,8 @@ def test_v1_and_debug_settings_never_supply_missing_v2_configuration() -> None: load_v2_runtime_database_settings(environment) with pytest.raises(V2ConfigurationError, match=V2_MIGRATION_DATABASE_URL): load_v2_migration_database_settings(environment) + with pytest.raises(V2ConfigurationError, match=V2_DATA_WRITE_DATABASE_URL): + load_v2_data_write_database_settings(environment) def test_configuration_errors_do_not_echo_secret_values() -> None: diff --git a/uv.lock b/uv.lock index c74718820..de2dde78d 100644 --- a/uv.lock +++ b/uv.lock @@ -2564,6 +2564,7 @@ dependencies = [ { name = "markupsafe" }, { name = "microdf-python" }, { name = "openai" }, + { name = "packaging" }, { name = "policyengine", extra = ["models"] }, { name = "policyengine-canada" }, { name = "policyengine-il" }, @@ -2613,6 +2614,7 @@ requires-dist = [ { name = "markupsafe", specifier = ">=3,<4" }, { name = "microdf-python", specifier = ">=1.0.0" }, { name = "openai" }, + { name = "packaging", specifier = ">=24,<27" }, { name = "policyengine", extras = ["models"], specifier = "==5.2.0" }, { name = "policyengine-canada", specifier = "==0.96.3" }, { name = "policyengine-il", specifier = "==0.1.0" }, From e71a384ead9eb170219a7f84389e99a0d7af4884 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sat, 29 Aug 2026 18:17:27 +0400 Subject: [PATCH 02/27] Fix repeated parameter value normalization --- .../data/v2/catalog/extraction.py | 21 +++++-- tests/unit/v2/test_catalog_extraction.py | 60 ++++++++++++++++++- 2 files changed, 72 insertions(+), 9 deletions(-) diff --git a/policyengine_api/data/v2/catalog/extraction.py b/policyengine_api/data/v2/catalog/extraction.py index fa82dcb0a..953054f95 100644 --- a/policyengine_api/data/v2/catalog/extraction.py +++ b/policyengine_api/data/v2/catalog/extraction.py @@ -308,7 +308,7 @@ def _normalize_parameter( ) parameter_id = _identifier("parameter", model_version_id, name) source_values: list[tuple[datetime, datetime | None, Any]] = [] - seen_starts: set[datetime] = set() + values_by_start: dict[datetime, tuple[datetime | None, Any]] = {} previous_start: datetime | None = None for source_value in getattr(source, "parameter_values", ()): start_date = _aware_datetime( @@ -325,15 +325,24 @@ def _normalize_parameter( raise CatalogExtractionError( f"parameter {name!r} has an unsupported JSON value: {error}" ) from error - if start_date in seen_starts: - raise CatalogExtractionError( - f"parameter {name!r} has a duplicate value start date" - ) + if start_date in values_by_start: + previous_end, previous_value = values_by_start[start_date] + if previous_value != value_json: + raise CatalogExtractionError( + f"parameter {name!r} has conflicting values at effective " + f"date {start_date.isoformat()}" + ) + if previous_end != end_date: + raise CatalogExtractionError( + f"parameter {name!r} exposes inconsistent intervals at effective " + f"date {start_date.isoformat()}" + ) + continue if previous_start is not None and start_date <= previous_start: raise CatalogExtractionError( f"parameter {name!r} values are not ordered oldest to newest" ) - seen_starts.add(start_date) + values_by_start[start_date] = (end_date, value_json) previous_start = start_date source_values.append((start_date, end_date, value_json)) diff --git a/tests/unit/v2/test_catalog_extraction.py b/tests/unit/v2/test_catalog_extraction.py index baa59d528..90a34239e 100644 --- a/tests/unit/v2/test_catalog_extraction.py +++ b/tests/unit/v2/test_catalog_extraction.py @@ -268,12 +268,66 @@ def test_rejects_parameter_values_that_are_not_oldest_to_newest() -> None: ) -def test_rejects_duplicate_parameter_value_start_dates() -> None: +def test_collapses_equal_parameter_values_at_the_same_effective_date() -> None: models = source_models() parameter = models["us"].parameters_by_name["gov.example.rate"] - parameter.parameter_values.append(parameter.parameter_values[1]) + repeated = parameter.parameter_values[1] + parameter.parameter_values.append( + SimpleNamespace( + value=repeated.value, + start_date=repeated.start_date, + end_date=repeated.end_date, + ) + ) + + catalog = extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + values = catalog.country("us").parameters[0].values + assert [(value.start_date, value.value_json) for value in values] == [ + (datetime(2025, 1, 1, tzinfo=timezone.utc), 0.1), + (datetime(2026, 1, 1, tzinfo=timezone.utc), 0.2), + ] + + +def test_rejects_conflicting_parameter_values_at_the_same_effective_date() -> None: + models = source_models() + parameter = models["us"].parameters_by_name["gov.example.rate"] + repeated = parameter.parameter_values[1] + parameter.parameter_values.append( + SimpleNamespace( + value=0.3, + start_date=repeated.start_date, + end_date=repeated.end_date, + ) + ) + + with pytest.raises(CatalogExtractionError, match="conflicting values"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + +def test_rejects_repeated_parameter_values_with_inconsistent_intervals() -> None: + models = source_models() + parameter = models["us"].parameters_by_name["gov.example.rate"] + repeated = parameter.parameter_values[1] + parameter.parameter_values.append( + SimpleNamespace( + value=repeated.value, + start_date=repeated.start_date, + end_date=repeated.start_date, + ) + ) - with pytest.raises(CatalogExtractionError, match="duplicate value start date"): + with pytest.raises(CatalogExtractionError, match="inconsistent intervals"): extract_catalog( bundle=bundle(), policyengine_version=POLICYENGINE_VERSION, From b9068acc432bb28c52fe00d82bce413ca017b64f Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sat, 29 Aug 2026 21:34:26 +0400 Subject: [PATCH 03/27] Update PolicyEngine.py to 5.2.0 --- docs/migration/stage-9-v2-metadata.md | 6 +++--- tests/integration/test_v2_catalog_installed.py | 3 ++- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/docs/migration/stage-9-v2-metadata.md b/docs/migration/stage-9-v2-metadata.md index 3176b270c..1f3b4b825 100644 --- a/docs/migration/stage-9-v2-metadata.md +++ b/docs/migration/stage-9-v2-metadata.md @@ -142,11 +142,11 @@ regions, model descriptions, current-law IDs, and time-period options. ## Production-scale qualification record The Stage 9 implementation qualification used the locked PolicyEngine.py -5.0.4 distribution. Its manifest selected PolicyEngine Core 3.30.1, +5.2.0 distribution. Its manifest selected PolicyEngine Core 3.30.1, PolicyEngine US 1.764.6, and PolicyEngine UK 2.90.2. Extraction produced 2 -models, 2 model versions, 6,649 variables, 27,826 named parameter nodes, 99,006 +models, 2 model versions, 6,649 variables, 27,813 named parameter nodes, 99,006 parameters, 1,172,130 parameter values, 2 logical input datasets, and 826 -regions. Publication took 33.916 seconds and added 7,979,842 bytes of measured +regions. Publication took 24.120 seconds and added 7,968,103 bytes of measured peak publisher memory. The US fallback summary reported 436 congressional districts, 333 places, and 51 states. diff --git a/tests/integration/test_v2_catalog_installed.py b/tests/integration/test_v2_catalog_installed.py index 05d79c47a..b09ce746a 100644 --- a/tests/integration/test_v2_catalog_installed.py +++ b/tests/integration/test_v2_catalog_installed.py @@ -40,7 +40,7 @@ def test_installed_policyengine_catalog_is_complete_and_bounded() -> None: "models": 2, "model_versions": 2, "variables": 6_649, - "parameter_nodes": 27_826, + "parameter_nodes": 27_813, "parameters": 99_006, "parameter_values": 1_172_130, "datasets": 2, @@ -54,6 +54,7 @@ def test_installed_policyengine_catalog_is_complete_and_bounded() -> None: assert country.model_version.version not in dict(expected_dependencies).values() assert country.variables assert country.parameter_nodes + assert all("__pycache__" not in node.name for node in country.parameter_nodes) assert country.parameters assert all( not dataset.is_output_dataset and dataset.storage_path is None From 75eafd651566f2200eacf10679947e9f37e0eff8 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sat, 29 Aug 2026 22:43:23 +0400 Subject: [PATCH 04/27] Add Stage 9 changelog fragment --- changelog.d/stage-9-v2-metadata.added.md | 1 + 1 file changed, 1 insertion(+) create mode 100644 changelog.d/stage-9-v2-metadata.added.md diff --git a/changelog.d/stage-9-v2-metadata.added.md b/changelog.d/stage-9-v2-metadata.added.md new file mode 100644 index 000000000..b00c3653f --- /dev/null +++ b/changelog.d/stage-9-v2-metadata.added.md @@ -0,0 +1 @@ +Add the versioned API v2 metadata catalog and Cloud Run preview endpoints. From c0865a4e0d37a90fd1579aed7a2f7c981971cb09 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 13:49:00 +0400 Subject: [PATCH 05/27] Fix v2 metadata catalog validation --- policyengine_api/constants.py | 17 ++++----- .../data/v2/catalog/extraction.py | 10 ++++++ policyengine_api/data/v2/catalog/query.py | 9 ++++- policyengine_api/dataset_display.py | 20 +++++++++++ tests/integration/test_v2_metadata_routes.py | 35 ++++++++++++++++++ tests/unit/v2/test_catalog_extraction.py | 17 +++++++++ tests/unit/v2/test_metadata_query.py | 36 +++++++++++++++++++ 7 files changed, 133 insertions(+), 11 deletions(-) create mode 100644 policyengine_api/dataset_display.py diff --git a/policyengine_api/constants.py b/policyengine_api/constants.py index 49ed4d96d..683fb4aa9 100644 --- a/policyengine_api/constants.py +++ b/policyengine_api/constants.py @@ -4,6 +4,11 @@ from importlib.metadata import distribution, distributions from pathlib import Path +from policyengine_api.dataset_display import ( + DEFAULT_DATASET_DISPLAY_LABEL, + get_dataset_display_label, +) + REPO = Path(__file__).parents[1] GET = "GET" POST = "POST" @@ -23,11 +28,7 @@ "uk": "policyengine-uk", "us": "policyengine-us", } -BUNDLE_DATASET_DISPLAY_LABELS = { - "populace_": "Microcosm", - "enhanced_frs_": "Enhanced FRS", -} -DEFAULT_BUNDLE_DATASET_LABEL = "Certified dataset" +DEFAULT_BUNDLE_DATASET_LABEL = DEFAULT_DATASET_DISPLAY_LABEL def _normalize_distribution_name(name: str | None) -> str: @@ -110,11 +111,7 @@ def get_bundle_default_dataset(country_id: str) -> str | None: def _bundle_dataset_display_label(default_dataset: object) -> str: - dataset_name = str(default_dataset or "") - for prefix, label in BUNDLE_DATASET_DISPLAY_LABELS.items(): - if dataset_name.startswith(prefix): - return label - return DEFAULT_BUNDLE_DATASET_LABEL + return get_dataset_display_label(default_dataset) def get_bundle_default_dataset_option(country_id: str) -> dict: diff --git a/policyengine_api/data/v2/catalog/extraction.py b/policyengine_api/data/v2/catalog/extraction.py index 953054f95..7d2302920 100644 --- a/policyengine_api/data/v2/catalog/extraction.py +++ b/policyengine_api/data/v2/catalog/extraction.py @@ -421,15 +421,25 @@ def _normalize_country( expected_country_package_version: str, ) -> CountryCatalog: source_model = getattr(source, "model", None) + expected_model_name = f"policyengine-{country_id}" model_name = _required_text( getattr(source_model, "id", None), field_name=f"{country_id} model name", maximum=32, ) + if model_name != expected_model_name: + raise CatalogExtractionError( + f"{country_id} public model identity does not match {expected_model_name}" + ) model_id = _identifier("model", model_name) model_version_id = _identifier("model-version", model_id, policyengine_version) model_package = getattr(source, "model_package", None) + if getattr(model_package, "name", None) != expected_model_name: + raise CatalogExtractionError( + f"{country_id} public model package identity does not match " + f"{expected_model_name}" + ) if getattr(model_package, "version", None) != expected_country_package_version: raise CatalogExtractionError( f"{country_id} public model does not match the PolicyEngine.py manifest" diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index bce8a4084..d19ddd05b 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -8,6 +8,7 @@ from sqlalchemy.exc import SQLAlchemyError from sqlmodel import Session, select +from policyengine_api.dataset_display import get_dataset_display_label from policyengine_api.data.v2.catalog.schemas import ( MetadataDataset, MetadataDatasetOption, @@ -195,6 +196,12 @@ def _read_metadata( raise MetadataCatalogUnavailableError( f"the {country_id} v2 metadata catalog is incomplete" ) + if {value.parameter_id for value in parameter_values} != { + parameter.id for parameter in parameters + }: + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 parameter values are incomplete" + ) dataset_ids = {region.default_dataset_id for region in regions} datasets = self._session.exec( @@ -335,7 +342,7 @@ def _read_metadata( datasets=[ MetadataDatasetOption( name=national_dataset.name, - label=national_dataset.description or national_dataset.name, + label=get_dataset_display_label(national_dataset.name), ) ], ), diff --git a/policyengine_api/dataset_display.py b/policyengine_api/dataset_display.py new file mode 100644 index 000000000..8859db018 --- /dev/null +++ b/policyengine_api/dataset_display.py @@ -0,0 +1,20 @@ +"""User-facing labels for certified PolicyEngine dataset families.""" + +from __future__ import annotations + + +DATASET_FAMILY_DISPLAY_LABELS = ( + ("populace_", "Microcosm"), + ("enhanced_frs_", "Enhanced FRS"), +) +DEFAULT_DATASET_DISPLAY_LABEL = "Certified dataset" + + +def get_dataset_display_label(dataset_name: object) -> str: + """Return the public label for an internal logical dataset name.""" + + normalized_name = str(dataset_name or "") + for prefix, label in DATASET_FAMILY_DISPLAY_LABELS: + if normalized_name.startswith(prefix): + return label + return DEFAULT_DATASET_DISPLAY_LABEL diff --git a/tests/integration/test_v2_metadata_routes.py b/tests/integration/test_v2_metadata_routes.py index 40a38b4ba..18d175f17 100644 --- a/tests/integration/test_v2_metadata_routes.py +++ b/tests/integration/test_v2_metadata_routes.py @@ -121,6 +121,12 @@ def test_postgres_preview_returns_complete_us_and_uk_catalogs_without_writes( assert us.status_code == uk.status_code == 200 assert us.json()["result"]["current_law_id"] == 2 assert uk.json()["result"]["current_law_id"] == 1 + assert us.json()["result"]["economy_options"]["datasets"][0]["label"] == ( + "Microcosm" + ) + assert uk.json()["result"]["economy_options"]["datasets"][0]["label"] == ( + "Enhanced FRS" + ) for country_id, response in (("us", us), ("uk", uk)): result = response.json()["result"] assert result["model"]["name"] == f"policyengine-{country_id}" @@ -244,3 +250,32 @@ def test_postgres_preview_distinguishes_invalid_and_absent_versions( assert absent.status_code == 404 assert absent.json()["status"] == "error" assert absent.json()["message"] + + +def test_postgres_preview_returns_typed_error_for_incomplete_parameter_values( + published_engine: Engine, +) -> None: + with published_engine.begin() as connection: + connection.execute( + text( + """ + DELETE FROM parameter_values + WHERE parameter_id IN ( + SELECT parameter.id + FROM parameters AS parameter + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = parameter.tax_benefit_model_version_id + JOIN tax_benefit_models AS model + ON model.id = model_version.model_id + WHERE model.name = 'policyengine-us' + ) + """ + ) + ) + + response = _client(published_engine).get("/v2/us/metadata") + + assert response.status_code == 503 + assert response.json()["status"] == "error" + assert response.json()["message"] + assert "result" not in response.json() diff --git a/tests/unit/v2/test_catalog_extraction.py b/tests/unit/v2/test_catalog_extraction.py index 90a34239e..4e07ede36 100644 --- a/tests/unit/v2/test_catalog_extraction.py +++ b/tests/unit/v2/test_catalog_extraction.py @@ -190,6 +190,23 @@ def test_rejects_incomplete_public_model_catalogs() -> None: ) +@pytest.mark.parametrize("identity_source", ["model", "model_package"]) +def test_rejects_inconsistent_country_model_identity(identity_source: str) -> None: + models = source_models() + if identity_source == "model": + models["us"].model.id = "unexpected-us-model" + else: + models["us"].model_package.name = "unexpected-us-package" + + with pytest.raises(CatalogExtractionError, match="identity does not match"): + extract_catalog( + bundle=bundle(), + policyengine_version=POLICYENGINE_VERSION, + models=models, + installed_version=installed_version, + ) + + def test_ignores_only_the_unnamed_structural_parameter_root() -> None: models = source_models() models["uk"].parameter_nodes_by_name[""] = SimpleNamespace( diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index dc0f4b194..2fa49bddc 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -297,6 +297,7 @@ def test_query_serializes_complete_typed_metadata_without_writes( assert [option.name for option in result.economy_options.datasets] == [ "populace_us_2024" ] + assert [option.label for option in result.economy_options.datasets] == ["Microcosm"] after = { model_class.__tablename__: len(catalog_session.exec(select(model_class)).all()) for model_class in model_classes @@ -419,6 +420,41 @@ def test_query_rejects_unsupported_and_incomplete_catalogs( service.get_metadata("uk") +def test_query_rejects_incomplete_parameter_values( + catalog_session: Session, +) -> None: + us_model_version = catalog_session.exec( + select(TaxBenefitModelVersion) + .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) + .where( + TaxBenefitModel.name == "policyengine-us", + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ).one() + us_parameter_ids = set( + catalog_session.exec( + select(Parameter.id).where( + Parameter.tax_benefit_model_version_id == us_model_version.id + ) + ).all() + ) + for value in catalog_session.exec( + select(ParameterValue).where(ParameterValue.parameter_id.in_(us_parameter_ids)) + ).all(): + catalog_session.delete(value) + catalog_session.commit() + + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + with pytest.raises( + MetadataCatalogUnavailableError, + match="parameter values are incomplete", + ): + service.get_metadata("us") + + def test_response_outcomes_are_discriminated_and_strict( catalog_session: Session, ) -> None: From 61690cf0e592195e09b11ec7acf9acc03f63786a Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:06:54 +0400 Subject: [PATCH 06/27] Include v2 integration tests in coverage --- .github/workflows/alembic-v2-check.yml | 19 ++++- .github/workflows/pr.yml | 2 + .github/workflows/push.yml | 2 + tests/unit/test_alembic_workflows.py | 7 ++ tests/unit/v2/test_metadata_query.py | 101 ++++++++++++++++++++++++- tests/unit/v2/test_metadata_routes.py | 70 ++++++++++++++++- 6 files changed, 195 insertions(+), 6 deletions(-) diff --git a/.github/workflows/alembic-v2-check.yml b/.github/workflows/alembic-v2-check.yml index 19da03199..4feb696b7 100644 --- a/.github/workflows/alembic-v2-check.yml +++ b/.github/workflows/alembic-v2-check.yml @@ -2,6 +2,9 @@ name: Alembic v2 and runtime-cache checks on: workflow_call: + secrets: + CODECOV_TOKEN: + required: false workflow_dispatch: jobs: @@ -53,14 +56,22 @@ jobs: - name: Require no ungenerated v2 schema or data operations run: uv run alembic -c alembic-v2.ini check - name: Verify the installed PolicyEngine.py catalog interface - run: uv run pytest -q tests/integration/test_v2_catalog_installed.py + run: uv run coverage run --branch -m pytest -q tests/integration/test_v2_catalog_installed.py env: RUN_V2_CATALOG_COMPATIBILITY: "1" - name: Test v2 metadata publication and preview routes - run: uv run pytest -q tests/integration/test_v2_catalog_publication.py tests/integration/test_v2_metadata_routes.py + run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication.py tests/integration/test_v2_metadata_routes.py - name: Qualify production-scale v2 metadata publication - run: uv run pytest -q tests/integration/test_v2_catalog_publication_qualification.py + run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication_qualification.py env: RUN_V2_CATALOG_PUBLICATION_QUALIFICATION: "1" - name: Test real Redis cross-instance semantics - run: uv run pytest -q tests/integration/test_runtime_cache_redis.py + run: uv run coverage run -a --branch -m pytest -q tests/integration/test_runtime_cache_redis.py + - name: Export v2 integration coverage + run: uv run coverage xml -i -o coverage-v2.xml + - name: Upload v2 integration coverage to Codecov + uses: codecov/codecov-action@v5 + with: + token: ${{ secrets.CODECOV_TOKEN }} + slug: PolicyEngine/policyengine-api + files: coverage-v2.xml diff --git a/.github/workflows/pr.yml b/.github/workflows/pr.yml index bd91d0b4a..ca1f5858a 100644 --- a/.github/workflows/pr.yml +++ b/.github/workflows/pr.yml @@ -57,6 +57,8 @@ jobs: alembic-v2-check: name: Alembic v2 and Redis qualification uses: ./.github/workflows/alembic-v2-check.yml + secrets: + CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} check-changelog: name: Check changelog fragment runs-on: ubuntu-latest diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index 6a41b3f96..9278e932d 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -37,6 +37,8 @@ jobs: name: Alembic v2 and Redis qualification if: github.repository == 'PolicyEngine/policyengine-api' uses: ./.github/workflows/alembic-v2-check.yml + secrets: + CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} ensure-staging-model-version-aligns-with-sim-api: name: Ensure staging model version aligns with simulation API diff --git a/tests/unit/test_alembic_workflows.py b/tests/unit/test_alembic_workflows.py index 8b84f4a1a..c4a7dd5f6 100644 --- a/tests/unit/test_alembic_workflows.py +++ b/tests/unit/test_alembic_workflows.py @@ -67,6 +67,7 @@ def test_pr_always_runs_reusable_alembic_check(): assert "dorny/paths-filter" not in workflow assert "alembic-v2-check:" in workflow assert "uses: ./.github/workflows/alembic-v2-check.yml" in v2_job + assert "CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}" in v2_job assert "needs:" not in v2_job assert "if:" not in v2_job @@ -79,6 +80,7 @@ def test_push_always_runs_lint_and_alembic_qualification_before_versioning(): assert "uses: ./.github/workflows/alembic-v1-check.yml" in workflow assert "alembic-v2-check:" in workflow assert "uses: ./.github/workflows/alembic-v2-check.yml" in workflow + assert "CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}" in workflow assert "needs: [lint, alembic-v1-check, alembic-v2-check]" in workflow assert "github.repository == 'PolicyEngine/policyengine-uk'" not in workflow @@ -139,6 +141,11 @@ def test_reusable_v2_check_uses_disposable_postgres_and_real_redis(): assert "RUN_V2_CATALOG_PUBLICATION_QUALIFICATION" in workflow assert "test_runtime_cache_redis.py" in workflow assert "uv sync --frozen" in workflow + assert workflow.count("coverage run --branch") == 1 + assert workflow.count("coverage run -a --branch") == 3 + assert "coverage xml -i -o coverage-v2.xml" in workflow + assert "codecov/codecov-action@v5" in workflow + assert "files: coverage-v2.xml" in workflow def test_release_migration_fails_closed_and_gates_both_staging_deploys(): diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index 2fa49bddc..eba0ec88f 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -8,6 +8,7 @@ from pydantic import TypeAdapter, ValidationError import pytest +from sqlalchemy import delete from sqlalchemy.pool import StaticPool from sqlmodel import Session, create_engine, select @@ -380,7 +381,13 @@ def test_query_rejects_invalid_or_absent_selected_versions( running_policyengine_version=POLICYENGINE_VERSION, ) - for invalid in ("", f" {POLICYENGINE_VERSION}", "not a version", "0.0.0"): + for invalid in ( + "", + f" {POLICYENGINE_VERSION}", + "not a version", + "0.0.0", + "1" * 129, + ): with pytest.raises(InvalidPolicyEngineVersionError): service.get_metadata("us", invalid) with pytest.raises(MetadataCatalogVersionNotFoundError): @@ -392,6 +399,98 @@ def test_query_rejects_invalid_or_absent_selected_versions( ).get_metadata("us") +def test_query_distinguishes_an_uninitialized_catalog_from_an_absent_version() -> None: + engine = create_engine("sqlite://") + V2_METADATA.create_all(engine) + try: + with Session(engine) as session: + service = V2MetadataQueryService( + session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + with pytest.raises( + MetadataCatalogUnavailableError, + match="not initialized", + ): + service.get_metadata("us") + with pytest.raises( + MetadataCatalogVersionNotFoundError, + match=( + f"PolicyEngine.py {POLICYENGINE_VERSION} is not published for us" + ), + ): + service.get_metadata("us", POLICYENGINE_VERSION) + finally: + engine.dispose() + + +def test_query_rejects_a_region_whose_dataset_is_absent( + catalog_session: Session, +) -> None: + national_region = catalog_session.exec( + select(Region).where(Region.code == "us") + ).one() + catalog_session.exec( + delete(Dataset).where(Dataset.id == national_region.default_dataset_id) + ) + catalog_session.commit() + + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + with pytest.raises( + MetadataCatalogUnavailableError, + match="region datasets are incomplete", + ): + service.get_metadata("us") + + +def test_query_requires_a_national_region(catalog_session: Session) -> None: + national_region = catalog_session.exec( + select(Region).where(Region.code == "us") + ).one() + catalog_session.delete(national_region) + catalog_session.commit() + + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + with pytest.raises( + MetadataCatalogUnavailableError, + match="national v2 region is absent", + ): + service.get_metadata("us") + + +def test_query_requires_nonempty_integer_time_periods( + catalog_session: Session, +) -> None: + us_version = catalog_session.exec( + select(TaxBenefitModelVersion) + .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) + .where( + TaxBenefitModel.name == "policyengine-us", + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ).one() + us_version.metadata_time_periods = [] + catalog_session.add(us_version) + catalog_session.commit() + + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + with pytest.raises( + MetadataCatalogUnavailableError, + match="model-version options are incomplete", + ): + service.get_metadata("us") + + def test_query_rejects_unsupported_and_incomplete_catalogs( catalog_session: Session, ) -> None: diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py index e5cac18e5..4cbc65b5e 100644 --- a/tests/unit/v2/test_metadata_routes.py +++ b/tests/unit/v2/test_metadata_routes.py @@ -30,6 +30,7 @@ MetadataVariable, ) from policyengine_api.data.v2.settings import V2ConfigurationError +from policyengine_api.fastapi_routes import dependencies as route_dependencies from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies from policyengine_api.migration_flags import ( RouteImplementation, @@ -38,9 +39,15 @@ class Reader: - def __init__(self, result: MetadataResult, error: Exception | None = None): + def __init__( + self, + result: MetadataResult, + error: Exception | None = None, + close_error: Exception | None = None, + ): self.result = result self.error = error + self.close_error = close_error self.calls = [] self.closed = False @@ -56,6 +63,8 @@ def get_metadata( def close(self) -> None: self.closed = True + if self.close_error is not None: + raise self.close_error def _result(country_id: str = "us") -> MetadataResult: @@ -311,6 +320,65 @@ def test_preview_passes_explicit_version_to_reader() -> None: assert reader.calls == [("us", "4.19.0")] +def test_preview_uses_default_reader_factory_when_none_is_injected( + monkeypatch: pytest.MonkeyPatch, +) -> None: + reader = Reader(_result()) + monkeypatch.setattr( + route_dependencies, + "_default_v2_metadata_reader_factory", + lambda: reader, + ) + + response = _client(None).get("/v2/us/metadata") + + assert response.status_code == 200 + assert reader.calls == [("us", None)] + assert reader.closed + + +def test_preview_ignores_reader_close_failure() -> None: + reader = Reader(_result(), close_error=RuntimeError("close failed")) + + response = _client(lambda: reader).get("/v2/us/metadata") + + assert response.status_code == 200 + assert response.json()["status"] == "ok" + assert reader.closed + + +def test_default_reader_uses_the_installed_policyengine_version( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from policyengine_api.data.v2 import database + from policyengine_api.data.v2.catalog import query + + session = object() + captured = {} + reader = object() + + def query_service(candidate_session, *, running_policyengine_version): + captured["session"] = candidate_session + captured["version"] = running_policyengine_version + return reader + + monkeypatch.setattr(database, "get_v2_session_factory", lambda: lambda: session) + monkeypatch.setattr(query, "V2MetadataQueryService", query_service) + monkeypatch.setattr( + route_dependencies.importlib_metadata, + "version", + lambda package: "5.2.0" if package == "policyengine" else "unexpected", + ) + route_dependencies._running_policyengine_version.cache_clear() + try: + result = route_dependencies._default_v2_metadata_reader_factory() + finally: + route_dependencies._running_policyengine_version.cache_clear() + + assert result is reader + assert captured == {"session": session, "version": "5.2.0"} + + def test_preview_reads_repeat_without_changing_v1_routing() -> None: reader = Reader(_result()) client = _client(lambda: reader) From 504cf0cfc5b43b88c5ecd49d56c67b1a41159ddc Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:18:14 +0400 Subject: [PATCH 07/27] Cover v2 configuration failure paths --- tests/unit/test_migration_flags.py | 18 ++++++++++++++++++ tests/unit/v2/test_settings.py | 11 +++++++++++ 2 files changed, 29 insertions(+) diff --git a/tests/unit/test_migration_flags.py b/tests/unit/test_migration_flags.py index 55dd39f78..b808a0941 100644 --- a/tests/unit/test_migration_flags.py +++ b/tests/unit/test_migration_flags.py @@ -108,6 +108,24 @@ def test_invalid_migration_flag_raises(monkeypatch): get_migration_context("policy") +@pytest.mark.parametrize( + "explicit_sources", + [ + {"db_write_source": "invalid", "db_read_source": None}, + {"db_write_source": None, "db_read_source": "invalid"}, + ], +) +def test_explicit_migration_context_rejects_invalid_database_sources( + explicit_sources, +): + with pytest.raises(ValueError, match="invalid explicit database"): + get_migration_context( + "metadata", + use_configured_db_sources=False, + **explicit_sources, + ) + + @pytest.mark.parametrize( ("path", "expected_group"), [ diff --git a/tests/unit/v2/test_settings.py b/tests/unit/v2/test_settings.py index 906a71c0c..e06e45b16 100644 --- a/tests/unit/v2/test_settings.py +++ b/tests/unit/v2/test_settings.py @@ -140,6 +140,17 @@ def fail_to_load(resource: str) -> str: assert "private failure" not in str(raised.value) +def test_runtime_rejects_an_empty_resolved_secret() -> None: + with pytest.raises(V2ConfigurationError, match="is empty"): + load_v2_runtime_database_settings( + { + **TARGET_ENVIRONMENT, + V2_RUNTIME_DATABASE_URL_SECRET_RESOURCE: RUNTIME_SECRET_RESOURCE, + }, + secret_loader=lambda resource: " ", + ) + + @pytest.mark.parametrize( "url", [ From ef011d423c937d9d7e0d86100d07ea56c7b4026f Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 14:45:10 +0400 Subject: [PATCH 08/27] Separate v2 integration from Alembic checks --- .github/workflows/alembic-v2-check.yml | 39 +----------- .github/workflows/pr.yml | 5 +- .github/workflows/push.yml | 14 ++++- .github/workflows/v2-integration-check.yml | 73 ++++++++++++++++++++++ tests/unit/test_alembic_workflows.py | 67 ++++++++++++++++---- 5 files changed, 146 insertions(+), 52 deletions(-) create mode 100644 .github/workflows/v2-integration-check.yml diff --git a/.github/workflows/alembic-v2-check.yml b/.github/workflows/alembic-v2-check.yml index 4feb696b7..4af38a66f 100644 --- a/.github/workflows/alembic-v2-check.yml +++ b/.github/workflows/alembic-v2-check.yml @@ -1,15 +1,12 @@ -name: Alembic v2 and runtime-cache checks +name: Alembic v2 schema checks on: workflow_call: - secrets: - CODECOV_TOKEN: - required: false workflow_dispatch: jobs: - postgres-redis-lifecycle: - name: V2 Postgres and Redis lifecycle + postgres-lifecycle: + name: V2 Postgres migration lifecycle runs-on: ubuntu-latest services: postgres: @@ -25,19 +22,9 @@ jobs: --health-interval=5s --health-timeout=5s --health-retries=20 - redis: - image: redis:7.2-alpine - ports: - - 6379:6379 - options: >- - --health-cmd="redis-cli ping" - --health-interval=5s - --health-timeout=5s - --health-retries=20 env: V2_MIGRATION_DATABASE_URL: postgresql+psycopg://postgres:policyengine_v2_test@127.0.0.1:5432/policyengine_v2_alembic_test V2_ALEMBIC_DISPOSABLE_TEST: "1" - RUNTIME_CACHE_TEST_URL: redis://127.0.0.1:6379/0 steps: - name: Checkout repo uses: actions/checkout@v4 @@ -55,23 +42,3 @@ jobs: run: uv run alembic -c alembic-v2.ini current --check-heads - name: Require no ungenerated v2 schema or data operations run: uv run alembic -c alembic-v2.ini check - - name: Verify the installed PolicyEngine.py catalog interface - run: uv run coverage run --branch -m pytest -q tests/integration/test_v2_catalog_installed.py - env: - RUN_V2_CATALOG_COMPATIBILITY: "1" - - name: Test v2 metadata publication and preview routes - run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication.py tests/integration/test_v2_metadata_routes.py - - name: Qualify production-scale v2 metadata publication - run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication_qualification.py - env: - RUN_V2_CATALOG_PUBLICATION_QUALIFICATION: "1" - - name: Test real Redis cross-instance semantics - run: uv run coverage run -a --branch -m pytest -q tests/integration/test_runtime_cache_redis.py - - name: Export v2 integration coverage - run: uv run coverage xml -i -o coverage-v2.xml - - name: Upload v2 integration coverage to Codecov - uses: codecov/codecov-action@v5 - with: - token: ${{ secrets.CODECOV_TOKEN }} - slug: PolicyEngine/policyengine-api - files: coverage-v2.xml diff --git a/.github/workflows/pr.yml b/.github/workflows/pr.yml index ca1f5858a..ffc5d0e0b 100644 --- a/.github/workflows/pr.yml +++ b/.github/workflows/pr.yml @@ -55,8 +55,11 @@ jobs: name: Alembic v1 qualification uses: ./.github/workflows/alembic-v1-check.yml alembic-v2-check: - name: Alembic v2 and Redis qualification + name: Alembic v2 qualification uses: ./.github/workflows/alembic-v2-check.yml + v2-integration-check: + name: V2 integration qualification + uses: ./.github/workflows/v2-integration-check.yml secrets: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} check-changelog: diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index 9278e932d..7bd174bc9 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -34,9 +34,14 @@ jobs: uses: ./.github/workflows/alembic-v1-check.yml alembic-v2-check: - name: Alembic v2 and Redis qualification + name: Alembic v2 qualification if: github.repository == 'PolicyEngine/policyengine-api' uses: ./.github/workflows/alembic-v2-check.yml + + v2-integration-check: + name: V2 integration qualification + if: github.repository == 'PolicyEngine/policyengine-api' + uses: ./.github/workflows/v2-integration-check.yml secrets: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} @@ -60,7 +65,11 @@ jobs: versioning: name: Update versioning - needs: [lint, alembic-v1-check, alembic-v2-check] + needs: + - lint + - alembic-v1-check + - alembic-v2-check + - v2-integration-check if: | (github.repository == 'PolicyEngine/policyengine-api') && !(github.event.head_commit.message == 'Update PolicyEngine API') @@ -103,6 +112,7 @@ jobs: - lint - alembic-v1-check - alembic-v2-check + - v2-integration-check if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') diff --git a/.github/workflows/v2-integration-check.yml b/.github/workflows/v2-integration-check.yml new file mode 100644 index 000000000..9bf9accee --- /dev/null +++ b/.github/workflows/v2-integration-check.yml @@ -0,0 +1,73 @@ +name: V2 integration checks + +on: + workflow_call: + secrets: + CODECOV_TOKEN: + required: false + workflow_dispatch: + +jobs: + catalog-runtime-integration: + name: V2 PostgreSQL and Redis integration + runs-on: ubuntu-latest + services: + postgres: + image: postgres:17 + env: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: policyengine_v2_test + POSTGRES_DB: policyengine_v2_alembic_test + ports: + - 5432:5432 + options: >- + --health-cmd="pg_isready -U postgres" + --health-interval=5s + --health-timeout=5s + --health-retries=20 + redis: + image: redis:7.2-alpine + ports: + - 6379:6379 + options: >- + --health-cmd="redis-cli ping" + --health-interval=5s + --health-timeout=5s + --health-retries=20 + env: + V2_MIGRATION_DATABASE_URL: postgresql+psycopg://postgres:policyengine_v2_test@127.0.0.1:5432/policyengine_v2_alembic_test + V2_ALEMBIC_DISPOSABLE_TEST: "1" + RUNTIME_CACHE_TEST_URL: redis://127.0.0.1:6379/0 + steps: + - name: Checkout repo + uses: actions/checkout@v4 + - name: Setup Python + uses: actions/setup-python@v5 + with: + python-version: "3.12" + - name: Setup uv + uses: astral-sh/setup-uv@v6 + - name: Install locked dependencies + run: uv sync --frozen + - name: Prepare the disposable v2 schema + run: uv run alembic -c alembic-v2.ini upgrade head + - name: Verify the installed PolicyEngine.py catalog interface + run: uv run coverage run --branch -m pytest -q tests/integration/test_v2_catalog_installed.py + env: + RUN_V2_CATALOG_COMPATIBILITY: "1" + - name: Test v2 metadata publication and preview routes + run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication.py tests/integration/test_v2_metadata_routes.py + - name: Qualify production-scale v2 metadata publication + run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication_qualification.py + env: + RUN_V2_CATALOG_PUBLICATION_QUALIFICATION: "1" + - name: Test real Redis cross-instance semantics + run: uv run coverage run -a --branch -m pytest -q tests/integration/test_runtime_cache_redis.py + - name: Export v2 integration coverage + run: uv run coverage xml -i -o coverage-v2.xml + - name: Upload v2 integration coverage to Codecov + uses: codecov/codecov-action@v5 + with: + token: ${{ secrets.CODECOV_TOKEN }} + slug: PolicyEngine/policyengine-api + files: coverage-v2.xml diff --git a/tests/unit/test_alembic_workflows.py b/tests/unit/test_alembic_workflows.py index c4a7dd5f6..4d481c985 100644 --- a/tests/unit/test_alembic_workflows.py +++ b/tests/unit/test_alembic_workflows.py @@ -54,10 +54,12 @@ def test_workflows_do_not_inline_long_shell_programs(): assert _long_inline_run_blocks() == [] -def test_pr_always_runs_reusable_alembic_check(): +def test_pr_always_runs_reusable_alembic_and_v2_integration_checks(): workflow = _workflow("pr.yml") - v2_job = workflow[workflow.index(" alembic-v2-check:") :] - v2_job = v2_job[: v2_job.index("\n check-changelog:")] + alembic_job = workflow[workflow.index(" alembic-v2-check:") :] + alembic_job = alembic_job[: alembic_job.index("\n v2-integration-check:")] + integration_job = workflow[workflow.index(" v2-integration-check:") :] + integration_job = integration_job[: integration_job.index("\n check-changelog:")] assert "alembic-v1-check:" in workflow assert "uses: ./.github/workflows/alembic-v1-check.yml" in workflow @@ -66,22 +68,40 @@ def test_pr_always_runs_reusable_alembic_check(): assert "detect-v2-platform-changes:" not in workflow assert "dorny/paths-filter" not in workflow assert "alembic-v2-check:" in workflow - assert "uses: ./.github/workflows/alembic-v2-check.yml" in v2_job - assert "CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}" in v2_job - assert "needs:" not in v2_job - assert "if:" not in v2_job - - -def test_push_always_runs_lint_and_alembic_qualification_before_versioning(): + assert "uses: ./.github/workflows/alembic-v2-check.yml" in alembic_job + assert "CODECOV_TOKEN" not in alembic_job + assert "v2-integration-check:" in workflow + assert "uses: ./.github/workflows/v2-integration-check.yml" in integration_job + assert "CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}" in integration_job + assert "needs:" not in alembic_job + assert "if:" not in alembic_job + assert "needs:" not in integration_job + assert "if:" not in integration_job + + +def test_push_requires_schema_and_v2_integration_checks_before_versioning(): workflow = _workflow("push.yml") + versioning_job = workflow[workflow.index(" versioning:") :] + versioning_job = versioning_job[: versioning_job.index("\n publish-git-tag:")] + tag_job = workflow[workflow.index(" publish-git-tag:") :] + tag_job = tag_job[: tag_job.index("\n migrate-v1-cloud-sql:")] assert "lint:" in workflow assert "alembic-v1-check:" in workflow assert "uses: ./.github/workflows/alembic-v1-check.yml" in workflow assert "alembic-v2-check:" in workflow assert "uses: ./.github/workflows/alembic-v2-check.yml" in workflow + assert "v2-integration-check:" in workflow + assert "uses: ./.github/workflows/v2-integration-check.yml" in workflow assert "CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}" in workflow - assert "needs: [lint, alembic-v1-check, alembic-v2-check]" in workflow + for required_job in ( + "lint", + "alembic-v1-check", + "alembic-v2-check", + "v2-integration-check", + ): + assert f"- {required_job}" in versioning_job + assert f"- {required_job}" in tag_job assert "github.repository == 'PolicyEngine/policyengine-uk'" not in workflow @@ -118,7 +138,7 @@ def test_reusable_alembic_check_uses_the_installed_python_environment(): assert "uv run" not in workflow -def test_reusable_v2_check_uses_disposable_postgres_and_real_redis(): +def test_reusable_v2_alembic_check_uses_only_disposable_postgres(): workflow = _workflow("alembic-v2-check.yml") lifecycle_script = ( REPO / ".github" / "scripts" / "test_alembic_v2_lifecycle.sh" @@ -127,12 +147,33 @@ def test_reusable_v2_check_uses_disposable_postgres_and_real_redis(): assert "workflow_call:" in workflow assert "workflow_dispatch:" in workflow assert "postgres:17" in workflow - assert "redis:7.2-alpine" in workflow assert "V2_ALEMBIC_DISPOSABLE_TEST" in workflow assert "alembic-v2.ini" in workflow assert "bash .github/scripts/test_alembic_v2_lifecycle.sh" in workflow assert "test_alembic_v2.py" in lifecycle_script assert "test_alembic_v2_lifecycle.py" in lifecycle_script + assert "redis:7.2-alpine" not in workflow + assert "RUNTIME_CACHE_TEST_URL" not in workflow + assert "test_v2_catalog_installed.py" not in workflow + assert "test_v2_catalog_publication.py" not in workflow + assert "test_v2_metadata_routes.py" not in workflow + assert "test_v2_catalog_publication_qualification.py" not in workflow + assert "test_runtime_cache_redis.py" not in workflow + assert "coverage run" not in workflow + assert "codecov/codecov-action" not in workflow + + +def test_reusable_v2_integration_check_uses_postgres_redis_and_coverage(): + workflow = _workflow("v2-integration-check.yml") + + assert "workflow_call:" in workflow + assert "workflow_dispatch:" in workflow + assert "postgres:17" in workflow + assert "redis:7.2-alpine" in workflow + assert "V2_ALEMBIC_DISPOSABLE_TEST" in workflow + assert "RUNTIME_CACHE_TEST_URL" in workflow + assert "alembic -c alembic-v2.ini upgrade head" in workflow + assert "test_alembic_v2_lifecycle.sh" not in workflow assert "test_v2_catalog_installed.py" in workflow assert "RUN_V2_CATALOG_COMPATIBILITY" in workflow assert "test_v2_catalog_publication.py" in workflow From 45b43fde6e5f1969cb8d6f462e6a70b89d8789ef Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 16:14:50 +0400 Subject: [PATCH 09/27] Rename v2 database seeding workflow --- .github/workflows/push.yml | 20 ++++++------- ...e-v2-metadata.yml => seed-v2-database.yml} | 8 ++--- docs/migration/stage-9-v2-metadata.md | 2 +- tests/unit/test_cloud_run_deploy_scripts.py | 2 +- tests/unit/v2/test_metadata_deployment.py | 29 ++++++++++--------- 5 files changed, 31 insertions(+), 30 deletions(-) rename .github/workflows/{initialize-v2-metadata.yml => seed-v2-database.yml} (89%) diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index 7bd174bc9..c54671259 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -165,15 +165,15 @@ jobs: if: always() run: bash .github/scripts/stop_cloud_sql_proxy.sh - initialize-v2-staging: - name: Initialize staging v2 metadata + seed-v2-staging-database: + name: Seed staging v2 database needs: - publish-git-tag - migrate-v1-cloud-sql if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') - uses: ./.github/workflows/initialize-v2-metadata.yml + uses: ./.github/workflows/seed-v2-database.yml with: deployment_environment: staging secrets: inherit @@ -185,7 +185,7 @@ jobs: - ensure-staging-model-version-aligns-with-sim-api - publish-git-tag - migrate-v1-cloud-sql - - initialize-v2-staging + - seed-v2-staging-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -281,7 +281,7 @@ jobs: - ensure-staging-model-version-aligns-with-sim-api - publish-git-tag - migrate-v1-cloud-sql - - initialize-v2-staging + - seed-v2-staging-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -519,13 +519,13 @@ jobs: - name: Check simulation API supports PolicyEngine bundle run: bash .github/check-policyengine-bundle-supported.sh - initialize-v2-production: - name: Initialize production v2 metadata + seed-v2-production-database: + name: Seed production v2 database needs: ensure-production-model-version-aligns-with-sim-api if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') - uses: ./.github/workflows/initialize-v2-metadata.yml + uses: ./.github/workflows/seed-v2-database.yml with: deployment_environment: production secrets: inherit @@ -533,7 +533,7 @@ jobs: deploy-production-candidate: name: Deploy production App Engine candidate runs-on: ubuntu-latest - needs: initialize-v2-production + needs: seed-v2-production-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -659,7 +659,7 @@ jobs: deploy-cloud-run-candidate: name: Deploy production Cloud Run candidate runs-on: ubuntu-latest - needs: initialize-v2-production + needs: seed-v2-production-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') diff --git a/.github/workflows/initialize-v2-metadata.yml b/.github/workflows/seed-v2-database.yml similarity index 89% rename from .github/workflows/initialize-v2-metadata.yml rename to .github/workflows/seed-v2-database.yml index e4b242a8d..aeed6d698 100644 --- a/.github/workflows/initialize-v2-metadata.yml +++ b/.github/workflows/seed-v2-database.yml @@ -1,4 +1,4 @@ -name: Initialize v2 metadata +name: Seed v2 database on: workflow_call: @@ -15,8 +15,8 @@ on: type: environment jobs: - initialize: - name: Upgrade and initialize v2 metadata + seed: + name: Upgrade schema and seed v2 database runs-on: ubuntu-latest timeout-minutes: 30 environment: ${{ inputs.deployment_environment }} @@ -38,7 +38,7 @@ jobs: run: bash .github/scripts/migrate_v2_metadata_schema.sh env: V2_MIGRATION_DATABASE_URL: ${{ secrets.V2_MIGRATION_DATABASE_URL }} - - name: Publish and validate the v2 metadata catalog + - name: Seed and validate the v2 metadata catalog run: uv run python scripts/initialize_v2_metadata.py env: V2_DATA_WRITE_DATABASE_URL: ${{ secrets.V2_DATA_WRITE_DATABASE_URL }} diff --git a/docs/migration/stage-9-v2-metadata.md b/docs/migration/stage-9-v2-metadata.md index 1f3b4b825..fc3f92859 100644 --- a/docs/migration/stage-9-v2-metadata.md +++ b/docs/migration/stage-9-v2-metadata.md @@ -48,7 +48,7 @@ catalog-publication credentials. ## Pre-activation sequence -The release workflow calls `.github/workflows/initialize-v2-metadata.yml` for +The release workflow calls `.github/workflows/seed-v2-database.yml` for the selected GitHub Environment before creating an API candidate. Its steps must execute in this order: diff --git a/tests/unit/test_cloud_run_deploy_scripts.py b/tests/unit/test_cloud_run_deploy_scripts.py index f2c7d6acd..97aa1608a 100644 --- a/tests/unit/test_cloud_run_deploy_scripts.py +++ b/tests/unit/test_cloud_run_deploy_scripts.py @@ -1785,7 +1785,7 @@ def test_push_workflow_staging_fully_gates_all_production_deployments(): docker_publish = _workflow_job_block(workflow, "docker") cloud_run_production = _workflow_job_block(workflow, "deploy-cloud-run-candidate") - production_initialization_dependency = "needs: initialize-v2-production" + production_initialization_dependency = "needs: seed-v2-production-database" assert production_initialization_dependency in app_engine_candidate assert 'APP_ENGINE_PROMOTE: "0"' in app_engine_candidate assert ( diff --git a/tests/unit/v2/test_metadata_deployment.py b/tests/unit/v2/test_metadata_deployment.py index cb6a7e806..7b8ee4ade 100644 --- a/tests/unit/v2/test_metadata_deployment.py +++ b/tests/unit/v2/test_metadata_deployment.py @@ -32,18 +32,20 @@ def _step(workflow: str, name: str, next_name: str | None = None) -> str: return workflow[start:end] -def test_reusable_initialization_workflow_separates_database_credentials() -> None: - workflow = _read(".github/workflows/initialize-v2-metadata.yml") +def test_reusable_seeding_workflow_separates_database_credentials() -> None: + workflow = _read(".github/workflows/seed-v2-database.yml") migration = _step( workflow, "Upgrade and verify the v2 schema", - "Publish and validate the v2 metadata catalog", + "Seed and validate the v2 metadata catalog", ) publication = _step( workflow, - "Publish and validate the v2 metadata catalog", + "Seed and validate the v2 metadata catalog", ) + assert not (REPO / ".github/workflows/initialize-v2-metadata.yml").exists() + assert workflow.startswith("name: Seed v2 database\n") assert "workflow_call:" in workflow assert "workflow_dispatch:" in workflow assert "environment: ${{ inputs.deployment_environment }}" in workflow @@ -65,23 +67,22 @@ def test_schema_upgrade_precedes_atomic_catalog_publication() -> None: assert upgrade < current < drift -def test_initialization_success_is_required_before_candidate_creation() -> None: +def test_seeding_success_is_required_before_candidate_creation() -> None: workflow = _read(".github/workflows/push.yml") - staging_initialization = _job(workflow, "initialize-v2-staging") - production_initialization = _job(workflow, "initialize-v2-production") + staging_seed = _job(workflow, "seed-v2-staging-database") + production_seed = _job(workflow, "seed-v2-production-database") - assert "deployment_environment: staging" in staging_initialization - assert "migrate-v1-cloud-sql" in staging_initialization + assert "deployment_environment: staging" in staging_seed + assert "migrate-v1-cloud-sql" in staging_seed for job_name in ("deploy-staging", "deploy-cloud-run-staging"): - assert "initialize-v2-staging" in _job(workflow, job_name) + assert "seed-v2-staging-database" in _job(workflow, job_name) assert ( - "needs: ensure-production-model-version-aligns-with-sim-api" - in production_initialization + "needs: ensure-production-model-version-aligns-with-sim-api" in production_seed ) - assert "deployment_environment: production" in production_initialization + assert "deployment_environment: production" in production_seed for job_name in ("deploy-production-candidate", "deploy-cloud-run-candidate"): - assert "needs: initialize-v2-production" in _job(workflow, job_name) + assert "needs: seed-v2-production-database" in _job(workflow, job_name) def test_stage_9_deployment_shell_script_is_syntax_valid() -> None: From 9c7253bb0aad3765bd2e840759b456ca0bb677b6 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 16:31:57 +0400 Subject: [PATCH 10/27] Split v2 catalog publication module --- .../data/v2/catalog/publication.py | 847 +----------------- .../v2/catalog/publication_reconciliation.py | 436 +++++++++ .../data/v2/catalog/publication_staging.py | 365 ++++++++ .../data/v2/catalog/publication_types.py | 37 + tests/unit/v2/test_catalog_publication.py | 21 +- 5 files changed, 880 insertions(+), 826 deletions(-) create mode 100644 policyengine_api/data/v2/catalog/publication_reconciliation.py create mode 100644 policyengine_api/data/v2/catalog/publication_staging.py create mode 100644 policyengine_api/data/v2/catalog/publication_types.py diff --git a/policyengine_api/data/v2/catalog/publication.py b/policyengine_api/data/v2/catalog/publication.py index 3ecf34937..e2addc599 100644 --- a/policyengine_api/data/v2/catalog/publication.py +++ b/policyengine_api/data/v2/catalog/publication.py @@ -2,368 +2,39 @@ from __future__ import annotations -from collections.abc import Callable, Iterable, Iterator, Sequence -from dataclasses import dataclass +from collections.abc import Callable import logging import time -from psycopg.types.json import Jsonb from sqlalchemy import Connection, Engine, text -from policyengine_api.data.v2.catalog.records import ( - CountryCatalog, - NormalizedCatalog, - iter_batches, +from policyengine_api.data.v2.catalog.publication_reconciliation import ( + assert_canonical_value_uniqueness, + assert_country_matches, + publish_new_country, + version_exists, ) +from policyengine_api.data.v2.catalog.publication_staging import ( + create_staging_tables, + stage_catalog, +) +from policyengine_api.data.v2.catalog.publication_types import ( + CatalogPublicationError, + PublicationEvidence, +) +from policyengine_api.data.v2.catalog.records import NormalizedCatalog EXPECTED_ALEMBIC_REVISION = "68b4a5ae5dc5" PUBLICATION_ADVISORY_LOCK_KEY = 8_629_020_026_090_001 -COPY_BATCH_SIZE = 10_000 LOGGER = logging.getLogger(__name__) - -class CatalogPublicationError(RuntimeError): - """Raised when publication cannot prove an atomic, complete result.""" - - -@dataclass(frozen=True, slots=True) -class PublicationEvidence: - """Non-secret facts emitted after a successful publication.""" - - policyengine_version: str - dependency_versions: tuple[tuple[str, str], ...] - entity_counts: dict[str, int] - fallback_summaries: tuple[tuple[str, str, int], ...] - elapsed_seconds: float - - def as_dict(self) -> dict[str, object]: - return { - "outcome": "ok", - "policyengine_version": self.policyengine_version, - "dependency_versions": dict(self.dependency_versions), - "entity_counts": self.entity_counts, - "fallback_summaries": [ - { - "country_id": country_id, - "region_type": region_type, - "count": count, - } - for country_id, region_type, count in self.fallback_summaries - ], - "elapsed_seconds": round(self.elapsed_seconds, 3), - } - - -TEMP_TABLE_STATEMENTS = ( - """ - CREATE TEMP TABLE stage_catalog_models ( - country_id text PRIMARY KEY, - id uuid NOT NULL, - name text NOT NULL, - description text, - version_id uuid NOT NULL, - version text NOT NULL, - version_description text, - current_law_id integer NOT NULL, - metadata_time_periods jsonb NOT NULL - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_variables ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - label text, - entity text NOT NULL, - description text, - data_type text, - possible_values jsonb, - default_value jsonb, - adds jsonb, - subtracts jsonb, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_parameter_nodes ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - label text, - description text, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_parameters ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - label text, - description text, - data_type text, - unit text, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_parameter_values ( - country_id text NOT NULL, - parameter_name text NOT NULL, - id uuid NOT NULL, - value_json jsonb NOT NULL, - start_date timestamptz NOT NULL, - end_date timestamptz, - PRIMARY KEY (country_id, parameter_name, start_date) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_datasets ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - description text, - year integer NOT NULL, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_regions ( - country_id text NOT NULL, - id uuid NOT NULL, - code text NOT NULL, - label text NOT NULL, - region_type text NOT NULL, - requires_filter boolean NOT NULL, - filter_field text, - filter_value text, - filter_strategy text, - parent_code text, - state_code text, - state_name text, - default_dataset_name text NOT NULL, - PRIMARY KEY (country_id, code) - ) ON COMMIT DROP - """, -) - - -COPY_COLUMNS = { - "stage_catalog_models": ( - "country_id", - "id", - "name", - "description", - "version_id", - "version", - "version_description", - "current_law_id", - "metadata_time_periods", - ), - "stage_catalog_variables": ( - "country_id", - "id", - "name", - "label", - "entity", - "description", - "data_type", - "possible_values", - "default_value", - "adds", - "subtracts", - ), - "stage_catalog_parameter_nodes": ( - "country_id", - "id", - "name", - "label", - "description", - ), - "stage_catalog_parameters": ( - "country_id", - "id", - "name", - "label", - "description", - "data_type", - "unit", - ), - "stage_catalog_parameter_values": ( - "country_id", - "parameter_name", - "id", - "value_json", - "start_date", - "end_date", - ), - "stage_catalog_datasets": ( - "country_id", - "id", - "name", - "description", - "year", - ), - "stage_catalog_regions": ( - "country_id", - "id", - "code", - "label", - "region_type", - "requires_filter", - "filter_field", - "filter_value", - "filter_strategy", - "parent_code", - "state_code", - "state_name", - "default_dataset_name", - ), -} - - -def _optional_json(value: object) -> Jsonb | None: - return None if value is None else Jsonb(value) - - -def _catalog_rows( - country: CountryCatalog, -) -> dict[str, Iterator[tuple[object, ...]]]: - dataset_names = {dataset.id: dataset.name for dataset in country.datasets} - - def model_rows() -> Iterator[tuple[object, ...]]: - yield ( - country.country_id, - country.model.id, - country.model.name, - country.model.description, - country.model_version.id, - country.model_version.version, - country.model_version.description, - country.model_version.current_law_id, - Jsonb(country.model_version.metadata_time_periods), - ) - - def variable_rows() -> Iterator[tuple[object, ...]]: - for batch in iter_batches(country.variables, batch_size=COPY_BATCH_SIZE): - for record in batch: - yield ( - country.country_id, - record.id, - record.name, - record.label, - record.entity, - record.description, - record.data_type, - _optional_json(record.possible_values), - Jsonb(record.default_value), - _optional_json(record.adds), - _optional_json(record.subtracts), - ) - - def parameter_node_rows() -> Iterator[tuple[object, ...]]: - for batch in iter_batches( - country.parameter_nodes, - batch_size=COPY_BATCH_SIZE, - ): - for record in batch: - yield ( - country.country_id, - record.id, - record.name, - record.label, - record.description, - ) - - def parameter_rows() -> Iterator[tuple[object, ...]]: - for batch in iter_batches(country.parameters, batch_size=COPY_BATCH_SIZE): - for record in batch: - yield ( - country.country_id, - record.id, - record.name, - record.label, - record.description, - record.data_type, - record.unit, - ) - - def parameter_value_rows() -> Iterator[tuple[object, ...]]: - parameter_names = { - parameter.id: parameter.name for parameter in country.parameters - } - for batch in country.parameter_value_batches(batch_size=COPY_BATCH_SIZE): - for record in batch: - yield ( - country.country_id, - parameter_names[record.parameter_id], - record.id, - Jsonb(record.value_json), - record.start_date, - record.end_date, - ) - - def dataset_rows() -> Iterator[tuple[object, ...]]: - for batch in iter_batches(country.datasets, batch_size=COPY_BATCH_SIZE): - for record in batch: - yield ( - country.country_id, - record.id, - record.name, - record.description, - record.year, - ) - - def region_rows() -> Iterator[tuple[object, ...]]: - for batch in iter_batches(country.regions, batch_size=COPY_BATCH_SIZE): - for record in batch: - yield ( - country.country_id, - record.id, - record.code, - record.label, - record.region_type, - record.requires_filter, - record.filter_field, - record.filter_value, - record.filter_strategy, - record.parent_code, - record.state_code, - record.state_name, - dataset_names[record.default_dataset_id], - ) - - return { - "stage_catalog_models": model_rows(), - "stage_catalog_variables": variable_rows(), - "stage_catalog_parameter_nodes": parameter_node_rows(), - "stage_catalog_parameters": parameter_rows(), - "stage_catalog_parameter_values": parameter_value_rows(), - "stage_catalog_datasets": dataset_rows(), - "stage_catalog_regions": region_rows(), - } - - -def _copy_rows( - connection: Connection, - *, - table_name: str, - columns: Sequence[str], - rows: Iterable[Sequence[object]], -) -> int: - """Write one bounded source stream through Psycopg COPY.""" - - raw_connection = connection.connection.driver_connection - statement = f"COPY {table_name} ({', '.join(columns)}) FROM STDIN" - count = 0 - with raw_connection.cursor() as cursor: - with cursor.copy(statement) as copy: - for row in rows: - copy.write_row(row) - count += 1 - return count +__all__ = [ + "CatalogPublicationError", + "PublicationEvidence", + "publish_catalog", +] def _verify_expected_revision(connection: Connection) -> None: @@ -390,43 +61,6 @@ def _acquire_publication_lock(connection: Connection) -> None: ).scalar_one() -def _create_staging_tables(connection: Connection) -> None: - for statement in TEMP_TABLE_STATEMENTS: - connection.execute(text(statement)) - - -def _stage_catalog( - connection: Connection, - catalog: NormalizedCatalog, - *, - checkpoint: Callable[[str, Connection], None] | None = None, -) -> dict[str, int]: - observed = {table_name: 0 for table_name in COPY_COLUMNS} - for country in catalog.countries: - for table_name, rows in _catalog_rows(country).items(): - observed[table_name] += _copy_rows( - connection, - table_name=table_name, - columns=COPY_COLUMNS[table_name], - rows=rows, - ) - if checkpoint is not None: - checkpoint("during_copy", connection) - expected = catalog.entity_counts() - expected_by_table = { - "stage_catalog_models": expected["models"], - "stage_catalog_variables": expected["variables"], - "stage_catalog_parameter_nodes": expected["parameter_nodes"], - "stage_catalog_parameters": expected["parameters"], - "stage_catalog_parameter_values": expected["parameter_values"], - "stage_catalog_datasets": expected["datasets"], - "stage_catalog_regions": expected["regions"], - } - if observed != expected_by_table: - raise CatalogPublicationError("COPY row counts differ from the catalog") - return observed - - def _protected_row_counts(connection: Connection) -> tuple[int, int, int, int]: return tuple( connection.execute(text(f"SELECT count(*) FROM {table_name}")).scalar_one() @@ -439,433 +73,6 @@ def _protected_row_counts(connection: Connection) -> tuple[int, int, int, int]: ) -def _version_exists(connection: Connection, country_id: str) -> bool: - return bool( - connection.execute( - text( - """ - SELECT EXISTS ( - SELECT 1 - FROM stage_catalog_models AS source - JOIN tax_benefit_models AS model - ON model.name = source.name - JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source.version - WHERE source.country_id = :country_id - ) - """ - ), - {"country_id": country_id}, - ).scalar_one() - ) - - -MODEL_DIFFERENCE_SQL = """ -SELECT EXISTS ( - SELECT 1 - FROM stage_catalog_models AS source - JOIN tax_benefit_models AS model ON model.name = source.name - JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source.version - WHERE source.country_id = :country_id - AND ( - model_version.description IS DISTINCT FROM source.version_description - OR model_version.current_law_id IS DISTINCT FROM source.current_law_id - OR model_version.metadata_time_periods::jsonb - IS DISTINCT FROM source.metadata_time_periods - ) -) -""" - - -VERSIONED_DIFFERENCE_SQL = ( - """ - WITH staged AS ( - SELECT name, label, entity, description, data_type, possible_values, - default_value, adds, subtracts - FROM stage_catalog_variables - WHERE country_id = :country_id - ), actual AS ( - SELECT variable.name, variable.label, variable.entity, - variable.description, variable.data_type, - variable.possible_values::jsonb, - variable.default_value::jsonb, - variable.adds::jsonb, variable.subtracts::jsonb - FROM variables AS variable - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = variable.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT name, label, description - FROM stage_catalog_parameter_nodes - WHERE country_id = :country_id - ), actual AS ( - SELECT node.name, node.label, node.description - FROM parameter_nodes AS node - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = node.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT name, label, description, data_type, unit - FROM stage_catalog_parameters - WHERE country_id = :country_id - ), actual AS ( - SELECT parameter.name, parameter.label, parameter.description, - parameter.data_type, parameter.unit - FROM parameters AS parameter - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = parameter.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT parameter_name, value_json, start_date, end_date - FROM stage_catalog_parameter_values - WHERE country_id = :country_id - ), actual AS ( - SELECT parameter.name, parameter_value.value_json::jsonb, - parameter_value.start_date, parameter_value.end_date - FROM parameter_values AS parameter_value - JOIN parameters AS parameter ON parameter.id = parameter_value.parameter_id - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = parameter.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - AND parameter_value.policy_id IS NULL - AND parameter_value.dynamic_id IS NULL - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, -) - - -VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL = ( - """ - WITH staged AS ( - SELECT name, description, year, false AS is_output_dataset, - NULL::text AS storage_path - FROM stage_catalog_datasets - WHERE country_id = :country_id - ), actual AS ( - SELECT dataset.name, dataset.description, dataset.year, - dataset.is_output_dataset, dataset.storage_path - FROM datasets AS dataset - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = dataset.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - AND NOT dataset.is_output_dataset - AND dataset.storage_path IS NULL - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT code, label, region_type, requires_filter, filter_field, - filter_value, filter_strategy, parent_code, state_code, - state_name, default_dataset_name - FROM stage_catalog_regions - WHERE country_id = :country_id - ), actual AS ( - SELECT region.code, region.label, region.region_type::text, - region.requires_filter, region.filter_field, - region.filter_value, region.filter_strategy, - region.parent_code, region.state_code, region.state_name, - dataset.name - FROM regions AS region - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = region.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN datasets AS dataset ON dataset.id = region.default_dataset_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, -) - - -def _assert_country_matches(connection: Connection, country_id: str) -> None: - statements = ( - MODEL_DIFFERENCE_SQL, - *VERSIONED_DIFFERENCE_SQL, - *VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL, - ) - for statement in statements: - differs = connection.execute( - text(statement), - {"country_id": country_id}, - ).scalar_one() - if differs: - raise CatalogPublicationError( - f"persisted {country_id} catalog differs from PolicyEngine.py" - ) - - -def _reject_dataset_role_conflicts( - connection: Connection, - country_id: str, -) -> None: - conflict = connection.execute( - text( - """ - SELECT EXISTS ( - SELECT 1 - FROM stage_catalog_datasets AS source_dataset - JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_dataset.country_id - JOIN tax_benefit_models AS model - ON model.name = source_model.name - JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version - JOIN datasets AS dataset - ON dataset.tax_benefit_model_version_id = model_version.id - AND dataset.name = source_dataset.name - WHERE source_dataset.country_id = :country_id - AND ( - dataset.is_output_dataset - OR dataset.storage_path IS NOT NULL - ) - ) - """ - ), - {"country_id": country_id}, - ).scalar_one() - if conflict: - raise CatalogPublicationError( - f"persisted {country_id} dataset identity is not an input dataset" - ) - - -INSERT_MODEL_SQL = """ -INSERT INTO tax_benefit_models ( - id, created_at, updated_at, name, description -) -SELECT id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, name, description -FROM stage_catalog_models -WHERE country_id = :country_id -ON CONFLICT (name) DO NOTHING -""" - -INSERT_MODEL_VERSION_SQL = """ -INSERT INTO tax_benefit_model_versions ( - id, created_at, model_id, version, description, current_law_id, - metadata_time_periods -) -SELECT source.version_id, CURRENT_TIMESTAMP, model.id, - source.version, source.version_description, source.current_law_id, - source.metadata_time_periods::json -FROM stage_catalog_models AS source -JOIN tax_benefit_models AS model ON model.name = source.name -WHERE source.country_id = :country_id -""" - -INSERT_DATASETS_SQL = """ -INSERT INTO datasets ( - id, created_at, updated_at, name, description, storage_path, year, - is_output_dataset, tax_benefit_model_version_id -) -SELECT source_dataset.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, - source_dataset.name, source_dataset.description, NULL, - source_dataset.year, false, model_version.id -FROM stage_catalog_datasets AS source_dataset -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_dataset.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_dataset.country_id = :country_id -""" - -INSERT_VARIABLES_SQL = """ -INSERT INTO variables ( - id, created_at, name, label, entity, description, data_type, - possible_values, default_value, adds, subtracts, - tax_benefit_model_version_id -) -SELECT source_variable.id, CURRENT_TIMESTAMP, source_variable.name, - source_variable.label, source_variable.entity, - source_variable.description, source_variable.data_type, - source_variable.possible_values::json, - source_variable.default_value::json, - source_variable.adds::json, source_variable.subtracts::json, - model_version.id -FROM stage_catalog_variables AS source_variable -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_variable.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_variable.country_id = :country_id -""" - -INSERT_PARAMETER_NODES_SQL = """ -INSERT INTO parameter_nodes ( - id, created_at, name, label, description, tax_benefit_model_version_id -) -SELECT source_node.id, CURRENT_TIMESTAMP, source_node.name, - source_node.label, source_node.description, model_version.id -FROM stage_catalog_parameter_nodes AS source_node -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_node.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_node.country_id = :country_id -""" - -INSERT_PARAMETERS_SQL = """ -INSERT INTO parameters ( - id, created_at, name, label, description, data_type, unit, - tax_benefit_model_version_id -) -SELECT source_parameter.id, CURRENT_TIMESTAMP, source_parameter.name, - source_parameter.label, source_parameter.description, - source_parameter.data_type, source_parameter.unit, model_version.id -FROM stage_catalog_parameters AS source_parameter -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_parameter.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_parameter.country_id = :country_id -""" - -INSERT_PARAMETER_VALUES_SQL = """ -INSERT INTO parameter_values ( - id, created_at, parameter_id, value_json, start_date, end_date, - policy_id, dynamic_id -) -SELECT source_value.id, CURRENT_TIMESTAMP, parameter.id, - source_value.value_json::json, source_value.start_date, - source_value.end_date, NULL, NULL -FROM stage_catalog_parameter_values AS source_value -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_value.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -JOIN parameters AS parameter - ON parameter.tax_benefit_model_version_id = model_version.id - AND parameter.name = source_value.parameter_name -WHERE source_value.country_id = :country_id -""" - -INSERT_REGIONS_SQL = """ -INSERT INTO regions ( - id, created_at, updated_at, code, label, region_type, requires_filter, - filter_field, filter_value, filter_strategy, parent_code, state_code, - state_name, tax_benefit_model_version_id, default_dataset_id -) -SELECT source_region.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, - source_region.code, source_region.label, - source_region.region_type::v2_region_type, - source_region.requires_filter, source_region.filter_field, - source_region.filter_value, source_region.filter_strategy, - source_region.parent_code, source_region.state_code, - source_region.state_name, model_version.id, dataset.id -FROM stage_catalog_regions AS source_region -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_region.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -JOIN datasets AS dataset - ON dataset.tax_benefit_model_version_id = model_version.id - AND dataset.name = source_region.default_dataset_name -WHERE source_region.country_id = :country_id -""" - - -SET_BASED_INSERT_SQL = ( - INSERT_MODEL_SQL, - INSERT_MODEL_VERSION_SQL, - INSERT_DATASETS_SQL, - INSERT_VARIABLES_SQL, - INSERT_PARAMETER_NODES_SQL, - INSERT_PARAMETERS_SQL, - INSERT_PARAMETER_VALUES_SQL, - INSERT_REGIONS_SQL, -) - - -def _publish_new_country(connection: Connection, country_id: str) -> None: - _reject_dataset_role_conflicts(connection, country_id) - for statement in SET_BASED_INSERT_SQL: - connection.execute(text(statement), {"country_id": country_id}) - - -def _assert_canonical_value_uniqueness(connection: Connection) -> None: - duplicates = connection.execute( - text( - """ - SELECT EXISTS ( - SELECT 1 - FROM parameter_values - WHERE policy_id IS NULL AND dynamic_id IS NULL - GROUP BY parameter_id, start_date - HAVING count(*) > 1 - ) - """ - ) - ).scalar_one() - if duplicates: - raise CatalogPublicationError( - "canonical parameter-value uniqueness validation failed" - ) - - def publish_catalog( engine: Engine, catalog: NormalizedCatalog, @@ -879,26 +86,26 @@ def publish_catalog( _verify_expected_revision(connection) _acquire_publication_lock(connection) before = _protected_row_counts(connection) - _create_staging_tables(connection) - _stage_catalog(connection, catalog, checkpoint=checkpoint) + create_staging_tables(connection) + stage_catalog(connection, catalog, checkpoint=checkpoint) if checkpoint is not None: checkpoint("after_copy", connection) existing: set[str] = set() for country in catalog.countries: - if _version_exists(connection, country.country_id): - _assert_country_matches(connection, country.country_id) + if version_exists(connection, country.country_id): + assert_country_matches(connection, country.country_id) existing.add(country.country_id) for country in catalog.countries: if country.country_id not in existing: - _publish_new_country(connection, country.country_id) + publish_new_country(connection, country.country_id) if checkpoint is not None: checkpoint("after_reconciliation", connection) for country in catalog.countries: - _assert_country_matches(connection, country.country_id) - _assert_canonical_value_uniqueness(connection) + assert_country_matches(connection, country.country_id) + assert_canonical_value_uniqueness(connection) if _protected_row_counts(connection) != before: raise CatalogPublicationError( "publication changed simulation, report, or dataset-version rows" diff --git a/policyengine_api/data/v2/catalog/publication_reconciliation.py b/policyengine_api/data/v2/catalog/publication_reconciliation.py new file mode 100644 index 000000000..a6c595475 --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication_reconciliation.py @@ -0,0 +1,436 @@ +"""Set-based reconciliation and validation for a staged v2 catalog.""" + +from __future__ import annotations + +from sqlalchemy import Connection, text + +from policyengine_api.data.v2.catalog.publication_types import ( + CatalogPublicationError, +) + + +def version_exists(connection: Connection, country_id: str) -> bool: + return bool( + connection.execute( + text( + """ + SELECT EXISTS ( + SELECT 1 + FROM stage_catalog_models AS source + JOIN tax_benefit_models AS model + ON model.name = source.name + JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source.version + WHERE source.country_id = :country_id + ) + """ + ), + {"country_id": country_id}, + ).scalar_one() + ) + + +MODEL_DIFFERENCE_SQL = """ +SELECT EXISTS ( + SELECT 1 + FROM stage_catalog_models AS source + JOIN tax_benefit_models AS model ON model.name = source.name + JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source.version + WHERE source.country_id = :country_id + AND ( + model_version.description IS DISTINCT FROM source.version_description + OR model_version.current_law_id IS DISTINCT FROM source.current_law_id + OR model_version.metadata_time_periods::jsonb + IS DISTINCT FROM source.metadata_time_periods + ) +) +""" + + +VERSIONED_DIFFERENCE_SQL = ( + """ + WITH staged AS ( + SELECT name, label, entity, description, data_type, possible_values, + default_value, adds, subtracts + FROM stage_catalog_variables + WHERE country_id = :country_id + ), actual AS ( + SELECT variable.name, variable.label, variable.entity, + variable.description, variable.data_type, + variable.possible_values::jsonb, + variable.default_value::jsonb, + variable.adds::jsonb, variable.subtracts::jsonb + FROM variables AS variable + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = variable.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT name, label, description + FROM stage_catalog_parameter_nodes + WHERE country_id = :country_id + ), actual AS ( + SELECT node.name, node.label, node.description + FROM parameter_nodes AS node + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = node.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT name, label, description, data_type, unit + FROM stage_catalog_parameters + WHERE country_id = :country_id + ), actual AS ( + SELECT parameter.name, parameter.label, parameter.description, + parameter.data_type, parameter.unit + FROM parameters AS parameter + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = parameter.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT parameter_name, value_json, start_date, end_date + FROM stage_catalog_parameter_values + WHERE country_id = :country_id + ), actual AS ( + SELECT parameter.name, parameter_value.value_json::jsonb, + parameter_value.start_date, parameter_value.end_date + FROM parameter_values AS parameter_value + JOIN parameters AS parameter ON parameter.id = parameter_value.parameter_id + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = parameter.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + AND parameter_value.policy_id IS NULL + AND parameter_value.dynamic_id IS NULL + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, +) + + +VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL = ( + """ + WITH staged AS ( + SELECT name, description, year, false AS is_output_dataset, + NULL::text AS storage_path + FROM stage_catalog_datasets + WHERE country_id = :country_id + ), actual AS ( + SELECT dataset.name, dataset.description, dataset.year, + dataset.is_output_dataset, dataset.storage_path + FROM datasets AS dataset + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = dataset.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + AND NOT dataset.is_output_dataset + AND dataset.storage_path IS NULL + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, + """ + WITH staged AS ( + SELECT code, label, region_type, requires_filter, filter_field, + filter_value, filter_strategy, parent_code, state_code, + state_name, default_dataset_name + FROM stage_catalog_regions + WHERE country_id = :country_id + ), actual AS ( + SELECT region.code, region.label, region.region_type::text, + region.requires_filter, region.filter_field, + region.filter_value, region.filter_strategy, + region.parent_code, region.state_code, region.state_name, + dataset.name + FROM regions AS region + JOIN tax_benefit_model_versions AS model_version + ON model_version.id = region.tax_benefit_model_version_id + JOIN tax_benefit_models AS model ON model.id = model_version.model_id + JOIN datasets AS dataset ON dataset.id = region.default_dataset_id + JOIN stage_catalog_models AS source + ON source.name = model.name AND source.version = model_version.version + WHERE source.country_id = :country_id + ), differences AS ( + (SELECT * FROM staged EXCEPT SELECT * FROM actual) + UNION ALL + (SELECT * FROM actual EXCEPT SELECT * FROM staged) + ) SELECT EXISTS (SELECT 1 FROM differences) + """, +) + + +def assert_country_matches(connection: Connection, country_id: str) -> None: + statements = ( + MODEL_DIFFERENCE_SQL, + *VERSIONED_DIFFERENCE_SQL, + *VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL, + ) + for statement in statements: + differs = connection.execute( + text(statement), + {"country_id": country_id}, + ).scalar_one() + if differs: + raise CatalogPublicationError( + f"persisted {country_id} catalog differs from PolicyEngine.py" + ) + + +def _reject_dataset_role_conflicts( + connection: Connection, + country_id: str, +) -> None: + conflict = connection.execute( + text( + """ + SELECT EXISTS ( + SELECT 1 + FROM stage_catalog_datasets AS source_dataset + JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_dataset.country_id + JOIN tax_benefit_models AS model + ON model.name = source_model.name + JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version + JOIN datasets AS dataset + ON dataset.tax_benefit_model_version_id = model_version.id + AND dataset.name = source_dataset.name + WHERE source_dataset.country_id = :country_id + AND ( + dataset.is_output_dataset + OR dataset.storage_path IS NOT NULL + ) + ) + """ + ), + {"country_id": country_id}, + ).scalar_one() + if conflict: + raise CatalogPublicationError( + f"persisted {country_id} dataset identity is not an input dataset" + ) + + +INSERT_MODEL_SQL = """ +INSERT INTO tax_benefit_models ( + id, created_at, updated_at, name, description +) +SELECT id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, name, description +FROM stage_catalog_models +WHERE country_id = :country_id +ON CONFLICT (name) DO NOTHING +""" + +INSERT_MODEL_VERSION_SQL = """ +INSERT INTO tax_benefit_model_versions ( + id, created_at, model_id, version, description, current_law_id, + metadata_time_periods +) +SELECT source.version_id, CURRENT_TIMESTAMP, model.id, + source.version, source.version_description, source.current_law_id, + source.metadata_time_periods::json +FROM stage_catalog_models AS source +JOIN tax_benefit_models AS model ON model.name = source.name +WHERE source.country_id = :country_id +""" + +INSERT_DATASETS_SQL = """ +INSERT INTO datasets ( + id, created_at, updated_at, name, description, storage_path, year, + is_output_dataset, tax_benefit_model_version_id +) +SELECT source_dataset.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, + source_dataset.name, source_dataset.description, NULL, + source_dataset.year, false, model_version.id +FROM stage_catalog_datasets AS source_dataset +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_dataset.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_dataset.country_id = :country_id +""" + +INSERT_VARIABLES_SQL = """ +INSERT INTO variables ( + id, created_at, name, label, entity, description, data_type, + possible_values, default_value, adds, subtracts, + tax_benefit_model_version_id +) +SELECT source_variable.id, CURRENT_TIMESTAMP, source_variable.name, + source_variable.label, source_variable.entity, + source_variable.description, source_variable.data_type, + source_variable.possible_values::json, + source_variable.default_value::json, + source_variable.adds::json, source_variable.subtracts::json, + model_version.id +FROM stage_catalog_variables AS source_variable +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_variable.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_variable.country_id = :country_id +""" + +INSERT_PARAMETER_NODES_SQL = """ +INSERT INTO parameter_nodes ( + id, created_at, name, label, description, tax_benefit_model_version_id +) +SELECT source_node.id, CURRENT_TIMESTAMP, source_node.name, + source_node.label, source_node.description, model_version.id +FROM stage_catalog_parameter_nodes AS source_node +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_node.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_node.country_id = :country_id +""" + +INSERT_PARAMETERS_SQL = """ +INSERT INTO parameters ( + id, created_at, name, label, description, data_type, unit, + tax_benefit_model_version_id +) +SELECT source_parameter.id, CURRENT_TIMESTAMP, source_parameter.name, + source_parameter.label, source_parameter.description, + source_parameter.data_type, source_parameter.unit, model_version.id +FROM stage_catalog_parameters AS source_parameter +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_parameter.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +WHERE source_parameter.country_id = :country_id +""" + +INSERT_PARAMETER_VALUES_SQL = """ +INSERT INTO parameter_values ( + id, created_at, parameter_id, value_json, start_date, end_date, + policy_id, dynamic_id +) +SELECT source_value.id, CURRENT_TIMESTAMP, parameter.id, + source_value.value_json::json, source_value.start_date, + source_value.end_date, NULL, NULL +FROM stage_catalog_parameter_values AS source_value +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_value.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +JOIN parameters AS parameter + ON parameter.tax_benefit_model_version_id = model_version.id + AND parameter.name = source_value.parameter_name +WHERE source_value.country_id = :country_id +""" + +INSERT_REGIONS_SQL = """ +INSERT INTO regions ( + id, created_at, updated_at, code, label, region_type, requires_filter, + filter_field, filter_value, filter_strategy, parent_code, state_code, + state_name, tax_benefit_model_version_id, default_dataset_id +) +SELECT source_region.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, + source_region.code, source_region.label, + source_region.region_type::v2_region_type, + source_region.requires_filter, source_region.filter_field, + source_region.filter_value, source_region.filter_strategy, + source_region.parent_code, source_region.state_code, + source_region.state_name, model_version.id, dataset.id +FROM stage_catalog_regions AS source_region +JOIN stage_catalog_models AS source_model + ON source_model.country_id = source_region.country_id +JOIN tax_benefit_models AS model ON model.name = source_model.name +JOIN tax_benefit_model_versions AS model_version + ON model_version.model_id = model.id + AND model_version.version = source_model.version +JOIN datasets AS dataset + ON dataset.tax_benefit_model_version_id = model_version.id + AND dataset.name = source_region.default_dataset_name +WHERE source_region.country_id = :country_id +""" + + +SET_BASED_INSERT_SQL = ( + INSERT_MODEL_SQL, + INSERT_MODEL_VERSION_SQL, + INSERT_DATASETS_SQL, + INSERT_VARIABLES_SQL, + INSERT_PARAMETER_NODES_SQL, + INSERT_PARAMETERS_SQL, + INSERT_PARAMETER_VALUES_SQL, + INSERT_REGIONS_SQL, +) + + +def publish_new_country(connection: Connection, country_id: str) -> None: + _reject_dataset_role_conflicts(connection, country_id) + for statement in SET_BASED_INSERT_SQL: + connection.execute(text(statement), {"country_id": country_id}) + + +def assert_canonical_value_uniqueness(connection: Connection) -> None: + duplicates = connection.execute( + text( + """ + SELECT EXISTS ( + SELECT 1 + FROM parameter_values + WHERE policy_id IS NULL AND dynamic_id IS NULL + GROUP BY parameter_id, start_date + HAVING count(*) > 1 + ) + """ + ) + ).scalar_one() + if duplicates: + raise CatalogPublicationError( + "canonical parameter-value uniqueness validation failed" + ) diff --git a/policyengine_api/data/v2/catalog/publication_staging.py b/policyengine_api/data/v2/catalog/publication_staging.py new file mode 100644 index 000000000..747023f1c --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication_staging.py @@ -0,0 +1,365 @@ +"""Temporary-table staging for v2 catalog publication.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Iterator, Sequence + +from psycopg.types.json import Jsonb +from sqlalchemy import Connection, text + +from policyengine_api.data.v2.catalog.publication_types import ( + CatalogPublicationError, +) +from policyengine_api.data.v2.catalog.records import ( + CountryCatalog, + NormalizedCatalog, + iter_batches, +) + + +COPY_BATCH_SIZE = 10_000 + +TEMP_TABLE_STATEMENTS = ( + """ + CREATE TEMP TABLE stage_catalog_models ( + country_id text PRIMARY KEY, + id uuid NOT NULL, + name text NOT NULL, + description text, + version_id uuid NOT NULL, + version text NOT NULL, + version_description text, + current_law_id integer NOT NULL, + metadata_time_periods jsonb NOT NULL + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_variables ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + label text, + entity text NOT NULL, + description text, + data_type text, + possible_values jsonb, + default_value jsonb, + adds jsonb, + subtracts jsonb, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_parameter_nodes ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + label text, + description text, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_parameters ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + label text, + description text, + data_type text, + unit text, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_parameter_values ( + country_id text NOT NULL, + parameter_name text NOT NULL, + id uuid NOT NULL, + value_json jsonb NOT NULL, + start_date timestamptz NOT NULL, + end_date timestamptz, + PRIMARY KEY (country_id, parameter_name, start_date) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_datasets ( + country_id text NOT NULL, + id uuid NOT NULL, + name text NOT NULL, + description text, + year integer NOT NULL, + PRIMARY KEY (country_id, name) + ) ON COMMIT DROP + """, + """ + CREATE TEMP TABLE stage_catalog_regions ( + country_id text NOT NULL, + id uuid NOT NULL, + code text NOT NULL, + label text NOT NULL, + region_type text NOT NULL, + requires_filter boolean NOT NULL, + filter_field text, + filter_value text, + filter_strategy text, + parent_code text, + state_code text, + state_name text, + default_dataset_name text NOT NULL, + PRIMARY KEY (country_id, code) + ) ON COMMIT DROP + """, +) + +COPY_COLUMNS = { + "stage_catalog_models": ( + "country_id", + "id", + "name", + "description", + "version_id", + "version", + "version_description", + "current_law_id", + "metadata_time_periods", + ), + "stage_catalog_variables": ( + "country_id", + "id", + "name", + "label", + "entity", + "description", + "data_type", + "possible_values", + "default_value", + "adds", + "subtracts", + ), + "stage_catalog_parameter_nodes": ( + "country_id", + "id", + "name", + "label", + "description", + ), + "stage_catalog_parameters": ( + "country_id", + "id", + "name", + "label", + "description", + "data_type", + "unit", + ), + "stage_catalog_parameter_values": ( + "country_id", + "parameter_name", + "id", + "value_json", + "start_date", + "end_date", + ), + "stage_catalog_datasets": ( + "country_id", + "id", + "name", + "description", + "year", + ), + "stage_catalog_regions": ( + "country_id", + "id", + "code", + "label", + "region_type", + "requires_filter", + "filter_field", + "filter_value", + "filter_strategy", + "parent_code", + "state_code", + "state_name", + "default_dataset_name", + ), +} + + +def _optional_json(value: object) -> Jsonb | None: + return None if value is None else Jsonb(value) + + +def _catalog_rows( + country: CountryCatalog, +) -> dict[str, Iterator[tuple[object, ...]]]: + dataset_names = {dataset.id: dataset.name for dataset in country.datasets} + + def model_rows() -> Iterator[tuple[object, ...]]: + yield ( + country.country_id, + country.model.id, + country.model.name, + country.model.description, + country.model_version.id, + country.model_version.version, + country.model_version.description, + country.model_version.current_law_id, + Jsonb(country.model_version.metadata_time_periods), + ) + + def variable_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.variables, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.label, + record.entity, + record.description, + record.data_type, + _optional_json(record.possible_values), + Jsonb(record.default_value), + _optional_json(record.adds), + _optional_json(record.subtracts), + ) + + def parameter_node_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches( + country.parameter_nodes, + batch_size=COPY_BATCH_SIZE, + ): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.label, + record.description, + ) + + def parameter_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.parameters, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.label, + record.description, + record.data_type, + record.unit, + ) + + def parameter_value_rows() -> Iterator[tuple[object, ...]]: + parameter_names = { + parameter.id: parameter.name for parameter in country.parameters + } + for batch in country.parameter_value_batches(batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + parameter_names[record.parameter_id], + record.id, + Jsonb(record.value_json), + record.start_date, + record.end_date, + ) + + def dataset_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.datasets, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.name, + record.description, + record.year, + ) + + def region_rows() -> Iterator[tuple[object, ...]]: + for batch in iter_batches(country.regions, batch_size=COPY_BATCH_SIZE): + for record in batch: + yield ( + country.country_id, + record.id, + record.code, + record.label, + record.region_type, + record.requires_filter, + record.filter_field, + record.filter_value, + record.filter_strategy, + record.parent_code, + record.state_code, + record.state_name, + dataset_names[record.default_dataset_id], + ) + + return { + "stage_catalog_models": model_rows(), + "stage_catalog_variables": variable_rows(), + "stage_catalog_parameter_nodes": parameter_node_rows(), + "stage_catalog_parameters": parameter_rows(), + "stage_catalog_parameter_values": parameter_value_rows(), + "stage_catalog_datasets": dataset_rows(), + "stage_catalog_regions": region_rows(), + } + + +def copy_rows( + connection: Connection, + *, + table_name: str, + columns: Sequence[str], + rows: Iterable[Sequence[object]], +) -> int: + """Write one bounded source stream through Psycopg COPY.""" + + raw_connection = connection.connection.driver_connection + statement = f"COPY {table_name} ({', '.join(columns)}) FROM STDIN" + count = 0 + with raw_connection.cursor() as cursor: + with cursor.copy(statement) as copy: + for row in rows: + copy.write_row(row) + count += 1 + return count + + +def create_staging_tables(connection: Connection) -> None: + for statement in TEMP_TABLE_STATEMENTS: + connection.execute(text(statement)) + + +def stage_catalog( + connection: Connection, + catalog: NormalizedCatalog, + *, + checkpoint: Callable[[str, Connection], None] | None = None, +) -> dict[str, int]: + observed = {table_name: 0 for table_name in COPY_COLUMNS} + for country in catalog.countries: + for table_name, rows in _catalog_rows(country).items(): + observed[table_name] += copy_rows( + connection, + table_name=table_name, + columns=COPY_COLUMNS[table_name], + rows=rows, + ) + if checkpoint is not None: + checkpoint("during_copy", connection) + expected = catalog.entity_counts() + expected_by_table = { + "stage_catalog_models": expected["models"], + "stage_catalog_variables": expected["variables"], + "stage_catalog_parameter_nodes": expected["parameter_nodes"], + "stage_catalog_parameters": expected["parameters"], + "stage_catalog_parameter_values": expected["parameter_values"], + "stage_catalog_datasets": expected["datasets"], + "stage_catalog_regions": expected["regions"], + } + if observed != expected_by_table: + raise CatalogPublicationError("COPY row counts differ from the catalog") + return observed diff --git a/policyengine_api/data/v2/catalog/publication_types.py b/policyengine_api/data/v2/catalog/publication_types.py new file mode 100644 index 000000000..b9e388fc0 --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication_types.py @@ -0,0 +1,37 @@ +"""Public result and error types for v2 catalog publication.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +class CatalogPublicationError(RuntimeError): + """Raised when publication cannot prove an atomic, complete result.""" + + +@dataclass(frozen=True, slots=True) +class PublicationEvidence: + """Non-secret facts emitted after a successful publication.""" + + policyengine_version: str + dependency_versions: tuple[tuple[str, str], ...] + entity_counts: dict[str, int] + fallback_summaries: tuple[tuple[str, str, int], ...] + elapsed_seconds: float + + def as_dict(self) -> dict[str, object]: + return { + "outcome": "ok", + "policyengine_version": self.policyengine_version, + "dependency_versions": dict(self.dependency_versions), + "entity_counts": self.entity_counts, + "fallback_summaries": [ + { + "country_id": country_id, + "region_type": region_type, + "count": count, + } + for country_id, region_type, count in self.fallback_summaries + ], + "elapsed_seconds": round(self.elapsed_seconds, 3), + } diff --git a/tests/unit/v2/test_catalog_publication.py b/tests/unit/v2/test_catalog_publication.py index b7e72ac6a..cf016f561 100644 --- a/tests/unit/v2/test_catalog_publication.py +++ b/tests/unit/v2/test_catalog_publication.py @@ -10,7 +10,11 @@ from alembic.script import ScriptDirectory import pytest -from policyengine_api.data.v2.catalog import publication +from policyengine_api.data.v2.catalog import ( + publication, + publication_reconciliation, + publication_staging, +) REPO = Path(__file__).parents[3] @@ -119,7 +123,7 @@ def test_copy_streams_rows_through_psycopg_without_an_orm_write() -> None: connection = FakeConnection() rows = ((index, f"row-{index}") for index in range(3)) - count = publication._copy_rows( + count = publication_staging.copy_rows( connection, table_name="stage_catalog_models", columns=("id", "name"), @@ -139,14 +143,19 @@ def test_copy_streams_rows_through_psycopg_without_an_orm_write() -> None: def test_publication_sql_uses_private_staging_and_set_based_inserts() -> None: assert all( "CREATE TEMP TABLE" in statement and "ON COMMIT DROP" in statement - for statement in publication.TEMP_TABLE_STATEMENTS + for statement in publication_staging.TEMP_TABLE_STATEMENTS ) assert all( - "INSERT INTO" in statement for statement in publication.SET_BASED_INSERT_SQL + "INSERT INTO" in statement + for statement in publication_reconciliation.SET_BASED_INSERT_SQL ) - assert all("SELECT" in statement for statement in publication.SET_BASED_INSERT_SQL) assert all( - "VALUES" not in statement for statement in publication.SET_BASED_INSERT_SQL + "SELECT" in statement + for statement in publication_reconciliation.SET_BASED_INSERT_SQL + ) + assert all( + "VALUES" not in statement + for statement in publication_reconciliation.SET_BASED_INSERT_SQL ) class Result: From 9140c2ae66ff4d22c8ce0222b9fdaf511b33c9fe Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 16:39:13 +0400 Subject: [PATCH 11/27] Document catalog publication lock ID --- policyengine_api/data/v2/catalog/publication.py | 1 + 1 file changed, 1 insertion(+) diff --git a/policyengine_api/data/v2/catalog/publication.py b/policyengine_api/data/v2/catalog/publication.py index e2addc599..aabc4e173 100644 --- a/policyengine_api/data/v2/catalog/publication.py +++ b/policyengine_api/data/v2/catalog/publication.py @@ -26,6 +26,7 @@ EXPECTED_ALEMBIC_REVISION = "68b4a5ae5dc5" +# Stable application-defined PostgreSQL lock ID shared by all v2 catalog publishers. PUBLICATION_ADVISORY_LOCK_KEY = 8_629_020_026_090_001 LOGGER = logging.getLogger(__name__) From f83b4b5fb7493cecb569262d3115820f6cb2dd51 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 16:57:31 +0400 Subject: [PATCH 12/27] Use SQLAlchemy for catalog publication --- .../data/v2/catalog/publication.py | 47 +- .../v2/catalog/publication_reconciliation.py | 994 +++++++++++------- .../data/v2/catalog/publication_staging.py | 332 +++--- tests/unit/v2/test_catalog_publication.py | 135 ++- 4 files changed, 873 insertions(+), 635 deletions(-) diff --git a/policyengine_api/data/v2/catalog/publication.py b/policyengine_api/data/v2/catalog/publication.py index aabc4e173..3e59306cf 100644 --- a/policyengine_api/data/v2/catalog/publication.py +++ b/policyengine_api/data/v2/catalog/publication.py @@ -6,7 +6,8 @@ import logging import time -from sqlalchemy import Connection, Engine, text +import sqlalchemy as sa +from sqlalchemy import Connection, Engine from policyengine_api.data.v2.catalog.publication_reconciliation import ( assert_canonical_value_uniqueness, @@ -23,12 +24,29 @@ PublicationEvidence, ) from policyengine_api.data.v2.catalog.records import NormalizedCatalog +from policyengine_api.data.v2.models import ( + DatasetVersion, + Report, + ReportRun, + Simulation, +) EXPECTED_ALEMBIC_REVISION = "68b4a5ae5dc5" # Stable application-defined PostgreSQL lock ID shared by all v2 catalog publishers. PUBLICATION_ADVISORY_LOCK_KEY = 8_629_020_026_090_001 +ALEMBIC_VERSION = sa.table( + "alembic_version", + sa.column("version_num", sa.String), +) +PROTECTED_TABLES = ( + DatasetVersion.__table__, + Simulation.__table__, + Report.__table__, + ReportRun.__table__, +) + LOGGER = logging.getLogger(__name__) __all__ = [ @@ -41,13 +59,10 @@ def _verify_expected_revision(connection: Connection) -> None: if connection.dialect.name != "postgresql": raise CatalogPublicationError("catalog publication requires PostgreSQL") - version_table = connection.execute( - text("SELECT to_regclass('public.alembic_version')") - ).scalar_one() - if version_table is None: + if not sa.inspect(connection).has_table(ALEMBIC_VERSION.name): raise CatalogPublicationError("the v2 Alembic revision table is absent") revisions = set( - connection.execute(text("SELECT version_num FROM alembic_version")).scalars() + connection.execute(sa.select(ALEMBIC_VERSION.c.version_num)).scalars() ) if revisions != {EXPECTED_ALEMBIC_REVISION}: raise CatalogPublicationError( @@ -57,20 +72,14 @@ def _verify_expected_revision(connection: Connection) -> None: def _acquire_publication_lock(connection: Connection) -> None: connection.execute( - text("SELECT pg_advisory_xact_lock(:lock_key)"), - {"lock_key": PUBLICATION_ADVISORY_LOCK_KEY}, + sa.select(sa.func.pg_advisory_xact_lock(PUBLICATION_ADVISORY_LOCK_KEY)) ).scalar_one() def _protected_row_counts(connection: Connection) -> tuple[int, int, int, int]: return tuple( - connection.execute(text(f"SELECT count(*) FROM {table_name}")).scalar_one() - for table_name in ( - "dataset_versions", - "simulations", - "reports", - "report_runs", - ) + connection.execute(sa.select(sa.func.count()).select_from(table)).scalar_one() + for table in PROTECTED_TABLES ) @@ -94,18 +103,18 @@ def publish_catalog( existing: set[str] = set() for country in catalog.countries: - if version_exists(connection, country.country_id): - assert_country_matches(connection, country.country_id) + if version_exists(connection, country): + assert_country_matches(connection, country) existing.add(country.country_id) for country in catalog.countries: if country.country_id not in existing: - publish_new_country(connection, country.country_id) + publish_new_country(connection, country) if checkpoint is not None: checkpoint("after_reconciliation", connection) for country in catalog.countries: - assert_country_matches(connection, country.country_id) + assert_country_matches(connection, country) assert_canonical_value_uniqueness(connection) if _protected_row_counts(connection) != before: raise CatalogPublicationError( diff --git a/policyengine_api/data/v2/catalog/publication_reconciliation.py b/policyengine_api/data/v2/catalog/publication_reconciliation.py index a6c595475..8ebac6d6e 100644 --- a/policyengine_api/data/v2/catalog/publication_reconciliation.py +++ b/policyengine_api/data/v2/catalog/publication_reconciliation.py @@ -2,435 +2,645 @@ from __future__ import annotations -from sqlalchemy import Connection, text +import sqlalchemy as sa +from sqlalchemy import Connection, Select +from sqlalchemy.dialects.postgresql import JSONB, insert as postgresql_insert +from policyengine_api.data.v2.catalog.publication_staging import ( + STAGE_CATALOG_DATASETS, + STAGE_CATALOG_MODELS, + STAGE_CATALOG_PARAMETER_NODES, + STAGE_CATALOG_PARAMETER_VALUES, + STAGE_CATALOG_PARAMETERS, + STAGE_CATALOG_REGIONS, + STAGE_CATALOG_VARIABLES, +) from policyengine_api.data.v2.catalog.publication_types import ( CatalogPublicationError, ) +from policyengine_api.data.v2.catalog.records import CountryCatalog +from policyengine_api.data.v2.models import ( + Dataset, + Parameter, + ParameterNode, + ParameterValue, + Region, + TaxBenefitModel, + TaxBenefitModelVersion, + Variable, +) -def version_exists(connection: Connection, country_id: str) -> bool: - return bool( - connection.execute( - text( - """ - SELECT EXISTS ( - SELECT 1 - FROM stage_catalog_models AS source - JOIN tax_benefit_models AS model - ON model.name = source.name - JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source.version - WHERE source.country_id = :country_id - ) - """ - ), - {"country_id": country_id}, - ).scalar_one() +DATASETS = Dataset.__table__ +MODELS = TaxBenefitModel.__table__ +MODEL_VERSIONS = TaxBenefitModelVersion.__table__ +PARAMETERS = Parameter.__table__ +PARAMETER_NODES = ParameterNode.__table__ +PARAMETER_VALUES = ParameterValue.__table__ +REGIONS = Region.__table__ +VARIABLES = Variable.__table__ + +COUNTRY_ID = sa.bindparam("country_id") + + +def version_exists(connection: Connection, country: CountryCatalog) -> bool: + statement = ( + sa.select(MODEL_VERSIONS.c.id) + .select_from( + MODEL_VERSIONS.join( + MODELS, + MODELS.c.id == MODEL_VERSIONS.c.model_id, + ) + ) + .where( + MODELS.c.name == country.model.name, + MODEL_VERSIONS.c.version == country.model_version.version, + ) + .limit(1) ) + return connection.execute(statement).first() is not None -MODEL_DIFFERENCE_SQL = """ -SELECT EXISTS ( - SELECT 1 - FROM stage_catalog_models AS source - JOIN tax_benefit_models AS model ON model.name = source.name - JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source.version - WHERE source.country_id = :country_id - AND ( - model_version.description IS DISTINCT FROM source.version_description - OR model_version.current_law_id IS DISTINCT FROM source.current_law_id - OR model_version.metadata_time_periods::jsonb - IS DISTINCT FROM source.metadata_time_periods - ) -) -""" - - -VERSIONED_DIFFERENCE_SQL = ( - """ - WITH staged AS ( - SELECT name, label, entity, description, data_type, possible_values, - default_value, adds, subtracts - FROM stage_catalog_variables - WHERE country_id = :country_id - ), actual AS ( - SELECT variable.name, variable.label, variable.entity, - variable.description, variable.data_type, - variable.possible_values::jsonb, - variable.default_value::jsonb, - variable.adds::jsonb, variable.subtracts::jsonb - FROM variables AS variable - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = variable.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT name, label, description - FROM stage_catalog_parameter_nodes - WHERE country_id = :country_id - ), actual AS ( - SELECT node.name, node.label, node.description - FROM parameter_nodes AS node - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = node.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT name, label, description, data_type, unit - FROM stage_catalog_parameters - WHERE country_id = :country_id - ), actual AS ( - SELECT parameter.name, parameter.label, parameter.description, - parameter.data_type, parameter.unit - FROM parameters AS parameter - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = parameter.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT parameter_name, value_json, start_date, end_date - FROM stage_catalog_parameter_values - WHERE country_id = :country_id - ), actual AS ( - SELECT parameter.name, parameter_value.value_json::jsonb, - parameter_value.start_date, parameter_value.end_date - FROM parameter_values AS parameter_value - JOIN parameters AS parameter ON parameter.id = parameter_value.parameter_id - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = parameter.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - AND parameter_value.policy_id IS NULL - AND parameter_value.dynamic_id IS NULL - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, -) +def _versioned_table_from(table: sa.Table) -> sa.Join: + return table.join( + MODEL_VERSIONS, + MODEL_VERSIONS.c.id == table.c.tax_benefit_model_version_id, + ).join( + MODELS, + MODELS.c.id == MODEL_VERSIONS.c.model_id, + ) -VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL = ( - """ - WITH staged AS ( - SELECT name, description, year, false AS is_output_dataset, - NULL::text AS storage_path - FROM stage_catalog_datasets - WHERE country_id = :country_id - ), actual AS ( - SELECT dataset.name, dataset.description, dataset.year, - dataset.is_output_dataset, dataset.storage_path - FROM datasets AS dataset - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = dataset.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - AND NOT dataset.is_output_dataset - AND dataset.storage_path IS NULL - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, - """ - WITH staged AS ( - SELECT code, label, region_type, requires_filter, filter_field, - filter_value, filter_strategy, parent_code, state_code, - state_name, default_dataset_name - FROM stage_catalog_regions - WHERE country_id = :country_id - ), actual AS ( - SELECT region.code, region.label, region.region_type::text, - region.requires_filter, region.filter_field, - region.filter_value, region.filter_strategy, - region.parent_code, region.state_code, region.state_name, - dataset.name - FROM regions AS region - JOIN tax_benefit_model_versions AS model_version - ON model_version.id = region.tax_benefit_model_version_id - JOIN tax_benefit_models AS model ON model.id = model_version.model_id - JOIN datasets AS dataset ON dataset.id = region.default_dataset_id - JOIN stage_catalog_models AS source - ON source.name = model.name AND source.version = model_version.version - WHERE source.country_id = :country_id - ), differences AS ( - (SELECT * FROM staged EXCEPT SELECT * FROM actual) - UNION ALL - (SELECT * FROM actual EXCEPT SELECT * FROM staged) - ) SELECT EXISTS (SELECT 1 FROM differences) - """, -) +def _comparison_pairs( + country: CountryCatalog, +) -> tuple[tuple[Select, Select], ...]: + staged_country = STAGE_CATALOG_MODELS.c.country_id == country.country_id + selected_version = sa.and_( + MODELS.c.name == country.model.name, + MODEL_VERSIONS.c.version == country.model_version.version, + ) + staged_model_version = sa.select( + STAGE_CATALOG_MODELS.c.version_description, + STAGE_CATALOG_MODELS.c.current_law_id, + STAGE_CATALOG_MODELS.c.metadata_time_periods, + ).where(staged_country) + actual_model_version = ( + sa.select( + MODEL_VERSIONS.c.description, + MODEL_VERSIONS.c.current_law_id, + sa.cast(MODEL_VERSIONS.c.metadata_time_periods, JSONB), + ) + .select_from( + MODEL_VERSIONS.join( + MODELS, + MODELS.c.id == MODEL_VERSIONS.c.model_id, + ) + ) + .where(selected_version) + ) -def assert_country_matches(connection: Connection, country_id: str) -> None: - statements = ( - MODEL_DIFFERENCE_SQL, - *VERSIONED_DIFFERENCE_SQL, - *VERSION_SCOPED_REFERENCE_DIFFERENCE_SQL, + staged_variables = sa.select( + STAGE_CATALOG_VARIABLES.c.name, + STAGE_CATALOG_VARIABLES.c.label, + STAGE_CATALOG_VARIABLES.c.entity, + STAGE_CATALOG_VARIABLES.c.description, + STAGE_CATALOG_VARIABLES.c.data_type, + STAGE_CATALOG_VARIABLES.c.possible_values, + STAGE_CATALOG_VARIABLES.c.default_value, + STAGE_CATALOG_VARIABLES.c.adds, + STAGE_CATALOG_VARIABLES.c.subtracts, + ).where(STAGE_CATALOG_VARIABLES.c.country_id == country.country_id) + actual_variables = ( + sa.select( + VARIABLES.c.name, + VARIABLES.c.label, + VARIABLES.c.entity, + VARIABLES.c.description, + VARIABLES.c.data_type, + sa.cast(VARIABLES.c.possible_values, JSONB), + sa.cast(VARIABLES.c.default_value, JSONB), + sa.cast(VARIABLES.c.adds, JSONB), + sa.cast(VARIABLES.c.subtracts, JSONB), + ) + .select_from(_versioned_table_from(VARIABLES)) + .where(selected_version) ) - for statement in statements: - differs = connection.execute( - text(statement), - {"country_id": country_id}, - ).scalar_one() - if differs: - raise CatalogPublicationError( - f"persisted {country_id} catalog differs from PolicyEngine.py" + + staged_parameter_nodes = sa.select( + STAGE_CATALOG_PARAMETER_NODES.c.name, + STAGE_CATALOG_PARAMETER_NODES.c.label, + STAGE_CATALOG_PARAMETER_NODES.c.description, + ).where(STAGE_CATALOG_PARAMETER_NODES.c.country_id == country.country_id) + actual_parameter_nodes = ( + sa.select( + PARAMETER_NODES.c.name, + PARAMETER_NODES.c.label, + PARAMETER_NODES.c.description, + ) + .select_from(_versioned_table_from(PARAMETER_NODES)) + .where(selected_version) + ) + + staged_parameters = sa.select( + STAGE_CATALOG_PARAMETERS.c.name, + STAGE_CATALOG_PARAMETERS.c.label, + STAGE_CATALOG_PARAMETERS.c.description, + STAGE_CATALOG_PARAMETERS.c.data_type, + STAGE_CATALOG_PARAMETERS.c.unit, + ).where(STAGE_CATALOG_PARAMETERS.c.country_id == country.country_id) + actual_parameters = ( + sa.select( + PARAMETERS.c.name, + PARAMETERS.c.label, + PARAMETERS.c.description, + PARAMETERS.c.data_type, + PARAMETERS.c.unit, + ) + .select_from(_versioned_table_from(PARAMETERS)) + .where(selected_version) + ) + + staged_parameter_values = sa.select( + STAGE_CATALOG_PARAMETER_VALUES.c.parameter_name, + STAGE_CATALOG_PARAMETER_VALUES.c.value_json, + STAGE_CATALOG_PARAMETER_VALUES.c.start_date, + STAGE_CATALOG_PARAMETER_VALUES.c.end_date, + ).where(STAGE_CATALOG_PARAMETER_VALUES.c.country_id == country.country_id) + actual_parameter_values = ( + sa.select( + PARAMETERS.c.name, + sa.cast(PARAMETER_VALUES.c.value_json, JSONB), + PARAMETER_VALUES.c.start_date, + PARAMETER_VALUES.c.end_date, + ) + .select_from( + PARAMETER_VALUES.join( + PARAMETERS, + PARAMETERS.c.id == PARAMETER_VALUES.c.parameter_id, + ) + .join( + MODEL_VERSIONS, + MODEL_VERSIONS.c.id == PARAMETERS.c.tax_benefit_model_version_id, + ) + .join(MODELS, MODELS.c.id == MODEL_VERSIONS.c.model_id) + ) + .where( + selected_version, + PARAMETER_VALUES.c.policy_id.is_(None), + PARAMETER_VALUES.c.dynamic_id.is_(None), + ) + ) + + staged_datasets = sa.select( + STAGE_CATALOG_DATASETS.c.name, + STAGE_CATALOG_DATASETS.c.description, + STAGE_CATALOG_DATASETS.c.year, + sa.literal(False), + sa.cast(sa.null(), DATASETS.c.storage_path.type), + ).where(STAGE_CATALOG_DATASETS.c.country_id == country.country_id) + actual_datasets = ( + sa.select( + DATASETS.c.name, + DATASETS.c.description, + DATASETS.c.year, + DATASETS.c.is_output_dataset, + DATASETS.c.storage_path, + ) + .select_from(_versioned_table_from(DATASETS)) + .where( + selected_version, + DATASETS.c.is_output_dataset.is_(False), + DATASETS.c.storage_path.is_(None), + ) + ) + + staged_regions = sa.select( + STAGE_CATALOG_REGIONS.c.code, + STAGE_CATALOG_REGIONS.c.label, + STAGE_CATALOG_REGIONS.c.region_type, + STAGE_CATALOG_REGIONS.c.requires_filter, + STAGE_CATALOG_REGIONS.c.filter_field, + STAGE_CATALOG_REGIONS.c.filter_value, + STAGE_CATALOG_REGIONS.c.filter_strategy, + STAGE_CATALOG_REGIONS.c.parent_code, + STAGE_CATALOG_REGIONS.c.state_code, + STAGE_CATALOG_REGIONS.c.state_name, + STAGE_CATALOG_REGIONS.c.default_dataset_name, + ).where(STAGE_CATALOG_REGIONS.c.country_id == country.country_id) + actual_regions = ( + sa.select( + REGIONS.c.code, + REGIONS.c.label, + sa.cast(REGIONS.c.region_type, sa.Text), + REGIONS.c.requires_filter, + REGIONS.c.filter_field, + REGIONS.c.filter_value, + REGIONS.c.filter_strategy, + REGIONS.c.parent_code, + REGIONS.c.state_code, + REGIONS.c.state_name, + DATASETS.c.name, + ) + .select_from( + _versioned_table_from(REGIONS).join( + DATASETS, + DATASETS.c.id == REGIONS.c.default_dataset_id, ) + ) + .where(selected_version) + ) + + return ( + (staged_model_version, actual_model_version), + (staged_variables, actual_variables), + (staged_parameter_nodes, actual_parameter_nodes), + (staged_parameters, actual_parameters), + (staged_parameter_values, actual_parameter_values), + (staged_datasets, actual_datasets), + (staged_regions, actual_regions), + ) + + +def _sets_differ( + connection: Connection, + staged: Select, + actual: Select, +) -> bool: + missing = sa.except_(staged, actual).subquery() + extra = sa.except_(actual, staged).subquery() + differences = sa.union_all( + sa.select(sa.literal(1)).select_from(missing), + sa.select(sa.literal(1)).select_from(extra), + ).limit(1) + return connection.execute(differences).first() is not None -def _reject_dataset_role_conflicts( +def assert_country_matches( connection: Connection, - country_id: str, + country: CountryCatalog, ) -> None: - conflict = connection.execute( - text( - """ - SELECT EXISTS ( - SELECT 1 - FROM stage_catalog_datasets AS source_dataset - JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_dataset.country_id - JOIN tax_benefit_models AS model - ON model.name = source_model.name - JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version - JOIN datasets AS dataset - ON dataset.tax_benefit_model_version_id = model_version.id - AND dataset.name = source_dataset.name - WHERE source_dataset.country_id = :country_id - AND ( - dataset.is_output_dataset - OR dataset.storage_path IS NOT NULL - ) + for staged, actual in _comparison_pairs(country): + if _sets_differ(connection, staged, actual): + raise CatalogPublicationError( + f"persisted {country.country_id} catalog differs from PolicyEngine.py" ) - """ - ), - {"country_id": country_id}, - ).scalar_one() - if conflict: - raise CatalogPublicationError( - f"persisted {country_id} dataset identity is not an input dataset" - ) -INSERT_MODEL_SQL = """ -INSERT INTO tax_benefit_models ( - id, created_at, updated_at, name, description +INSERT_MODEL = ( + postgresql_insert(MODELS) + .from_select( + ["id", "created_at", "updated_at", "name", "description"], + sa.select( + STAGE_CATALOG_MODELS.c.id, + sa.func.current_timestamp(), + sa.func.current_timestamp(), + STAGE_CATALOG_MODELS.c.name, + STAGE_CATALOG_MODELS.c.description, + ).where(STAGE_CATALOG_MODELS.c.country_id == COUNTRY_ID), + ) + .on_conflict_do_nothing(index_elements=[MODELS.c.name]) ) -SELECT id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, name, description -FROM stage_catalog_models -WHERE country_id = :country_id -ON CONFLICT (name) DO NOTHING -""" - -INSERT_MODEL_VERSION_SQL = """ -INSERT INTO tax_benefit_model_versions ( - id, created_at, model_id, version, description, current_law_id, - metadata_time_periods + +INSERT_MODEL_VERSION = sa.insert(MODEL_VERSIONS).from_select( + [ + "id", + "created_at", + "model_id", + "version", + "description", + "current_law_id", + "metadata_time_periods", + ], + sa.select( + STAGE_CATALOG_MODELS.c.version_id, + sa.func.current_timestamp(), + MODELS.c.id, + STAGE_CATALOG_MODELS.c.version, + STAGE_CATALOG_MODELS.c.version_description, + STAGE_CATALOG_MODELS.c.current_law_id, + sa.cast( + STAGE_CATALOG_MODELS.c.metadata_time_periods, + MODEL_VERSIONS.c.metadata_time_periods.type, + ), + ) + .select_from( + STAGE_CATALOG_MODELS.join( + MODELS, + MODELS.c.name == STAGE_CATALOG_MODELS.c.name, + ) + ) + .where(STAGE_CATALOG_MODELS.c.country_id == COUNTRY_ID), ) -SELECT source.version_id, CURRENT_TIMESTAMP, model.id, - source.version, source.version_description, source.current_law_id, - source.metadata_time_periods::json -FROM stage_catalog_models AS source -JOIN tax_benefit_models AS model ON model.name = source.name -WHERE source.country_id = :country_id -""" - -INSERT_DATASETS_SQL = """ -INSERT INTO datasets ( - id, created_at, updated_at, name, description, storage_path, year, - is_output_dataset, tax_benefit_model_version_id + +INSERT_DATASETS = sa.insert(DATASETS).from_select( + [ + "id", + "created_at", + "updated_at", + "name", + "description", + "storage_path", + "year", + "is_output_dataset", + "tax_benefit_model_version_id", + ], + sa.select( + STAGE_CATALOG_DATASETS.c.id, + sa.func.current_timestamp(), + sa.func.current_timestamp(), + STAGE_CATALOG_DATASETS.c.name, + STAGE_CATALOG_DATASETS.c.description, + sa.cast(sa.null(), DATASETS.c.storage_path.type), + STAGE_CATALOG_DATASETS.c.year, + sa.literal(False), + MODEL_VERSIONS.c.id, + ) + .select_from( + STAGE_CATALOG_DATASETS.join( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_MODELS.c.country_id == STAGE_CATALOG_DATASETS.c.country_id, + ) + .join(MODELS, MODELS.c.name == STAGE_CATALOG_MODELS.c.name) + .join( + MODEL_VERSIONS, + sa.and_( + MODEL_VERSIONS.c.model_id == MODELS.c.id, + MODEL_VERSIONS.c.version == STAGE_CATALOG_MODELS.c.version, + ), + ) + ) + .where(STAGE_CATALOG_DATASETS.c.country_id == COUNTRY_ID), ) -SELECT source_dataset.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, - source_dataset.name, source_dataset.description, NULL, - source_dataset.year, false, model_version.id -FROM stage_catalog_datasets AS source_dataset -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_dataset.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_dataset.country_id = :country_id -""" - -INSERT_VARIABLES_SQL = """ -INSERT INTO variables ( - id, created_at, name, label, entity, description, data_type, - possible_values, default_value, adds, subtracts, - tax_benefit_model_version_id + +INSERT_VARIABLES = sa.insert(VARIABLES).from_select( + [ + "id", + "created_at", + "name", + "label", + "entity", + "description", + "data_type", + "possible_values", + "default_value", + "adds", + "subtracts", + "tax_benefit_model_version_id", + ], + sa.select( + STAGE_CATALOG_VARIABLES.c.id, + sa.func.current_timestamp(), + STAGE_CATALOG_VARIABLES.c.name, + STAGE_CATALOG_VARIABLES.c.label, + STAGE_CATALOG_VARIABLES.c.entity, + STAGE_CATALOG_VARIABLES.c.description, + STAGE_CATALOG_VARIABLES.c.data_type, + sa.cast( + STAGE_CATALOG_VARIABLES.c.possible_values, + VARIABLES.c.possible_values.type, + ), + sa.cast( + STAGE_CATALOG_VARIABLES.c.default_value, + VARIABLES.c.default_value.type, + ), + sa.cast(STAGE_CATALOG_VARIABLES.c.adds, VARIABLES.c.adds.type), + sa.cast( + STAGE_CATALOG_VARIABLES.c.subtracts, + VARIABLES.c.subtracts.type, + ), + MODEL_VERSIONS.c.id, + ) + .select_from( + STAGE_CATALOG_VARIABLES.join( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_MODELS.c.country_id == STAGE_CATALOG_VARIABLES.c.country_id, + ) + .join(MODELS, MODELS.c.name == STAGE_CATALOG_MODELS.c.name) + .join( + MODEL_VERSIONS, + sa.and_( + MODEL_VERSIONS.c.model_id == MODELS.c.id, + MODEL_VERSIONS.c.version == STAGE_CATALOG_MODELS.c.version, + ), + ) + ) + .where(STAGE_CATALOG_VARIABLES.c.country_id == COUNTRY_ID), ) -SELECT source_variable.id, CURRENT_TIMESTAMP, source_variable.name, - source_variable.label, source_variable.entity, - source_variable.description, source_variable.data_type, - source_variable.possible_values::json, - source_variable.default_value::json, - source_variable.adds::json, source_variable.subtracts::json, - model_version.id -FROM stage_catalog_variables AS source_variable -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_variable.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_variable.country_id = :country_id -""" - -INSERT_PARAMETER_NODES_SQL = """ -INSERT INTO parameter_nodes ( - id, created_at, name, label, description, tax_benefit_model_version_id + +INSERT_PARAMETER_NODES = sa.insert(PARAMETER_NODES).from_select( + [ + "id", + "created_at", + "name", + "label", + "description", + "tax_benefit_model_version_id", + ], + sa.select( + STAGE_CATALOG_PARAMETER_NODES.c.id, + sa.func.current_timestamp(), + STAGE_CATALOG_PARAMETER_NODES.c.name, + STAGE_CATALOG_PARAMETER_NODES.c.label, + STAGE_CATALOG_PARAMETER_NODES.c.description, + MODEL_VERSIONS.c.id, + ) + .select_from( + STAGE_CATALOG_PARAMETER_NODES.join( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_MODELS.c.country_id + == STAGE_CATALOG_PARAMETER_NODES.c.country_id, + ) + .join(MODELS, MODELS.c.name == STAGE_CATALOG_MODELS.c.name) + .join( + MODEL_VERSIONS, + sa.and_( + MODEL_VERSIONS.c.model_id == MODELS.c.id, + MODEL_VERSIONS.c.version == STAGE_CATALOG_MODELS.c.version, + ), + ) + ) + .where(STAGE_CATALOG_PARAMETER_NODES.c.country_id == COUNTRY_ID), ) -SELECT source_node.id, CURRENT_TIMESTAMP, source_node.name, - source_node.label, source_node.description, model_version.id -FROM stage_catalog_parameter_nodes AS source_node -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_node.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_node.country_id = :country_id -""" - -INSERT_PARAMETERS_SQL = """ -INSERT INTO parameters ( - id, created_at, name, label, description, data_type, unit, - tax_benefit_model_version_id + +INSERT_PARAMETERS = sa.insert(PARAMETERS).from_select( + [ + "id", + "created_at", + "name", + "label", + "description", + "data_type", + "unit", + "tax_benefit_model_version_id", + ], + sa.select( + STAGE_CATALOG_PARAMETERS.c.id, + sa.func.current_timestamp(), + STAGE_CATALOG_PARAMETERS.c.name, + STAGE_CATALOG_PARAMETERS.c.label, + STAGE_CATALOG_PARAMETERS.c.description, + STAGE_CATALOG_PARAMETERS.c.data_type, + STAGE_CATALOG_PARAMETERS.c.unit, + MODEL_VERSIONS.c.id, + ) + .select_from( + STAGE_CATALOG_PARAMETERS.join( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_MODELS.c.country_id == STAGE_CATALOG_PARAMETERS.c.country_id, + ) + .join(MODELS, MODELS.c.name == STAGE_CATALOG_MODELS.c.name) + .join( + MODEL_VERSIONS, + sa.and_( + MODEL_VERSIONS.c.model_id == MODELS.c.id, + MODEL_VERSIONS.c.version == STAGE_CATALOG_MODELS.c.version, + ), + ) + ) + .where(STAGE_CATALOG_PARAMETERS.c.country_id == COUNTRY_ID), ) -SELECT source_parameter.id, CURRENT_TIMESTAMP, source_parameter.name, - source_parameter.label, source_parameter.description, - source_parameter.data_type, source_parameter.unit, model_version.id -FROM stage_catalog_parameters AS source_parameter -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_parameter.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -WHERE source_parameter.country_id = :country_id -""" - -INSERT_PARAMETER_VALUES_SQL = """ -INSERT INTO parameter_values ( - id, created_at, parameter_id, value_json, start_date, end_date, - policy_id, dynamic_id + +INSERT_PARAMETER_VALUES = sa.insert(PARAMETER_VALUES).from_select( + [ + "id", + "created_at", + "parameter_id", + "value_json", + "start_date", + "end_date", + "policy_id", + "dynamic_id", + ], + sa.select( + STAGE_CATALOG_PARAMETER_VALUES.c.id, + sa.func.current_timestamp(), + PARAMETERS.c.id, + sa.cast( + STAGE_CATALOG_PARAMETER_VALUES.c.value_json, + PARAMETER_VALUES.c.value_json.type, + ), + STAGE_CATALOG_PARAMETER_VALUES.c.start_date, + STAGE_CATALOG_PARAMETER_VALUES.c.end_date, + sa.cast(sa.null(), PARAMETER_VALUES.c.policy_id.type), + sa.cast(sa.null(), PARAMETER_VALUES.c.dynamic_id.type), + ) + .select_from( + STAGE_CATALOG_PARAMETER_VALUES.join( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_MODELS.c.country_id + == STAGE_CATALOG_PARAMETER_VALUES.c.country_id, + ) + .join(MODELS, MODELS.c.name == STAGE_CATALOG_MODELS.c.name) + .join( + MODEL_VERSIONS, + sa.and_( + MODEL_VERSIONS.c.model_id == MODELS.c.id, + MODEL_VERSIONS.c.version == STAGE_CATALOG_MODELS.c.version, + ), + ) + .join( + PARAMETERS, + sa.and_( + PARAMETERS.c.tax_benefit_model_version_id == MODEL_VERSIONS.c.id, + PARAMETERS.c.name == STAGE_CATALOG_PARAMETER_VALUES.c.parameter_name, + ), + ) + ) + .where(STAGE_CATALOG_PARAMETER_VALUES.c.country_id == COUNTRY_ID), ) -SELECT source_value.id, CURRENT_TIMESTAMP, parameter.id, - source_value.value_json::json, source_value.start_date, - source_value.end_date, NULL, NULL -FROM stage_catalog_parameter_values AS source_value -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_value.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -JOIN parameters AS parameter - ON parameter.tax_benefit_model_version_id = model_version.id - AND parameter.name = source_value.parameter_name -WHERE source_value.country_id = :country_id -""" - -INSERT_REGIONS_SQL = """ -INSERT INTO regions ( - id, created_at, updated_at, code, label, region_type, requires_filter, - filter_field, filter_value, filter_strategy, parent_code, state_code, - state_name, tax_benefit_model_version_id, default_dataset_id + +INSERT_REGIONS = sa.insert(REGIONS).from_select( + [ + "id", + "created_at", + "updated_at", + "code", + "label", + "region_type", + "requires_filter", + "filter_field", + "filter_value", + "filter_strategy", + "parent_code", + "state_code", + "state_name", + "tax_benefit_model_version_id", + "default_dataset_id", + ], + sa.select( + STAGE_CATALOG_REGIONS.c.id, + sa.func.current_timestamp(), + sa.func.current_timestamp(), + STAGE_CATALOG_REGIONS.c.code, + STAGE_CATALOG_REGIONS.c.label, + sa.cast(STAGE_CATALOG_REGIONS.c.region_type, REGIONS.c.region_type.type), + STAGE_CATALOG_REGIONS.c.requires_filter, + STAGE_CATALOG_REGIONS.c.filter_field, + STAGE_CATALOG_REGIONS.c.filter_value, + STAGE_CATALOG_REGIONS.c.filter_strategy, + STAGE_CATALOG_REGIONS.c.parent_code, + STAGE_CATALOG_REGIONS.c.state_code, + STAGE_CATALOG_REGIONS.c.state_name, + MODEL_VERSIONS.c.id, + DATASETS.c.id, + ) + .select_from( + STAGE_CATALOG_REGIONS.join( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_MODELS.c.country_id == STAGE_CATALOG_REGIONS.c.country_id, + ) + .join(MODELS, MODELS.c.name == STAGE_CATALOG_MODELS.c.name) + .join( + MODEL_VERSIONS, + sa.and_( + MODEL_VERSIONS.c.model_id == MODELS.c.id, + MODEL_VERSIONS.c.version == STAGE_CATALOG_MODELS.c.version, + ), + ) + .join( + DATASETS, + sa.and_( + DATASETS.c.tax_benefit_model_version_id == MODEL_VERSIONS.c.id, + DATASETS.c.name == STAGE_CATALOG_REGIONS.c.default_dataset_name, + ), + ) + ) + .where(STAGE_CATALOG_REGIONS.c.country_id == COUNTRY_ID), ) -SELECT source_region.id, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP, - source_region.code, source_region.label, - source_region.region_type::v2_region_type, - source_region.requires_filter, source_region.filter_field, - source_region.filter_value, source_region.filter_strategy, - source_region.parent_code, source_region.state_code, - source_region.state_name, model_version.id, dataset.id -FROM stage_catalog_regions AS source_region -JOIN stage_catalog_models AS source_model - ON source_model.country_id = source_region.country_id -JOIN tax_benefit_models AS model ON model.name = source_model.name -JOIN tax_benefit_model_versions AS model_version - ON model_version.model_id = model.id - AND model_version.version = source_model.version -JOIN datasets AS dataset - ON dataset.tax_benefit_model_version_id = model_version.id - AND dataset.name = source_region.default_dataset_name -WHERE source_region.country_id = :country_id -""" - - -SET_BASED_INSERT_SQL = ( - INSERT_MODEL_SQL, - INSERT_MODEL_VERSION_SQL, - INSERT_DATASETS_SQL, - INSERT_VARIABLES_SQL, - INSERT_PARAMETER_NODES_SQL, - INSERT_PARAMETERS_SQL, - INSERT_PARAMETER_VALUES_SQL, - INSERT_REGIONS_SQL, + +SET_BASED_INSERT_STATEMENTS = ( + INSERT_MODEL, + INSERT_MODEL_VERSION, + INSERT_DATASETS, + INSERT_VARIABLES, + INSERT_PARAMETER_NODES, + INSERT_PARAMETERS, + INSERT_PARAMETER_VALUES, + INSERT_REGIONS, ) -def publish_new_country(connection: Connection, country_id: str) -> None: - _reject_dataset_role_conflicts(connection, country_id) - for statement in SET_BASED_INSERT_SQL: - connection.execute(text(statement), {"country_id": country_id}) +def publish_new_country(connection: Connection, country: CountryCatalog) -> None: + for statement in SET_BASED_INSERT_STATEMENTS: + connection.execute(statement, {"country_id": country.country_id}) def assert_canonical_value_uniqueness(connection: Connection) -> None: - duplicates = connection.execute( - text( - """ - SELECT EXISTS ( - SELECT 1 - FROM parameter_values - WHERE policy_id IS NULL AND dynamic_id IS NULL - GROUP BY parameter_id, start_date - HAVING count(*) > 1 - ) - """ + duplicates = ( + sa.select(PARAMETER_VALUES.c.parameter_id) + .where( + PARAMETER_VALUES.c.policy_id.is_(None), + PARAMETER_VALUES.c.dynamic_id.is_(None), + ) + .group_by( + PARAMETER_VALUES.c.parameter_id, + PARAMETER_VALUES.c.start_date, ) - ).scalar_one() - if duplicates: + .having(sa.func.count() > 1) + .limit(1) + ) + if connection.execute(duplicates).first() is not None: raise CatalogPublicationError( "canonical parameter-value uniqueness validation failed" ) diff --git a/policyengine_api/data/v2/catalog/publication_staging.py b/policyengine_api/data/v2/catalog/publication_staging.py index 747023f1c..325d3ac3c 100644 --- a/policyengine_api/data/v2/catalog/publication_staging.py +++ b/policyengine_api/data/v2/catalog/publication_staging.py @@ -4,8 +4,11 @@ from collections.abc import Callable, Iterable, Iterator, Sequence +from psycopg import sql from psycopg.types.json import Jsonb -from sqlalchemy import Connection, text +import sqlalchemy as sa +from sqlalchemy import Connection, MetaData, Table +from sqlalchemy.dialects.postgresql import JSONB from policyengine_api.data.v2.catalog.publication_types import ( CatalogPublicationError, @@ -19,171 +22,126 @@ COPY_BATCH_SIZE = 10_000 -TEMP_TABLE_STATEMENTS = ( - """ - CREATE TEMP TABLE stage_catalog_models ( - country_id text PRIMARY KEY, - id uuid NOT NULL, - name text NOT NULL, - description text, - version_id uuid NOT NULL, - version text NOT NULL, - version_description text, - current_law_id integer NOT NULL, - metadata_time_periods jsonb NOT NULL - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_variables ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - label text, - entity text NOT NULL, - description text, - data_type text, - possible_values jsonb, - default_value jsonb, - adds jsonb, - subtracts jsonb, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_parameter_nodes ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - label text, - description text, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_parameters ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - label text, - description text, - data_type text, - unit text, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_parameter_values ( - country_id text NOT NULL, - parameter_name text NOT NULL, - id uuid NOT NULL, - value_json jsonb NOT NULL, - start_date timestamptz NOT NULL, - end_date timestamptz, - PRIMARY KEY (country_id, parameter_name, start_date) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_datasets ( - country_id text NOT NULL, - id uuid NOT NULL, - name text NOT NULL, - description text, - year integer NOT NULL, - PRIMARY KEY (country_id, name) - ) ON COMMIT DROP - """, - """ - CREATE TEMP TABLE stage_catalog_regions ( - country_id text NOT NULL, - id uuid NOT NULL, - code text NOT NULL, - label text NOT NULL, - region_type text NOT NULL, - requires_filter boolean NOT NULL, - filter_field text, - filter_value text, - filter_strategy text, - parent_code text, - state_code text, - state_name text, - default_dataset_name text NOT NULL, - PRIMARY KEY (country_id, code) - ) ON COMMIT DROP - """, +STAGING_METADATA = MetaData() + +STAGE_CATALOG_MODELS = Table( + "stage_catalog_models", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("name", sa.Text, nullable=False), + sa.Column("description", sa.Text), + sa.Column("version_id", sa.Uuid, nullable=False), + sa.Column("version", sa.Text, nullable=False), + sa.Column("version_description", sa.Text), + sa.Column("current_law_id", sa.Integer, nullable=False), + sa.Column("metadata_time_periods", JSONB, nullable=False), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", ) -COPY_COLUMNS = { - "stage_catalog_models": ( - "country_id", - "id", - "name", - "description", - "version_id", - "version", - "version_description", - "current_law_id", - "metadata_time_periods", - ), - "stage_catalog_variables": ( - "country_id", - "id", - "name", - "label", - "entity", - "description", - "data_type", - "possible_values", - "default_value", - "adds", - "subtracts", - ), - "stage_catalog_parameter_nodes": ( - "country_id", - "id", - "name", - "label", - "description", - ), - "stage_catalog_parameters": ( - "country_id", - "id", - "name", - "label", - "description", - "data_type", - "unit", - ), - "stage_catalog_parameter_values": ( - "country_id", - "parameter_name", - "id", - "value_json", +STAGE_CATALOG_VARIABLES = Table( + "stage_catalog_variables", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("name", sa.Text, primary_key=True), + sa.Column("label", sa.Text), + sa.Column("entity", sa.Text, nullable=False), + sa.Column("description", sa.Text), + sa.Column("data_type", sa.Text), + sa.Column("possible_values", JSONB), + sa.Column("default_value", JSONB, nullable=False), + sa.Column("adds", JSONB), + sa.Column("subtracts", JSONB), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", +) + +STAGE_CATALOG_PARAMETER_NODES = Table( + "stage_catalog_parameter_nodes", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("name", sa.Text, primary_key=True), + sa.Column("label", sa.Text), + sa.Column("description", sa.Text), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", +) + +STAGE_CATALOG_PARAMETERS = Table( + "stage_catalog_parameters", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("name", sa.Text, primary_key=True), + sa.Column("label", sa.Text), + sa.Column("description", sa.Text), + sa.Column("data_type", sa.Text), + sa.Column("unit", sa.Text), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", +) + +STAGE_CATALOG_PARAMETER_VALUES = Table( + "stage_catalog_parameter_values", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("parameter_name", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("value_json", JSONB, nullable=False), + sa.Column( "start_date", - "end_date", - ), - "stage_catalog_datasets": ( - "country_id", - "id", - "name", - "description", - "year", - ), - "stage_catalog_regions": ( - "country_id", - "id", - "code", - "label", - "region_type", - "requires_filter", - "filter_field", - "filter_value", - "filter_strategy", - "parent_code", - "state_code", - "state_name", - "default_dataset_name", + sa.DateTime(timezone=True), + primary_key=True, ), -} + sa.Column("end_date", sa.DateTime(timezone=True)), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", +) + +STAGE_CATALOG_DATASETS = Table( + "stage_catalog_datasets", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("name", sa.Text, primary_key=True), + sa.Column("description", sa.Text), + sa.Column("year", sa.Integer, nullable=False), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", +) + +STAGE_CATALOG_REGIONS = Table( + "stage_catalog_regions", + STAGING_METADATA, + sa.Column("country_id", sa.Text, primary_key=True), + sa.Column("id", sa.Uuid, nullable=False), + sa.Column("code", sa.Text, primary_key=True), + sa.Column("label", sa.Text, nullable=False), + sa.Column("region_type", sa.Text, nullable=False), + sa.Column("requires_filter", sa.Boolean, nullable=False), + sa.Column("filter_field", sa.Text), + sa.Column("filter_value", sa.Text), + sa.Column("filter_strategy", sa.Text), + sa.Column("parent_code", sa.Text), + sa.Column("state_code", sa.Text), + sa.Column("state_name", sa.Text), + sa.Column("default_dataset_name", sa.Text, nullable=False), + prefixes=["TEMPORARY"], + postgresql_on_commit="DROP", +) + +STAGING_TABLES = ( + STAGE_CATALOG_MODELS, + STAGE_CATALOG_VARIABLES, + STAGE_CATALOG_PARAMETER_NODES, + STAGE_CATALOG_PARAMETERS, + STAGE_CATALOG_PARAMETER_VALUES, + STAGE_CATALOG_DATASETS, + STAGE_CATALOG_REGIONS, +) def _optional_json(value: object) -> Jsonb | None: @@ -192,7 +150,7 @@ def _optional_json(value: object) -> Jsonb | None: def _catalog_rows( country: CountryCatalog, -) -> dict[str, Iterator[tuple[object, ...]]]: +) -> dict[Table, Iterator[tuple[object, ...]]]: dataset_names = {dataset.id: dataset.name for dataset in country.datasets} def model_rows() -> Iterator[tuple[object, ...]]: @@ -298,27 +256,29 @@ def region_rows() -> Iterator[tuple[object, ...]]: ) return { - "stage_catalog_models": model_rows(), - "stage_catalog_variables": variable_rows(), - "stage_catalog_parameter_nodes": parameter_node_rows(), - "stage_catalog_parameters": parameter_rows(), - "stage_catalog_parameter_values": parameter_value_rows(), - "stage_catalog_datasets": dataset_rows(), - "stage_catalog_regions": region_rows(), + STAGE_CATALOG_MODELS: model_rows(), + STAGE_CATALOG_VARIABLES: variable_rows(), + STAGE_CATALOG_PARAMETER_NODES: parameter_node_rows(), + STAGE_CATALOG_PARAMETERS: parameter_rows(), + STAGE_CATALOG_PARAMETER_VALUES: parameter_value_rows(), + STAGE_CATALOG_DATASETS: dataset_rows(), + STAGE_CATALOG_REGIONS: region_rows(), } def copy_rows( connection: Connection, *, - table_name: str, - columns: Sequence[str], + table: Table, rows: Iterable[Sequence[object]], ) -> int: """Write one bounded source stream through Psycopg COPY.""" raw_connection = connection.connection.driver_connection - statement = f"COPY {table_name} ({', '.join(columns)}) FROM STDIN" + statement = sql.SQL("COPY {} ({}) FROM STDIN").format( + sql.Identifier(table.name), + sql.SQL(", ").join(sql.Identifier(column.name) for column in table.columns), + ) count = 0 with raw_connection.cursor() as cursor: with cursor.copy(statement) as copy: @@ -329,8 +289,7 @@ def copy_rows( def create_staging_tables(connection: Connection) -> None: - for statement in TEMP_TABLE_STATEMENTS: - connection.execute(text(statement)) + STAGING_METADATA.create_all(connection, checkfirst=False) def stage_catalog( @@ -339,26 +298,25 @@ def stage_catalog( *, checkpoint: Callable[[str, Connection], None] | None = None, ) -> dict[str, int]: - observed = {table_name: 0 for table_name in COPY_COLUMNS} + observed = {table.name: 0 for table in STAGING_TABLES} for country in catalog.countries: - for table_name, rows in _catalog_rows(country).items(): - observed[table_name] += copy_rows( + for table, rows in _catalog_rows(country).items(): + observed[table.name] += copy_rows( connection, - table_name=table_name, - columns=COPY_COLUMNS[table_name], + table=table, rows=rows, ) if checkpoint is not None: checkpoint("during_copy", connection) expected = catalog.entity_counts() expected_by_table = { - "stage_catalog_models": expected["models"], - "stage_catalog_variables": expected["variables"], - "stage_catalog_parameter_nodes": expected["parameter_nodes"], - "stage_catalog_parameters": expected["parameters"], - "stage_catalog_parameter_values": expected["parameter_values"], - "stage_catalog_datasets": expected["datasets"], - "stage_catalog_regions": expected["regions"], + STAGE_CATALOG_MODELS.name: expected["models"], + STAGE_CATALOG_VARIABLES.name: expected["variables"], + STAGE_CATALOG_PARAMETER_NODES.name: expected["parameter_nodes"], + STAGE_CATALOG_PARAMETERS.name: expected["parameters"], + STAGE_CATALOG_PARAMETER_VALUES.name: expected["parameter_values"], + STAGE_CATALOG_DATASETS.name: expected["datasets"], + STAGE_CATALOG_REGIONS.name: expected["regions"], } if observed != expected_by_table: raise CatalogPublicationError("COPY row counts differ from the catalog") diff --git a/tests/unit/v2/test_catalog_publication.py b/tests/unit/v2/test_catalog_publication.py index cf016f561..6d09cafb5 100644 --- a/tests/unit/v2/test_catalog_publication.py +++ b/tests/unit/v2/test_catalog_publication.py @@ -9,12 +9,27 @@ from alembic.config import Config from alembic.script import ScriptDirectory import pytest +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql +from sqlalchemy.schema import CreateTable +from sqlalchemy.sql.elements import TextClause from policyengine_api.data.v2.catalog import ( publication, publication_reconciliation, publication_staging, ) +from policyengine_api.data.v2.models import ( + Dataset, + Parameter, + ParameterNode, + ParameterValue, + Region, + TaxBenefitModel, + TaxBenefitModelVersion, + Variable, +) +from tests.fixtures.v2_catalog import normalized_catalog REPO = Path(__file__).parents[3] @@ -87,33 +102,40 @@ def scalars(self): class _RevisionConnection: - def __init__(self, *, dialect: str, version_table=None, revisions=()): + def __init__(self, *, dialect: str, revisions=()): self.dialect = type("Dialect", (), {"name": dialect})() - self.results = iter( - ( - _ScalarResult(value=version_table), - _ScalarResult(values=revisions), - ) - ) + self.result = _ScalarResult(values=revisions) def execute(self, _statement): - return next(self.results) + return self.result + + +class _Inspector: + def __init__(self, table_exists: bool): + self.table_exists = table_exists + + def has_table(self, table_name: str) -> bool: + assert table_name == "alembic_version" + return self.table_exists def test_revision_check_rejects_non_postgres_missing_and_wrong_revisions() -> None: with pytest.raises(publication.CatalogPublicationError, match="PostgreSQL"): publication._verify_expected_revision(_RevisionConnection(dialect="sqlite")) - with pytest.raises(publication.CatalogPublicationError, match="table is absent"): - publication._verify_expected_revision( - _RevisionConnection(dialect="postgresql", version_table=None) - ) + with ( + patch.object(publication.sa, "inspect", return_value=_Inspector(False)), + pytest.raises(publication.CatalogPublicationError, match="table is absent"), + ): + publication._verify_expected_revision(_RevisionConnection(dialect="postgresql")) - with pytest.raises(publication.CatalogPublicationError, match="expected"): + with ( + patch.object(publication.sa, "inspect", return_value=_Inspector(True)), + pytest.raises(publication.CatalogPublicationError, match="expected"), + ): publication._verify_expected_revision( _RevisionConnection( dialect="postgresql", - version_table="alembic_version", revisions=("wrong-revision",), ) ) @@ -122,17 +144,24 @@ def test_revision_check_rejects_non_postgres_missing_and_wrong_revisions() -> No def test_copy_streams_rows_through_psycopg_without_an_orm_write() -> None: connection = FakeConnection() rows = ((index, f"row-{index}") for index in range(3)) + copy_table = sa.Table( + "stage_catalog_models", + sa.MetaData(), + sa.Column("id", sa.Integer), + sa.Column("name", sa.Text), + ) count = publication_staging.copy_rows( connection, - table_name="stage_catalog_models", - columns=("id", "name"), + table=copy_table, rows=rows, ) cursor = connection.connection.driver_connection.selected_cursor assert count == 3 - assert cursor.statement == ("COPY stage_catalog_models (id, name) FROM STDIN") + assert cursor.statement.as_string() == ( + 'COPY "stage_catalog_models" ("id", "name") FROM STDIN' + ) assert cursor.copy_operation.rows == [ (0, "row-0"), (1, "row-1"), @@ -140,22 +169,51 @@ def test_copy_streams_rows_through_psycopg_without_an_orm_write() -> None: ] -def test_publication_sql_uses_private_staging_and_set_based_inserts() -> None: - assert all( - "CREATE TEMP TABLE" in statement and "ON COMMIT DROP" in statement - for statement in publication_staging.TEMP_TABLE_STATEMENTS +def test_publication_uses_mapped_tables_and_sqlalchemy_statements() -> None: + dialect = postgresql.dialect() + assert ( + publication_reconciliation.DATASETS, + publication_reconciliation.MODELS, + publication_reconciliation.MODEL_VERSIONS, + publication_reconciliation.PARAMETERS, + publication_reconciliation.PARAMETER_NODES, + publication_reconciliation.PARAMETER_VALUES, + publication_reconciliation.REGIONS, + publication_reconciliation.VARIABLES, + ) == ( + Dataset.__table__, + TaxBenefitModel.__table__, + TaxBenefitModelVersion.__table__, + Parameter.__table__, + ParameterNode.__table__, + ParameterValue.__table__, + Region.__table__, + Variable.__table__, ) - assert all( - "INSERT INTO" in statement - for statement in publication_reconciliation.SET_BASED_INSERT_SQL + staging_ddl = tuple( + str(CreateTable(table).compile(dialect=dialect)) + for table in publication_staging.STAGING_TABLES ) assert all( - "SELECT" in statement - for statement in publication_reconciliation.SET_BASED_INSERT_SQL + "CREATE TEMPORARY TABLE" in statement and "ON COMMIT DROP" in statement + for statement in staging_ddl + ) + + inserts = publication_reconciliation.SET_BASED_INSERT_STATEMENTS + assert all(not isinstance(statement, TextClause) for statement in inserts) + compiled_inserts = tuple( + str(statement.compile(dialect=dialect)) for statement in inserts ) + assert all("INSERT INTO" in statement for statement in compiled_inserts) + assert all("SELECT" in statement for statement in compiled_inserts) + assert all("VALUES" not in statement for statement in compiled_inserts) + + catalog = normalized_catalog() + comparisons = publication_reconciliation._comparison_pairs(catalog.country("us")) assert all( - "VALUES" not in statement - for statement in publication_reconciliation.SET_BASED_INSERT_SQL + isinstance(statement, sa.sql.Select) and not isinstance(statement, TextClause) + for pair in comparisons + for statement in pair ) class Result: @@ -163,20 +221,23 @@ def scalar_one(self): return None class Connection: - statement = "" - parameters = {} + statement = None - def execute(self, statement, parameters): - self.statement = str(statement) - self.parameters = parameters + def execute(self, statement): + self.statement = statement return Result() connection = Connection() publication._acquire_publication_lock(connection) - assert "pg_advisory_xact_lock" in connection.statement - assert connection.parameters == { - "lock_key": publication.PUBLICATION_ADVISORY_LOCK_KEY - } + assert not isinstance(connection.statement, TextClause) + compiled_lock = str( + connection.statement.compile( + dialect=dialect, + compile_kwargs={"literal_binds": True}, + ) + ) + assert "pg_advisory_xact_lock" in compiled_lock + assert str(publication.PUBLICATION_ADVISORY_LOCK_KEY) in compiled_lock def test_completion_evidence_contains_only_reviewed_non_secret_fields() -> None: From 854299c73c7f1ce34959d8bf3ee607896e4a8ec0 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:30:34 +0400 Subject: [PATCH 13/27] Define split v2 metadata contracts --- policyengine_api/data/v2/catalog/query.py | 61 +++++++ policyengine_api/data/v2/catalog/schemas.py | 172 +++++++++++++++++++- 2 files changed, 232 insertions(+), 1 deletion(-) diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index d19ddd05b..953600219 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections import defaultdict +from dataclasses import dataclass from packaging.version import InvalidVersion, Version from sqlalchemy.exc import SQLAlchemyError @@ -55,6 +56,16 @@ class MetadataCatalogVersionNotFoundError(LookupError): """Raised when an explicitly selected catalog version is absent.""" +@dataclass(frozen=True) +class SelectedCatalog: + """One country catalog selected by its canonical PolicyEngine.py version.""" + + country_id: str + policyengine_version: str + model: TaxBenefitModel + model_version: TaxBenefitModelVersion + + def validate_policyengine_version(value: str) -> str: """Return one bounded canonical PEP 440 version string.""" @@ -93,6 +104,56 @@ def close(self) -> None: self._session.close() + def select_catalog( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> SelectedCatalog: + """Select exactly one initialized country catalog.""" + + if country_id not in SUPPORTED_PREVIEW_COUNTRIES: + raise UnsupportedPreviewCountryError(country_id) + explicit_version = policyengine_version is not None + selected_version = ( + validate_policyengine_version(policyengine_version) + if explicit_version + else self._running_policyengine_version + ) + try: + row = self._session.exec( + select(TaxBenefitModel, TaxBenefitModelVersion) + .join( + TaxBenefitModelVersion, + TaxBenefitModelVersion.model_id == TaxBenefitModel.id, + ) + .where( + TaxBenefitModel.name == f"policyengine-{country_id}", + TaxBenefitModelVersion.version == selected_version, + ) + ).one_or_none() + except SQLAlchemyError as error: + raise MetadataCatalogUnavailableError( + "the v2 metadata catalog cannot be queried" + ) from error + + if row is None: + if explicit_version: + raise MetadataCatalogVersionNotFoundError( + f"PolicyEngine.py {selected_version} is not published " + f"for {country_id}" + ) + raise MetadataCatalogUnavailableError( + f"the running PolicyEngine.py {selected_version} catalog " + f"is absent for {country_id}" + ) + model, model_version = row + return SelectedCatalog( + country_id=country_id, + policyengine_version=selected_version, + model=model, + model_version=model_version, + ) + def get_metadata( self, country_id: str, diff --git a/policyengine_api/data/v2/catalog/schemas.py b/policyengine_api/data/v2/catalog/schemas.py index e3bb198d0..ec4cc3ee8 100644 --- a/policyengine_api/data/v2/catalog/schemas.py +++ b/policyengine_api/data/v2/catalog/schemas.py @@ -4,7 +4,7 @@ from datetime import datetime from enum import StrEnum -from typing import Annotated, Literal +from typing import Annotated, Generic, Literal, TypeVar from uuid import UUID from pydantic import BaseModel, ConfigDict, Field, JsonValue, StringConstraints @@ -75,6 +75,23 @@ class MetadataParameter(StrictResponseModel): values: list[MetadataParameterValue] +class MetadataParameterSummary(StrictResponseModel): + id: UUID + name: str + label: str | None + description: str | None + data_type: str | None + unit: str | None + + +class MetadataCanonicalParameterValue(StrictResponseModel): + id: UUID + parameter_id: UUID + value: JsonValue + start_date: datetime + end_date: datetime | None + + class MetadataDataset(StrictResponseModel): id: UUID name: str @@ -122,6 +139,159 @@ class MetadataEconomyOptions(StrictResponseModel): datasets: list[MetadataDatasetOption] +class MetadataModelVersionDetail(MetadataModelVersion): + current_law_id: int + metadata_time_periods: list[int] + + +class MetadataParameterChild(StrictResponseModel): + path: str + label: str + type: Literal["node", "parameter"] + child_count: int | None = None + parameter: MetadataParameterSummary | None = None + + +ResourceT = TypeVar("ResourceT") + + +class MetadataPageResult(StrictResponseModel, Generic[ResourceT]): + policyengine_version: str + items: list[ResourceT] + offset: int + limit: int + has_more: bool + + +class MetadataDetailResult(StrictResponseModel, Generic[ResourceT]): + policyengine_version: str + item: ResourceT + + +class MetadataModelSelectionResult(StrictResponseModel): + policyengine_version: str + model: MetadataModel + model_version: MetadataModelVersionDetail + + +class MetadataEconomyOptionsResult(StrictResponseModel): + policyengine_version: str + current_law_id: int + region: list[MetadataRegionOption] + time_period: list[MetadataTimePeriodOption] + datasets: list[MetadataDatasetOption] + + +class MetadataResourceSuccessResponse(StrictResponseModel, Generic[ResourceT]): + status: Literal["ok"] = "ok" + message: None = None + result: ResourceT + + +class MetadataModelPageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataModel]] +): + pass + + +class MetadataModelDetailResponse( + MetadataResourceSuccessResponse[MetadataDetailResult[MetadataModel]] +): + pass + + +class MetadataModelSelectionResponse( + MetadataResourceSuccessResponse[MetadataModelSelectionResult] +): + pass + + +class MetadataModelVersionPageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataModelVersionDetail]] +): + pass + + +class MetadataModelVersionDetailResponse( + MetadataResourceSuccessResponse[MetadataDetailResult[MetadataModelVersionDetail]] +): + pass + + +class MetadataVariablePageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataVariable]] +): + pass + + +class MetadataVariableDetailResponse( + MetadataResourceSuccessResponse[MetadataDetailResult[MetadataVariable]] +): + pass + + +class MetadataParameterPageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataParameterSummary]] +): + pass + + +class MetadataParameterDetailResponse( + MetadataResourceSuccessResponse[MetadataDetailResult[MetadataParameterSummary]] +): + pass + + +class MetadataParameterChildPageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataParameterChild]] +): + pass + + +class MetadataParameterValuePageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataCanonicalParameterValue]] +): + pass + + +class MetadataParameterValueDetailResponse( + MetadataResourceSuccessResponse[ + MetadataDetailResult[MetadataCanonicalParameterValue] + ] +): + pass + + +class MetadataDatasetPageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataDataset]] +): + pass + + +class MetadataDatasetDetailResponse( + MetadataResourceSuccessResponse[MetadataDetailResult[MetadataDataset]] +): + pass + + +class MetadataRegionPageResponse( + MetadataResourceSuccessResponse[MetadataPageResult[MetadataRegion]] +): + pass + + +class MetadataRegionDetailResponse( + MetadataResourceSuccessResponse[MetadataDetailResult[MetadataRegion]] +): + pass + + +class MetadataEconomyOptionsResponse( + MetadataResourceSuccessResponse[MetadataEconomyOptionsResult] +): + pass + + class MetadataResult(StrictResponseModel): current_law_id: int model: MetadataModel From dd6d97e724d91d5cec3137fb63e0b5908ec55749 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:35:07 +0400 Subject: [PATCH 14/27] Add bounded v2 metadata queries --- policyengine_api/data/v2/catalog/query.py | 689 ++++++++++++++++++++++ 1 file changed, 689 insertions(+) diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index 953600219..700a9d5cf 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -4,20 +4,32 @@ from collections import defaultdict from dataclasses import dataclass +from datetime import datetime, timezone +from typing import TypeVar +from uuid import UUID from packaging.version import InvalidVersion, Version +import sqlalchemy as sa from sqlalchemy.exc import SQLAlchemyError from sqlmodel import Session, select from policyengine_api.dataset_display import get_dataset_display_label from policyengine_api.data.v2.catalog.schemas import ( + MetadataCanonicalParameterValue, MetadataDataset, MetadataDatasetOption, + MetadataDetailResult, MetadataEconomyOptions, + MetadataEconomyOptionsResult, MetadataModel, MetadataModelVersion, + MetadataModelSelectionResult, + MetadataModelVersionDetail, + MetadataPageResult, MetadataParameter, + MetadataParameterChild, MetadataParameterNode, + MetadataParameterSummary, MetadataParameterValue, MetadataRegion, MetadataRegionOption, @@ -56,6 +68,14 @@ class MetadataCatalogVersionNotFoundError(LookupError): """Raised when an explicitly selected catalog version is absent.""" +class MetadataResourceNotFoundError(LookupError): + """Raised when a selected catalog does not contain a requested resource.""" + + +class InvalidMetadataPageError(ValueError): + """Raised when collection pagination is outside the documented bounds.""" + + @dataclass(frozen=True) class SelectedCatalog: """One country catalog selected by its canonical PolicyEngine.py version.""" @@ -66,6 +86,134 @@ class SelectedCatalog: model_version: TaxBenefitModelVersion +ResourceT = TypeVar("ResourceT") + + +def _page( + selected: SelectedCatalog, + rows: list[ResourceT], + *, + offset: int, + limit: int, +) -> MetadataPageResult[ResourceT]: + return MetadataPageResult( + policyengine_version=selected.policyengine_version, + items=rows[:limit], + offset=offset, + limit=limit, + has_more=len(rows) > limit, + ) + + +def validate_metadata_page(offset: int, limit: int) -> tuple[int, int]: + if offset < 0: + raise InvalidMetadataPageError("offset must be at least 0") + if not 1 <= limit <= 500: + raise InvalidMetadataPageError("limit must be between 1 and 500") + return offset, limit + + +def _metadata_model(selected: SelectedCatalog) -> MetadataModel: + return MetadataModel( + id=selected.model.id, + name=selected.model.name, + description=selected.model_version.description, + ) + + +def _metadata_model_version(selected: SelectedCatalog) -> MetadataModelVersionDetail: + return MetadataModelVersionDetail( + id=selected.model_version.id, + model_id=selected.model.id, + version=selected.model_version.version, + description=selected.model_version.description, + current_law_id=selected.model_version.current_law_id, + metadata_time_periods=selected.model_version.metadata_time_periods, + ) + + +def _metadata_variable(variable: Variable) -> MetadataVariable: + return MetadataVariable( + id=variable.id, + name=variable.name, + label=variable.label, + entity=variable.entity, + description=variable.description, + data_type=variable.data_type, + possible_values=variable.possible_values, + default_value=variable.default_value, + adds=variable.adds, + subtracts=variable.subtracts, + ) + + +def _metadata_parameter(parameter: Parameter) -> MetadataParameterSummary: + return MetadataParameterSummary( + id=parameter.id, + name=parameter.name, + label=parameter.label, + description=parameter.description, + data_type=parameter.data_type, + unit=parameter.unit, + ) + + +def _metadata_parameter_value( + value: ParameterValue, +) -> MetadataCanonicalParameterValue: + return MetadataCanonicalParameterValue( + id=value.id, + parameter_id=value.parameter_id, + value=value.value_json, + start_date=value.start_date, + end_date=value.end_date, + ) + + +def _metadata_dataset(dataset: Dataset) -> MetadataDataset: + return MetadataDataset( + id=dataset.id, + name=dataset.name, + description=dataset.description, + year=dataset.year, + ) + + +def _metadata_region(region: Region) -> MetadataRegion: + return MetadataRegion( + id=region.id, + code=region.code, + label=region.label, + region_type=region.region_type.value, + requires_filter=region.requires_filter, + filter_field=region.filter_field, + filter_value=region.filter_value, + filter_strategy=region.filter_strategy, + parent_code=region.parent_code, + state_code=region.state_code, + state_name=region.state_name, + default_dataset_id=region.default_dataset_id, + ) + + +def _escaped_like(value: str) -> str: + return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + + +def _child_path_expression(column: object, prefix: str, dialect: str) -> object: + remainder = sa.func.substr(column, len(prefix) + 1) + dot_position = ( + sa.func.instr(remainder, ".") + if dialect == "sqlite" + else sa.func.strpos(remainder, ".") + ) + segment = sa.case( + (dot_position > 0, sa.func.substr(remainder, 1, dot_position - 1)), + else_=remainder, + ) + return sa.literal(prefix) + segment + + def validate_policyengine_version(value: str) -> str: """Return one bounded canonical PEP 440 version string.""" @@ -154,6 +302,547 @@ def select_catalog( model_version=model_version, ) + def _resource_rows(self, statement: object) -> list: + try: + return list(self._session.exec(statement).all()) + except SQLAlchemyError as error: + raise MetadataCatalogUnavailableError( + "the v2 metadata catalog cannot be queried" + ) from error + + def list_models( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataModel]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + rows = [_metadata_model(selected)] if offset == 0 else [] + return _page(selected, rows, offset=offset, limit=limit) + + def get_model( + self, + country_id: str, + model_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataModel]: + selected = self.select_catalog(country_id, policyengine_version) + if selected.model.id != model_id: + raise MetadataResourceNotFoundError(f"model {model_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_model(selected), + ) + + def get_model_by_country( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataModelSelectionResult: + selected = self.select_catalog(country_id, policyengine_version) + return MetadataModelSelectionResult( + policyengine_version=selected.policyengine_version, + model=_metadata_model(selected), + model_version=_metadata_model_version(selected), + ) + + def list_model_versions( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataModelVersionDetail]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + rows = [_metadata_model_version(selected)] if offset == 0 else [] + return _page(selected, rows, offset=offset, limit=limit) + + def get_model_version( + self, + country_id: str, + version_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataModelVersionDetail]: + selected = self.select_catalog(country_id, policyengine_version) + if selected.model_version.id != version_id: + raise MetadataResourceNotFoundError( + f"model version {version_id} was not found" + ) + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_model_version(selected), + ) + + def list_variables( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + search: str | None = None, + ) -> MetadataPageResult[MetadataVariable]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + statement = select(Variable).where( + Variable.tax_benefit_model_version_id == selected.model_version.id + ) + if search: + pattern = f"%{_escaped_like(search)}%" + statement = statement.where( + sa.or_( + Variable.name.ilike(pattern, escape="\\"), + Variable.label.ilike(pattern, escape="\\"), + Variable.description.ilike(pattern, escape="\\"), + ) + ) + rows = self._resource_rows( + statement.order_by(Variable.name).offset(offset).limit(limit + 1) + ) + return _page( + selected, + [_metadata_variable(row) for row in rows], + offset=offset, + limit=limit, + ) + + def get_variable( + self, + country_id: str, + variable_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataVariable]: + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(Variable).where( + Variable.id == variable_id, + Variable.tax_benefit_model_version_id == selected.model_version.id, + ) + ) + if not rows: + raise MetadataResourceNotFoundError(f"variable {variable_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_variable(rows[0]), + ) + + def list_parameters( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + search: str | None = None, + ) -> MetadataPageResult[MetadataParameterSummary]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + statement = select(Parameter).where( + Parameter.tax_benefit_model_version_id == selected.model_version.id + ) + if search: + pattern = f"%{_escaped_like(search)}%" + statement = statement.where( + sa.or_( + Parameter.name.ilike(pattern, escape="\\"), + Parameter.label.ilike(pattern, escape="\\"), + Parameter.description.ilike(pattern, escape="\\"), + ) + ) + rows = self._resource_rows( + statement.order_by(Parameter.name).offset(offset).limit(limit + 1) + ) + return _page( + selected, + [_metadata_parameter(row) for row in rows], + offset=offset, + limit=limit, + ) + + def get_parameter( + self, + country_id: str, + parameter_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataParameterSummary]: + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(Parameter).where( + Parameter.id == parameter_id, + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ) + ) + if not rows: + raise MetadataResourceNotFoundError( + f"parameter {parameter_id} was not found" + ) + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_parameter(rows[0]), + ) + + def list_parameter_children( + self, + country_id: str, + policyengine_version: str | None = None, + *, + parent_path: str = "", + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataParameterChild]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + version_id = selected.model_version.id + prefix = f"{parent_path}." if parent_path else "" + escaped_prefix = _escaped_like(prefix) + dialect = self._session.get_bind().dialect.name + node_child_path = _child_path_expression(ParameterNode.name, prefix, dialect) + parameter_child_path = _child_path_expression( + Parameter.name, + prefix, + dialect, + ) + paths = sa.union( + select(node_child_path.label("path")).where( + ParameterNode.tax_benefit_model_version_id == version_id, + ParameterNode.name.like(f"{escaped_prefix}%", escape="\\"), + ), + select(parameter_child_path.label("path")).where( + Parameter.tax_benefit_model_version_id == version_id, + Parameter.name.like(f"{escaped_prefix}%", escape="\\"), + ), + ).subquery() + descendant_count = ( + select(sa.func.count(Parameter.id)) + .where( + Parameter.tax_benefit_model_version_id == version_id, + sa.func.substr( + Parameter.name, + 1, + sa.func.length(paths.c.path) + 1, + ) + == paths.c.path + ".", + ) + .correlate(paths) + .scalar_subquery() + ) + is_node = sa.or_(descendant_count > 0, Parameter.id.is_(None)) + statement = ( + select( + paths.c.path, + sa.func.coalesce(ParameterNode.label, Parameter.label).label("label"), + sa.case((is_node, "node"), else_="parameter").label("type"), + sa.case((is_node, descendant_count), else_=None).label("child_count"), + Parameter.id.label("parameter_id"), + Parameter.label.label("parameter_label"), + Parameter.description.label("parameter_description"), + Parameter.data_type.label("parameter_data_type"), + Parameter.unit.label("parameter_unit"), + ) + .select_from( + paths.outerjoin( + ParameterNode, + sa.and_( + ParameterNode.name == paths.c.path, + ParameterNode.tax_benefit_model_version_id == version_id, + ), + ).outerjoin( + Parameter, + sa.and_( + Parameter.name == paths.c.path, + Parameter.tax_benefit_model_version_id == version_id, + ), + ) + ) + .order_by(paths.c.path) + .offset(offset) + .limit(limit + 1) + ) + rows = self._resource_rows(statement) + items = [] + for row in rows: + parameter = None + if row.type == "parameter": + parameter = MetadataParameterSummary( + id=row.parameter_id, + name=row.path, + label=row.parameter_label, + description=row.parameter_description, + data_type=row.parameter_data_type, + unit=row.parameter_unit, + ) + items.append( + MetadataParameterChild( + path=row.path, + label=row.label or row.path.rsplit(".", 1)[-1], + type=row.type, + child_count=row.child_count, + parameter=parameter, + ) + ) + return _page(selected, items, offset=offset, limit=limit) + + def list_parameter_values( + self, + country_id: str, + policyengine_version: str | None = None, + *, + parameter_id: UUID | None = None, + current: bool = False, + offset: int = 0, + limit: int = 100, + now: datetime | None = None, + ) -> MetadataPageResult[MetadataCanonicalParameterValue]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + statement = ( + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ) + ) + if parameter_id is not None: + statement = statement.where(ParameterValue.parameter_id == parameter_id) + if current: + selected_time = now or datetime.now(timezone.utc) + statement = statement.where( + ParameterValue.start_date <= selected_time, + sa.or_( + ParameterValue.end_date.is_(None), + ParameterValue.end_date > selected_time, + ), + ) + rows = self._resource_rows( + statement.order_by( + Parameter.name, + ParameterValue.start_date.desc(), + ParameterValue.id, + ) + .offset(offset) + .limit(limit + 1) + ) + return _page( + selected, + [_metadata_parameter_value(row) for row in rows], + offset=offset, + limit=limit, + ) + + def get_parameter_value( + self, + country_id: str, + value_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataCanonicalParameterValue]: + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + ParameterValue.id == value_id, + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ) + ) + if not rows: + raise MetadataResourceNotFoundError( + f"parameter value {value_id} was not found" + ) + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_parameter_value(rows[0]), + ) + + def list_datasets( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataDataset]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(Dataset) + .where( + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ) + .order_by(Dataset.name) + .offset(offset) + .limit(limit + 1) + ) + return _page( + selected, + [_metadata_dataset(row) for row in rows], + offset=offset, + limit=limit, + ) + + def get_dataset( + self, + country_id: str, + dataset_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataDataset]: + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(Dataset).where( + Dataset.id == dataset_id, + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ) + ) + if not rows: + raise MetadataResourceNotFoundError(f"dataset {dataset_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_dataset(rows[0]), + ) + + def list_regions( + self, + country_id: str, + policyengine_version: str | None = None, + *, + region_type: str | None = None, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataRegion]: + validate_metadata_page(offset, limit) + selected = self.select_catalog(country_id, policyengine_version) + statement = select(Region).where( + Region.tax_benefit_model_version_id == selected.model_version.id + ) + if region_type is not None: + statement = statement.where(Region.region_type == region_type) + rows = self._resource_rows( + statement.order_by(Region.code).offset(offset).limit(limit + 1) + ) + return _page( + selected, + [_metadata_region(row) for row in rows], + offset=offset, + limit=limit, + ) + + def get_region( + self, + country_id: str, + region_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataRegion]: + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(Region).where( + Region.id == region_id, + Region.tax_benefit_model_version_id == selected.model_version.id, + ) + ) + if not rows: + raise MetadataResourceNotFoundError(f"region {region_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_region(rows[0]), + ) + + def get_region_by_code( + self, + country_id: str, + region_code: str, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataRegion]: + selected = self.select_catalog(country_id, policyengine_version) + rows = self._resource_rows( + select(Region).where( + Region.code == region_code, + Region.tax_benefit_model_version_id == selected.model_version.id, + ) + ) + if not rows: + raise MetadataResourceNotFoundError(f"region {region_code!r} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_metadata_region(rows[0]), + ) + + def get_economy_options( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataEconomyOptionsResult: + selected = self.select_catalog(country_id, policyengine_version) + regions = self._resource_rows( + select(Region) + .where(Region.tax_benefit_model_version_id == selected.model_version.id) + .order_by(Region.code) + ) + national_region = next( + (region for region in regions if region.code == country_id), + None, + ) + if national_region is None: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 region is absent" + ) + datasets = self._resource_rows( + select(Dataset).where( + Dataset.id == national_region.default_dataset_id, + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ) + ) + if len(datasets) != 1: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 dataset is absent" + ) + time_periods = selected.model_version.metadata_time_periods + if ( + not isinstance(selected.model_version.current_law_id, int) + or not isinstance(time_periods, list) + or not time_periods + or any(not isinstance(year, int) for year in time_periods) + ): + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 model-version options are incomplete" + ) + national_dataset = datasets[0] + return MetadataEconomyOptionsResult( + policyengine_version=selected.policyengine_version, + current_law_id=selected.model_version.current_law_id, + region=[ + MetadataRegionOption( + name=region.code, + label=region.label, + type=region.region_type.value, + ) + for region in regions + ], + time_period=[ + MetadataTimePeriodOption(name=year, label=str(year)) + for year in time_periods + ], + datasets=[ + MetadataDatasetOption( + name=national_dataset.name, + label=get_dataset_display_label(national_dataset.name), + ) + ], + ) + def get_metadata( self, country_id: str, From 1ee83a585474db9bd8d80145ea4efaf6dad22391 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:38:42 +0400 Subject: [PATCH 15/27] Expose split v2 metadata routes --- policyengine_api/asgi_factory.py | 15 ++ .../fastapi_routes/v2_metadata.py | 12 ++ .../fastapi_routes/v2_metadata_common.py | 97 +++++++++ .../fastapi_routes/v2_metadata_geography.py | 165 +++++++++++++++ .../fastapi_routes/v2_metadata_models.py | 189 ++++++++++++++++++ .../fastapi_routes/v2_metadata_parameters.py | 153 ++++++++++++++ tests/unit/v2/test_metadata_routes.py | 18 ++ 7 files changed, 649 insertions(+) create mode 100644 policyengine_api/fastapi_routes/v2_metadata_common.py create mode 100644 policyengine_api/fastapi_routes/v2_metadata_geography.py create mode 100644 policyengine_api/fastapi_routes/v2_metadata_models.py create mode 100644 policyengine_api/fastapi_routes/v2_metadata_parameters.py diff --git a/policyengine_api/asgi_factory.py b/policyengine_api/asgi_factory.py index 89d96143b..9a1a76435 100644 --- a/policyengine_api/asgi_factory.py +++ b/policyengine_api/asgi_factory.py @@ -8,6 +8,8 @@ from a2wsgi import WSGIMiddleware 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 @@ -109,6 +111,19 @@ async def add_headers_to_unhandled_errors( _apply_shared_response_headers(request, response, request_id) return response + @app.exception_handler(RequestValidationError) + async def typed_v2_request_validation_error( + request: Request, + error: RequestValidationError, + ) -> Response: + if request.url.path.startswith("/v2/"): + from policyengine_api.fastapi_routes.v2_metadata_common import ( + error_response, + ) + + return error_response(422, "Invalid v2 metadata request") + return await request_validation_exception_handler(request, error) + @app.middleware("http") async def add_cors_for_native_routes(request, call_next): started_at = time.time() diff --git a/policyengine_api/fastapi_routes/v2_metadata.py b/policyengine_api/fastapi_routes/v2_metadata.py index 410815eb0..75b290a6e 100644 --- a/policyengine_api/fastapi_routes/v2_metadata.py +++ b/policyengine_api/fastapi_routes/v2_metadata.py @@ -16,6 +16,15 @@ ) from policyengine_api.data.v2.settings import V2ConfigurationError from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies +from policyengine_api.fastapi_routes.v2_metadata_geography import ( + build_v2_metadata_geography_router, +) +from policyengine_api.fastapi_routes.v2_metadata_models import ( + build_v2_metadata_model_router, +) +from policyengine_api.fastapi_routes.v2_metadata_parameters import ( + build_v2_metadata_parameter_router, +) ERROR_RESPONSES = { @@ -56,6 +65,9 @@ def build_v2_metadata_router( """Build isolated preview routes without loading v2 configuration.""" router = APIRouter() + router.include_router(build_v2_metadata_model_router(dependencies)) + router.include_router(build_v2_metadata_parameter_router(dependencies)) + router.include_router(build_v2_metadata_geography_router(dependencies)) @router.get( "/v2/openapi.json", diff --git a/policyengine_api/fastapi_routes/v2_metadata_common.py b/policyengine_api/fastapi_routes/v2_metadata_common.py new file mode 100644 index 000000000..de54e75c4 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_common.py @@ -0,0 +1,97 @@ +"""Shared response handling for dormant v2 metadata resources.""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import TypeVar + +from pydantic import BaseModel +from starlette.responses import JSONResponse + +from policyengine_api.data.v2.catalog.query import ( + InvalidMetadataPageError, + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, + MetadataResourceNotFoundError, + UnsupportedPreviewCountryError, +) +from policyengine_api.data.v2.catalog.schemas import MetadataErrorResponse +from policyengine_api.data.v2.settings import V2ConfigurationError +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies + + +ERROR_RESPONSES = { + 400: { + "model": MetadataErrorResponse, + "description": "The resource request or PolicyEngine.py version is invalid.", + }, + 404: { + "model": MetadataErrorResponse, + "description": "The requested catalog resource is absent.", + }, + 405: { + "model": MetadataErrorResponse, + "description": "The dormant v2 metadata resources support GET only.", + }, + 422: { + "model": MetadataErrorResponse, + "description": "The request parameters do not match the resource schema.", + }, + 500: { + "model": MetadataErrorResponse, + "description": "The resource query failed internally.", + }, + 503: { + "model": MetadataErrorResponse, + "description": "The initialized v2 catalog is unavailable.", + }, +} + + +def error_response(status_code: int, message: str) -> JSONResponse: + error = MetadataErrorResponse(message=message) + return JSONResponse( + status_code=status_code, + content=error.model_dump(mode="json"), + ) + + +ResponseT = TypeVar("ResponseT", bound=BaseModel) + + +def read_resource( + dependencies: NativeRouteDependencies, + response_type: type[ResponseT], + operation: Callable[[object], object], +) -> ResponseT | JSONResponse: + reader = None + try: + factory = dependencies.v2_metadata_reader_factory + if factory is None: + from policyengine_api.fastapi_routes.dependencies import ( + _default_v2_metadata_reader_factory, + ) + + factory = _default_v2_metadata_reader_factory + reader = factory() + return response_type(result=operation(reader)) + except (InvalidMetadataPageError, InvalidPolicyEngineVersionError) as error: + return error_response(400, str(error)) + except UnsupportedPreviewCountryError as error: + return error_response(400, f"Unsupported country: {error}") + except ( + MetadataCatalogVersionNotFoundError, + MetadataResourceNotFoundError, + ) as error: + return error_response(404, str(error)) + except (V2ConfigurationError, MetadataCatalogUnavailableError): + return error_response(503, "V2 metadata catalog is unavailable") + except Exception: # noqa: BLE001 - preview must return typed errors + return error_response(500, "V2 metadata query failed") + finally: + if reader is not None: + try: + reader.close() + except Exception: + pass diff --git a/policyengine_api/fastapi_routes/v2_metadata_geography.py b/policyengine_api/fastapi_routes/v2_metadata_geography.py new file mode 100644 index 000000000..8114a74c0 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_geography.py @@ -0,0 +1,165 @@ +"""Dataset, region, and economy-option preview routes.""" + +from __future__ import annotations + +from typing import Annotated +from uuid import UUID + +from fastapi import APIRouter, Query +from starlette.responses import JSONResponse + +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDatasetDetailResponse, + MetadataDatasetPageResponse, + MetadataEconomyOptionsResponse, + MetadataRegionDetailResponse, + MetadataRegionPageResponse, + MetadataRegionType, +) +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies +from policyengine_api.fastapi_routes.v2_metadata_common import ( + ERROR_RESPONSES, + read_resource, +) + + +Offset = Annotated[int, Query(ge=0)] +Limit = Annotated[int, Query(ge=1, le=500)] + + +def build_v2_metadata_geography_router( + dependencies: NativeRouteDependencies, +) -> APIRouter: + router = APIRouter(prefix="/v2") + + @router.get( + "/datasets", + response_model=MetadataDatasetPageResponse, + responses=ERROR_RESPONSES, + summary="List logical inputs from a selected PolicyEngine.py catalog", + ) + def list_datasets( + country_id: str, + policyengine_version: str | None = None, + offset: Offset = 0, + limit: Limit = 100, + ) -> MetadataDatasetPageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataDatasetPageResponse, + lambda reader: reader.list_datasets( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + ) + + @router.get( + "/datasets/{dataset_id}", + response_model=MetadataDatasetDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one logical input from a selected PolicyEngine.py catalog", + ) + def get_dataset( + dataset_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataDatasetDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataDatasetDetailResponse, + lambda reader: reader.get_dataset( + country_id, + dataset_id, + policyengine_version, + ), + ) + + @router.get( + "/regions", + response_model=MetadataRegionPageResponse, + responses=ERROR_RESPONSES, + summary="List regions from a selected PolicyEngine.py catalog", + ) + def list_regions( + country_id: str, + policyengine_version: str | None = None, + region_type: MetadataRegionType | None = None, + offset: Offset = 0, + limit: Limit = 100, + ) -> MetadataRegionPageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataRegionPageResponse, + lambda reader: reader.list_regions( + country_id, + policyengine_version, + region_type=region_type.value if region_type is not None else None, + offset=offset, + limit=limit, + ), + ) + + @router.get( + "/regions/by-code/{region_code:path}", + response_model=MetadataRegionDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one region by code from a selected catalog", + ) + def get_region_by_code( + region_code: str, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataRegionDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataRegionDetailResponse, + lambda reader: reader.get_region_by_code( + country_id, + region_code, + policyengine_version, + ), + ) + + @router.get( + "/regions/{region_id}", + response_model=MetadataRegionDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one region from a selected PolicyEngine.py catalog", + ) + def get_region( + region_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataRegionDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataRegionDetailResponse, + lambda reader: reader.get_region( + country_id, + region_id, + policyengine_version, + ), + ) + + @router.get( + "/economy-options", + response_model=MetadataEconomyOptionsResponse, + responses=ERROR_RESPONSES, + summary="Get compact economy-selection options from a selected catalog", + ) + def get_economy_options( + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataEconomyOptionsResponse | JSONResponse: + return read_resource( + dependencies, + MetadataEconomyOptionsResponse, + lambda reader: reader.get_economy_options( + country_id, + policyengine_version, + ), + ) + + return router diff --git a/policyengine_api/fastapi_routes/v2_metadata_models.py b/policyengine_api/fastapi_routes/v2_metadata_models.py new file mode 100644 index 000000000..e9ab4d6f8 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_models.py @@ -0,0 +1,189 @@ +"""Tax-benefit model, version, and variable preview routes.""" + +from __future__ import annotations + +from typing import Annotated +from uuid import UUID + +from fastapi import APIRouter, Query +from starlette.responses import JSONResponse + +from policyengine_api.data.v2.catalog.schemas import ( + MetadataModelDetailResponse, + MetadataModelPageResponse, + MetadataModelSelectionResponse, + MetadataModelVersionDetailResponse, + MetadataModelVersionPageResponse, + MetadataVariableDetailResponse, + MetadataVariablePageResponse, +) +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies +from policyengine_api.fastapi_routes.v2_metadata_common import ( + ERROR_RESPONSES, + read_resource, +) + + +Offset = Annotated[int, Query(ge=0)] +Limit = Annotated[int, Query(ge=1, le=500)] + + +def build_v2_metadata_model_router( + dependencies: NativeRouteDependencies, +) -> APIRouter: + router = APIRouter(prefix="/v2") + + @router.get( + "/tax-benefit-models", + response_model=MetadataModelPageResponse, + responses=ERROR_RESPONSES, + summary="List models for one PolicyEngine.py catalog", + ) + def list_models( + country_id: str, + policyengine_version: str | None = None, + offset: Offset = 0, + limit: Limit = 100, + ) -> MetadataModelPageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataModelPageResponse, + lambda reader: reader.list_models( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + ) + + @router.get( + "/tax-benefit-models/by-country/{country_id}", + response_model=MetadataModelSelectionResponse, + responses=ERROR_RESPONSES, + summary="Get a country model and selected PolicyEngine.py version", + ) + def get_model_by_country( + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataModelSelectionResponse | JSONResponse: + return read_resource( + dependencies, + MetadataModelSelectionResponse, + lambda reader: reader.get_model_by_country( + country_id, + policyengine_version, + ), + ) + + @router.get( + "/tax-benefit-models/{model_id}", + response_model=MetadataModelDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one model from a selected PolicyEngine.py catalog", + ) + def get_model( + model_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataModelDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataModelDetailResponse, + lambda reader: reader.get_model( + country_id, + model_id, + policyengine_version, + ), + ) + + @router.get( + "/tax-benefit-model-versions", + response_model=MetadataModelVersionPageResponse, + responses=ERROR_RESPONSES, + summary="List selected PolicyEngine.py model versions", + ) + def list_model_versions( + country_id: str, + policyengine_version: str | None = None, + offset: Offset = 0, + limit: Limit = 100, + ) -> MetadataModelVersionPageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataModelVersionPageResponse, + lambda reader: reader.list_model_versions( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + ) + + @router.get( + "/tax-benefit-model-versions/{version_id}", + response_model=MetadataModelVersionDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one selected PolicyEngine.py model version", + ) + def get_model_version( + version_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataModelVersionDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataModelVersionDetailResponse, + lambda reader: reader.get_model_version( + country_id, + version_id, + policyengine_version, + ), + ) + + @router.get( + "/variables", + response_model=MetadataVariablePageResponse, + responses=ERROR_RESPONSES, + summary="List variables from a selected PolicyEngine.py catalog", + ) + def list_variables( + country_id: str, + policyengine_version: str | None = None, + offset: Offset = 0, + limit: Limit = 100, + search: str | None = None, + ) -> MetadataVariablePageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataVariablePageResponse, + lambda reader: reader.list_variables( + country_id, + policyengine_version, + offset=offset, + limit=limit, + search=search, + ), + ) + + @router.get( + "/variables/{variable_id}", + response_model=MetadataVariableDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one variable from a selected PolicyEngine.py catalog", + ) + def get_variable( + variable_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataVariableDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataVariableDetailResponse, + lambda reader: reader.get_variable( + country_id, + variable_id, + policyengine_version, + ), + ) + + return router diff --git a/policyengine_api/fastapi_routes/v2_metadata_parameters.py b/policyengine_api/fastapi_routes/v2_metadata_parameters.py new file mode 100644 index 000000000..d94c52fb4 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_parameters.py @@ -0,0 +1,153 @@ +"""Parameter and canonical parameter-value preview routes.""" + +from __future__ import annotations + +from typing import Annotated +from uuid import UUID + +from fastapi import APIRouter, Query +from starlette.responses import JSONResponse + +from policyengine_api.data.v2.catalog.schemas import ( + MetadataParameterChildPageResponse, + MetadataParameterDetailResponse, + MetadataParameterPageResponse, + MetadataParameterValueDetailResponse, + MetadataParameterValuePageResponse, +) +from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies +from policyengine_api.fastapi_routes.v2_metadata_common import ( + ERROR_RESPONSES, + read_resource, +) + + +Offset = Annotated[int, Query(ge=0)] +Limit = Annotated[int, Query(ge=1, le=500)] + + +def build_v2_metadata_parameter_router( + dependencies: NativeRouteDependencies, +) -> APIRouter: + router = APIRouter(prefix="/v2") + + @router.get( + "/parameters", + response_model=MetadataParameterPageResponse, + responses=ERROR_RESPONSES, + summary="List parameters from a selected PolicyEngine.py catalog", + ) + def list_parameters( + country_id: str, + policyengine_version: str | None = None, + offset: Offset = 0, + limit: Limit = 100, + search: str | None = None, + ) -> MetadataParameterPageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataParameterPageResponse, + lambda reader: reader.list_parameters( + country_id, + policyengine_version, + offset=offset, + limit=limit, + search=search, + ), + ) + + @router.get( + "/parameters/children", + response_model=MetadataParameterChildPageResponse, + responses=ERROR_RESPONSES, + summary="List direct children of one parameter path", + ) + def list_parameter_children( + country_id: str, + parent_path: str = "", + policyengine_version: str | None = None, + offset: Offset = 0, + limit: Limit = 100, + ) -> MetadataParameterChildPageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataParameterChildPageResponse, + lambda reader: reader.list_parameter_children( + country_id, + policyengine_version, + parent_path=parent_path, + offset=offset, + limit=limit, + ), + ) + + @router.get( + "/parameters/{parameter_id}", + response_model=MetadataParameterDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one parameter from a selected PolicyEngine.py catalog", + ) + def get_parameter( + parameter_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataParameterDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataParameterDetailResponse, + lambda reader: reader.get_parameter( + country_id, + parameter_id, + policyengine_version, + ), + ) + + @router.get( + "/parameter-values", + response_model=MetadataParameterValuePageResponse, + responses=ERROR_RESPONSES, + summary="List canonical values from a selected PolicyEngine.py catalog", + ) + def list_parameter_values( + country_id: str, + policyengine_version: str | None = None, + parameter_id: UUID | None = None, + current: bool = False, + offset: Offset = 0, + limit: Limit = 100, + ) -> MetadataParameterValuePageResponse | JSONResponse: + return read_resource( + dependencies, + MetadataParameterValuePageResponse, + lambda reader: reader.list_parameter_values( + country_id, + policyengine_version, + parameter_id=parameter_id, + current=current, + offset=offset, + limit=limit, + ), + ) + + @router.get( + "/parameter-values/{value_id}", + response_model=MetadataParameterValueDetailResponse, + responses=ERROR_RESPONSES, + summary="Get one canonical value from a selected PolicyEngine.py catalog", + ) + def get_parameter_value( + value_id: UUID, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataParameterValueDetailResponse | JSONResponse: + return read_resource( + dependencies, + MetadataParameterValueDetailResponse, + lambda reader: reader.get_parameter_value( + country_id, + value_id, + policyengine_version, + ), + ) + + return router diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py index 4cbc65b5e..ff7c892ed 100644 --- a/tests/unit/v2/test_metadata_routes.py +++ b/tests/unit/v2/test_metadata_routes.py @@ -400,8 +400,26 @@ def test_openapi_references_explicit_preview_response_schemas() -> None: assert response.headers["content-type"].startswith("application/json") schema = response.json() assert set(schema["paths"]) == { + "/v2/datasets", + "/v2/datasets/{dataset_id}", + "/v2/economy-options", + "/v2/parameters", + "/v2/parameters/children", + "/v2/parameters/{parameter_id}", + "/v2/parameter-values", + "/v2/parameter-values/{value_id}", + "/v2/regions", + "/v2/regions/by-code/{region_code}", + "/v2/regions/{region_id}", + "/v2/tax-benefit-models", + "/v2/tax-benefit-models/by-country/{country_id}", + "/v2/tax-benefit-models/{model_id}", + "/v2/tax-benefit-model-versions", + "/v2/tax-benefit-model-versions/{version_id}", "/v2/us/metadata", "/v2/uk/metadata", + "/v2/variables", + "/v2/variables/{variable_id}", "/v2/{country_id}/metadata", } From a5d30be0e37787af249963257139ac61cf84e169 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:41:32 +0400 Subject: [PATCH 16/27] Test split v2 metadata reads --- tests/unit/v2/test_metadata_query.py | 249 +++++++++++++++++++++++++- tests/unit/v2/test_metadata_routes.py | 248 +++++++++++++++++++++++++ 2 files changed, 496 insertions(+), 1 deletion(-) diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index eba0ec88f..b42af29b2 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -3,19 +3,22 @@ from __future__ import annotations import ast +from datetime import datetime, timezone from pathlib import Path from uuid import uuid4 from pydantic import TypeAdapter, ValidationError import pytest -from sqlalchemy import delete +from sqlalchemy import delete, event from sqlalchemy.pool import StaticPool from sqlmodel import Session, create_engine, select from policyengine_api.data.v2.catalog.query import ( + InvalidMetadataPageError, InvalidPolicyEngineVersionError, MetadataCatalogUnavailableError, MetadataCatalogVersionNotFoundError, + MetadataResourceNotFoundError, UnsupportedPreviewCountryError, V2MetadataQueryService, ) @@ -621,3 +624,247 @@ def test_query_module_imports_no_policyengine_or_v1_metadata_source() -> None: ) for module in imported ) + + +def test_resource_collection_uses_bounded_pagination_without_counting( + catalog_session: Session, +) -> None: + model_version = catalog_session.exec( + select(TaxBenefitModelVersion) + .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) + .where( + TaxBenefitModel.name == "policyengine-us", + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ).one() + catalog_session.add( + Variable( + id=uuid4(), + tax_benefit_model_version_id=model_version.id, + name="pension_income", + label="Pension income", + entity="person", + description="Pension income before tax", + data_type="float", + possible_values=None, + default_value=0, + adds=None, + subtracts=None, + ) + ) + catalog_session.commit() + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, *_args) -> None: + statements.append(statement.lower()) + + bind = catalog_session.get_bind() + event.listen(bind, "before_cursor_execute", record_statement) + try: + result = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ).list_variables("us", offset=0, limit=1) + finally: + event.remove(bind, "before_cursor_execute", record_statement) + + assert [item.name for item in result.items] == ["employment_income"] + assert result.policyengine_version == POLICYENGINE_VERSION + assert result.offset == 0 + assert result.limit == 1 + assert result.has_more is True + assert len(statements) == 2 + assert all("count(" not in statement for statement in statements) + assert "limit ? offset ?" in statements[-1] + + +@pytest.mark.parametrize( + ("offset", "limit"), + [(-1, 100), (0, 0), (0, 501)], +) +def test_resource_collection_rejects_out_of_range_pages_before_querying( + catalog_session: Session, + offset: int, + limit: int, +) -> None: + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, *_args) -> None: + statements.append(statement) + + bind = catalog_session.get_bind() + event.listen(bind, "before_cursor_execute", record_statement) + try: + with pytest.raises(InvalidMetadataPageError): + V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ).list_variables("us", offset=offset, limit=limit) + finally: + event.remove(bind, "before_cursor_execute", record_statement) + + assert statements == [] + + +def test_parameter_collection_does_not_query_or_embed_values( + catalog_session: Session, +) -> None: + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, *_args) -> None: + statements.append(statement.lower()) + + bind = catalog_session.get_bind() + event.listen(bind, "before_cursor_execute", record_statement) + try: + result = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ).list_parameters("us") + finally: + event.remove(bind, "before_cursor_execute", record_statement) + + assert [item.name for item in result.items] == ["gov.example.rate"] + assert "values" not in result.items[0].model_dump() + assert len(statements) == 2 + assert all("parameter_values" not in statement for statement in statements) + + +def test_parameter_values_are_separate_canonical_resources( + catalog_session: Session, +) -> None: + parameter = catalog_session.exec( + select(Parameter) + .join( + TaxBenefitModelVersion, + TaxBenefitModelVersion.id == Parameter.tax_benefit_model_version_id, + ) + .where( + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + Parameter.name == "gov.example.rate", + ) + ).first() + override_id = uuid4() + catalog_session.add( + ParameterValue( + id=uuid4(), + parameter_id=parameter.id, + value_json=0.9, + start_date=datetime(2026, 1, 1, tzinfo=timezone.utc), + end_date=None, + policy_id=override_id, + ) + ) + catalog_session.commit() + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + all_values = service.list_parameter_values( + "us", + parameter_id=parameter.id, + ) + current_value = service.list_parameter_values( + "us", + parameter_id=parameter.id, + current=True, + now=datetime(2026, 6, 1, tzinfo=timezone.utc), + ) + + assert [item.value for item in all_values.items] == [0.2, 0.1] + assert [item.value for item in current_value.items] == [0.2] + assert all(item.parameter_id == parameter.id for item in all_values.items) + + +def test_parameter_children_are_loaded_one_level_at_a_time( + catalog_session: Session, +) -> None: + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + root = service.list_parameter_children("us") + government = service.list_parameter_children("us", parent_path="gov") + example = service.list_parameter_children("us", parent_path="gov.example") + + assert [(item.path, item.type) for item in root.items] == [("gov", "node")] + assert [(item.path, item.type) for item in government.items] == [ + ("gov.example", "node") + ] + assert [(item.path, item.type) for item in example.items] == [ + ("gov.example.rate", "parameter") + ] + assert root.items[0].child_count == 1 + assert example.items[0].parameter is not None + assert example.items[0].parameter.name == "gov.example.rate" + + +def test_resource_filters_and_details_remain_inside_selected_catalog( + catalog_session: Session, +) -> None: + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + variables = service.list_variables("us", search="employment") + parameters = service.list_parameters("us", search="example rate") + states = service.list_regions("us", region_type="state") + region = service.get_region_by_code("us", "state/ca") + + assert [item.name for item in variables.items] == ["employment_income"] + assert [item.name for item in parameters.items] == ["gov.example.rate"] + assert [item.code for item in states.items] == ["state/ca"] + assert region.item.code == "state/ca" + with pytest.raises(MetadataResourceNotFoundError): + service.get_region("us", uuid4()) + + +def test_each_resource_result_identifies_an_exact_selected_version( + catalog_session: Session, +) -> None: + _add_country_version( + catalog_session, + policyengine_version="5.0.5", + current_law_id=22, + time_periods=[2041, 2040], + ) + service = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + default_variables = service.list_variables("us") + selected_variables = service.list_variables("us", "5.0.5") + selected_options = service.get_economy_options("us", "5.0.5") + + assert default_variables.policyengine_version == POLICYENGINE_VERSION + assert selected_variables.policyengine_version == "5.0.5" + assert selected_options.policyengine_version == "5.0.5" + assert selected_options.current_law_id == 22 + assert [item.name for item in selected_options.time_period] == [2041, 2040] + + +def test_economy_options_reads_only_regions_and_the_national_dataset( + catalog_session: Session, +) -> None: + statements: list[str] = [] + + def record_statement(_connection, _cursor, statement, *_args) -> None: + statements.append(statement.lower()) + + bind = catalog_session.get_bind() + event.listen(bind, "before_cursor_execute", record_statement) + try: + result = V2MetadataQueryService( + catalog_session, + running_policyengine_version=POLICYENGINE_VERSION, + ).get_economy_options("us") + finally: + event.remove(bind, "before_cursor_execute", record_statement) + + assert len(statements) == 3 + assert all("variables" not in statement for statement in statements) + assert all("parameters" not in statement for statement in statements) + assert [dataset.label for dataset in result.datasets] == ["Microcosm"] diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py index ff7c892ed..ea6b6fc8d 100644 --- a/tests/unit/v2/test_metadata_routes.py +++ b/tests/unit/v2/test_metadata_routes.py @@ -11,17 +11,28 @@ from policyengine_api.asgi_factory import create_asgi_app from policyengine_api.data.v2.catalog.query import ( + InvalidMetadataPageError, InvalidPolicyEngineVersionError, MetadataCatalogUnavailableError, MetadataCatalogVersionNotFoundError, + MetadataResourceNotFoundError, + UnsupportedPreviewCountryError, ) from policyengine_api.data.v2.catalog.schemas import ( + MetadataCanonicalParameterValue, MetadataDataset, + MetadataDetailResult, MetadataEconomyOptions, + MetadataEconomyOptionsResult, MetadataModel, + MetadataModelSelectionResult, MetadataModelVersion, + MetadataModelVersionDetail, + MetadataPageResult, MetadataParameter, + MetadataParameterChild, MetadataParameterNode, + MetadataParameterSummary, MetadataParameterValue, MetadataRegion, MetadataRegionOption, @@ -67,6 +78,34 @@ def close(self) -> None: raise self.close_error +class ResourceReader: + def __init__( + self, + results: dict[str, object], + *, + error: Exception | None = None, + ): + self.results = results + self.error = error + self.calls: list[tuple[str, tuple, dict]] = [] + self.closed = False + + def __getattr__(self, name: str): + if name not in self.results: + raise AttributeError(name) + + def read(*args, **kwargs): + self.calls.append((name, args, kwargs)) + if self.error is not None: + raise self.error + return self.results[name] + + return read + + def close(self) -> None: + self.closed = True + + def _result(country_id: str = "us") -> MetadataResult: model_id = uuid4() version_id = uuid4() @@ -164,6 +203,74 @@ def _result(country_id: str = "us") -> MetadataResult: ) +def _resource_results(country_id: str = "us") -> dict[str, object]: + combined = _result(country_id) + version = combined.model_version.version + model_version = MetadataModelVersionDetail( + **combined.model_version.model_dump(), + current_law_id=combined.current_law_id, + metadata_time_periods=[2026], + ) + parameter = combined.parameters[0] + parameter_summary = MetadataParameterSummary( + **parameter.model_dump(exclude={"values"}) + ) + parameter_value = MetadataCanonicalParameterValue( + parameter_id=parameter.id, + **parameter.values[0].model_dump(), + ) + + def page(items: list[object]) -> MetadataPageResult: + return MetadataPageResult( + policyengine_version=version, + items=items, + offset=0, + limit=100, + has_more=False, + ) + + def detail(item: object) -> MetadataDetailResult: + return MetadataDetailResult(policyengine_version=version, item=item) + + return { + "list_models": page([combined.model]), + "get_model": detail(combined.model), + "get_model_by_country": MetadataModelSelectionResult( + policyengine_version=version, + model=combined.model, + model_version=model_version, + ), + "list_model_versions": page([model_version]), + "get_model_version": detail(model_version), + "list_variables": page(combined.variables), + "get_variable": detail(combined.variables[0]), + "list_parameters": page([parameter_summary]), + "list_parameter_children": page( + [ + MetadataParameterChild( + path=parameter_summary.name, + label=parameter_summary.label or parameter_summary.name, + type="parameter", + parameter=parameter_summary, + ) + ] + ), + "get_parameter": detail(parameter_summary), + "list_parameter_values": page([parameter_value]), + "get_parameter_value": detail(parameter_value), + "list_datasets": page(combined.datasets), + "get_dataset": detail(combined.datasets[0]), + "list_regions": page(combined.regions), + "get_region_by_code": detail(combined.regions[0]), + "get_region": detail(combined.regions[0]), + "get_economy_options": MetadataEconomyOptionsResult( + policyengine_version=version, + current_law_id=combined.current_law_id, + **combined.economy_options.model_dump(), + ), + } + + def _client(factory) -> TestClient: flask_app = Flask(__name__) @@ -455,3 +562,144 @@ def test_openapi_references_explicit_preview_response_schemas() -> None: ] == "#/components/schemas/MetadataErrorResponse" ) + + +def test_each_split_resource_route_returns_its_typed_result() -> None: + reader = ResourceReader(_resource_results()) + client = _client(lambda: reader) + combined = _result() + parameter = combined.parameters[0] + requests = [ + ("list_models", "/v2/tax-benefit-models?country_id=us"), + ("get_model_by_country", "/v2/tax-benefit-models/by-country/us"), + ( + "get_model", + f"/v2/tax-benefit-models/{combined.model.id}?country_id=us", + ), + ( + "list_model_versions", + "/v2/tax-benefit-model-versions?country_id=us", + ), + ( + "get_model_version", + f"/v2/tax-benefit-model-versions/{combined.model_version.id}?country_id=us", + ), + ("list_variables", "/v2/variables?country_id=us"), + ( + "get_variable", + f"/v2/variables/{combined.variables[0].id}?country_id=us", + ), + ("list_parameters", "/v2/parameters?country_id=us"), + ( + "list_parameter_children", + "/v2/parameters/children?country_id=us&parent_path=gov.example", + ), + ("get_parameter", f"/v2/parameters/{parameter.id}?country_id=us"), + ("list_parameter_values", "/v2/parameter-values?country_id=us"), + ( + "get_parameter_value", + f"/v2/parameter-values/{parameter.values[0].id}?country_id=us", + ), + ("list_datasets", "/v2/datasets?country_id=us"), + ( + "get_dataset", + f"/v2/datasets/{combined.datasets[0].id}?country_id=us", + ), + ("list_regions", "/v2/regions?country_id=us"), + ( + "get_region_by_code", + f"/v2/regions/by-code/{combined.regions[0].code}?country_id=us", + ), + ( + "get_region", + f"/v2/regions/{combined.regions[0].id}?country_id=us", + ), + ("get_economy_options", "/v2/economy-options?country_id=us"), + ] + + for expected_method, path in requests: + response = client.get(path) + assert response.status_code == 200, (path, response.text) + payload = response.json() + assert payload["status"] == "ok" + assert payload["message"] is None + assert payload["result"]["policyengine_version"] == "4.20.3" + assert reader.calls[-1][0] == expected_method + + assert reader.closed + + +def test_split_route_forwards_version_filters_and_pagination() -> None: + reader = ResourceReader(_resource_results()) + response = _client(lambda: reader).get( + "/v2/variables", + params={ + "country_id": "uk", + "policyengine_version": "4.19.0", + "offset": 25, + "limit": 50, + "search": "income", + }, + ) + + assert response.status_code == 200 + assert reader.calls == [ + ( + "list_variables", + ("uk", "4.19.0"), + {"offset": 25, "limit": 50, "search": "income"}, + ) + ] + assert reader.closed + + +@pytest.mark.parametrize( + "params", + [ + {}, + {"country_id": "us", "offset": -1}, + {"country_id": "us", "limit": 0}, + {"country_id": "us", "limit": 501}, + ], +) +def test_split_route_validation_failures_use_the_error_schema(params: dict) -> None: + calls = [] + + def factory(): + calls.append("called") + return ResourceReader(_resource_results()) + + response = _client(factory).get("/v2/variables", params=params) + + assert response.status_code == 422 + assert response.json() == { + "status": "error", + "message": "Invalid v2 metadata request", + } + assert calls == [] + + +@pytest.mark.parametrize( + ("error", "expected_status"), + [ + (InvalidMetadataPageError("invalid page"), 400), + (InvalidPolicyEngineVersionError("invalid version"), 400), + (UnsupportedPreviewCountryError("ca"), 400), + (MetadataResourceNotFoundError("missing variable"), 404), + (MetadataCatalogVersionNotFoundError("missing version"), 404), + (MetadataCatalogUnavailableError("unavailable"), 503), + (RuntimeError("private query detail"), 500), + ], +) +def test_split_route_query_failures_use_documented_error_statuses( + error: Exception, + expected_status: int, +) -> None: + reader = ResourceReader(_resource_results(), error=error) + response = _client(lambda: reader).get("/v2/variables?country_id=us") + + assert response.status_code == expected_status + assert response.json()["status"] == "error" + assert response.json()["message"] + assert "private query detail" not in response.text + assert reader.closed From 1c17e10204672ff8ad57dd5aaf71cbf314e18043 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:48:12 +0400 Subject: [PATCH 17/27] Remove combined v2 metadata response --- .../data/v2/catalog/catalog_selection.py | 119 +++ .../data/v2/catalog/parameter_tree_query.py | 131 +++ policyengine_api/data/v2/catalog/query.py | 498 +----------- policyengine_api/data/v2/catalog/schemas.py | 56 +- .../fastapi_routes/dependencies.py | 15 +- .../fastapi_routes/v2_metadata.py | 117 +-- tests/unit/v2/test_metadata_query.py | 735 +++++------------ tests/unit/v2/test_metadata_routes.py | 748 +++++++----------- 8 files changed, 790 insertions(+), 1629 deletions(-) create mode 100644 policyengine_api/data/v2/catalog/catalog_selection.py create mode 100644 policyengine_api/data/v2/catalog/parameter_tree_query.py diff --git a/policyengine_api/data/v2/catalog/catalog_selection.py b/policyengine_api/data/v2/catalog/catalog_selection.py new file mode 100644 index 000000000..7014210c9 --- /dev/null +++ b/policyengine_api/data/v2/catalog/catalog_selection.py @@ -0,0 +1,119 @@ +"""PolicyEngine.py catalog version selection for v2 resource reads.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from packaging.version import InvalidVersion, Version +from sqlalchemy.exc import SQLAlchemyError +from sqlmodel import Session, select + +from policyengine_api.data.v2.models import ( + TaxBenefitModel, + TaxBenefitModelVersion, +) + + +SUPPORTED_PREVIEW_COUNTRIES = frozenset({"us", "uk"}) + + +class MetadataCatalogUnavailableError(RuntimeError): + """Raised when an initialized catalog cannot be read.""" + + +class UnsupportedPreviewCountryError(ValueError): + """Raised when a country has no Stage 9 resource catalog.""" + + +class InvalidPolicyEngineVersionError(ValueError): + """Raised when an explicit version is not a canonical package version.""" + + +class MetadataCatalogVersionNotFoundError(LookupError): + """Raised when an explicitly selected catalog version is absent.""" + + +@dataclass(frozen=True) +class SelectedCatalog: + """One country catalog selected by its canonical PolicyEngine.py version.""" + + country_id: str + policyengine_version: str + model: TaxBenefitModel + model_version: TaxBenefitModelVersion + + +def validate_policyengine_version(value: str) -> str: + """Return one bounded canonical PEP 440 version string.""" + + if not isinstance(value, str) or not value or value != value.strip(): + raise InvalidPolicyEngineVersionError( + "policyengine_version must be a non-empty canonical version" + ) + if len(value) > 128: + raise InvalidPolicyEngineVersionError( + "policyengine_version must be at most 128 characters" + ) + try: + parsed = Version(value) + except InvalidVersion as error: + raise InvalidPolicyEngineVersionError( + "policyengine_version must be a canonical PEP 440 version" + ) from error + if str(parsed) != value or parsed == Version("0.0.0"): + raise InvalidPolicyEngineVersionError( + "policyengine_version must be a canonical non-placeholder version" + ) + return value + + +def select_catalog( + session: Session, + *, + country_id: str, + running_policyengine_version: str, + policyengine_version: str | None = None, +) -> SelectedCatalog: + """Select one country catalog using the requested or running package version.""" + + if country_id not in SUPPORTED_PREVIEW_COUNTRIES: + raise UnsupportedPreviewCountryError(country_id) + explicit_version = policyengine_version is not None + selected_version = ( + validate_policyengine_version(policyengine_version) + if explicit_version + else running_policyengine_version + ) + try: + row = session.exec( + select(TaxBenefitModel, TaxBenefitModelVersion) + .join( + TaxBenefitModelVersion, + TaxBenefitModelVersion.model_id == TaxBenefitModel.id, + ) + .where( + TaxBenefitModel.name == f"policyengine-{country_id}", + TaxBenefitModelVersion.version == selected_version, + ) + ).one_or_none() + except SQLAlchemyError as error: + raise MetadataCatalogUnavailableError( + "the v2 metadata catalog cannot be queried" + ) from error + + if row is None: + if explicit_version: + raise MetadataCatalogVersionNotFoundError( + f"PolicyEngine.py {selected_version} is not published for {country_id}" + ) + raise MetadataCatalogUnavailableError( + f"the running PolicyEngine.py {selected_version} catalog " + f"is absent for {country_id}" + ) + model, model_version = row + return SelectedCatalog( + country_id=country_id, + policyengine_version=selected_version, + model=model, + model_version=model_version, + ) diff --git a/policyengine_api/data/v2/catalog/parameter_tree_query.py b/policyengine_api/data/v2/catalog/parameter_tree_query.py new file mode 100644 index 000000000..5d2c25eed --- /dev/null +++ b/policyengine_api/data/v2/catalog/parameter_tree_query.py @@ -0,0 +1,131 @@ +"""SQL query and row conversion for direct parameter-tree children.""" + +from __future__ import annotations + +from uuid import UUID + +import sqlalchemy as sa +from sqlmodel import select + +from policyengine_api.data.v2.catalog.schemas import ( + MetadataParameterChild, + MetadataParameterSummary, +) +from policyengine_api.data.v2.models import Parameter, ParameterNode + + +def _escaped_like(value: str) -> str: + return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + + +def _child_path(column: object, prefix: str, dialect: str) -> object: + remainder = sa.func.substr(column, len(prefix) + 1) + dot_position = ( + sa.func.instr(remainder, ".") + if dialect == "sqlite" + else sa.func.strpos(remainder, ".") + ) + segment = sa.case( + (dot_position > 0, sa.func.substr(remainder, 1, dot_position - 1)), + else_=remainder, + ) + return sa.literal(prefix) + segment + + +def parameter_children_query( + *, + model_version_id: UUID, + parent_path: str, + dialect: str, + offset: int, + limit: int, +) -> object: + """Build one bounded query for a parameter path's direct children.""" + + prefix = f"{parent_path}." if parent_path else "" + escaped_prefix = _escaped_like(prefix) + node_child_path = _child_path(ParameterNode.name, prefix, dialect) + parameter_child_path = _child_path(Parameter.name, prefix, dialect) + paths = sa.union( + select(node_child_path.label("path")).where( + ParameterNode.tax_benefit_model_version_id == model_version_id, + ParameterNode.name.like(f"{escaped_prefix}%", escape="\\"), + ), + select(parameter_child_path.label("path")).where( + Parameter.tax_benefit_model_version_id == model_version_id, + Parameter.name.like(f"{escaped_prefix}%", escape="\\"), + ), + ).subquery() + descendant_count = ( + select(sa.func.count(Parameter.id)) + .where( + Parameter.tax_benefit_model_version_id == model_version_id, + sa.func.substr( + Parameter.name, + 1, + sa.func.length(paths.c.path) + 1, + ) + == paths.c.path + ".", + ) + .correlate(paths) + .scalar_subquery() + ) + is_node = sa.or_(descendant_count > 0, Parameter.id.is_(None)) + return ( + select( + paths.c.path, + sa.func.coalesce(ParameterNode.label, Parameter.label).label("label"), + sa.case((is_node, "node"), else_="parameter").label("type"), + sa.case((is_node, descendant_count), else_=None).label("child_count"), + Parameter.id.label("parameter_id"), + Parameter.label.label("parameter_label"), + Parameter.description.label("parameter_description"), + Parameter.data_type.label("parameter_data_type"), + Parameter.unit.label("parameter_unit"), + ) + .select_from( + paths.outerjoin( + ParameterNode, + sa.and_( + ParameterNode.name == paths.c.path, + ParameterNode.tax_benefit_model_version_id == model_version_id, + ), + ).outerjoin( + Parameter, + sa.and_( + Parameter.name == paths.c.path, + Parameter.tax_benefit_model_version_id == model_version_id, + ), + ) + ) + .order_by(paths.c.path) + .offset(offset) + .limit(limit + 1) + ) + + +def parameter_children_from_rows(rows: list) -> list[MetadataParameterChild]: + """Convert parameter-tree query rows into typed direct-child records.""" + + items = [] + for row in rows: + parameter = None + if row.type == "parameter": + parameter = MetadataParameterSummary( + id=row.parameter_id, + name=row.path, + label=row.parameter_label, + description=row.parameter_description, + data_type=row.parameter_data_type, + unit=row.parameter_unit, + ) + items.append( + MetadataParameterChild( + path=row.path, + label=row.label or row.path.rsplit(".", 1)[-1], + type=row.type, + child_count=row.child_count, + parameter=parameter, + ) + ) + return items diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index 700a9d5cf..8e51cbd14 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -2,70 +2,65 @@ from __future__ import annotations -from collections import defaultdict -from dataclasses import dataclass from datetime import datetime, timezone from typing import TypeVar from uuid import UUID -from packaging.version import InvalidVersion, Version import sqlalchemy as sa from sqlalchemy.exc import SQLAlchemyError from sqlmodel import Session, select from policyengine_api.dataset_display import get_dataset_display_label +from policyengine_api.data.v2.catalog.catalog_selection import ( + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, + SelectedCatalog, + UnsupportedPreviewCountryError, + select_catalog, + validate_policyengine_version, +) +from policyengine_api.data.v2.catalog.parameter_tree_query import ( + parameter_children_from_rows, + parameter_children_query, +) from policyengine_api.data.v2.catalog.schemas import ( MetadataCanonicalParameterValue, MetadataDataset, MetadataDatasetOption, MetadataDetailResult, - MetadataEconomyOptions, MetadataEconomyOptionsResult, MetadataModel, - MetadataModelVersion, MetadataModelSelectionResult, MetadataModelVersionDetail, MetadataPageResult, - MetadataParameter, MetadataParameterChild, - MetadataParameterNode, MetadataParameterSummary, - MetadataParameterValue, MetadataRegion, MetadataRegionOption, - MetadataResult, MetadataTimePeriodOption, MetadataVariable, ) from policyengine_api.data.v2.models import ( Dataset, Parameter, - ParameterNode, ParameterValue, Region, - TaxBenefitModel, - TaxBenefitModelVersion, Variable, ) -SUPPORTED_PREVIEW_COUNTRIES = frozenset({"us", "uk"}) - - -class MetadataCatalogUnavailableError(RuntimeError): - """Raised when a complete initialized catalog cannot be read.""" - - -class UnsupportedPreviewCountryError(ValueError): - """Raised when a country has no Stage 9 preview catalog.""" - - -class InvalidPolicyEngineVersionError(ValueError): - """Raised when an explicit version is not a canonical package version.""" - - -class MetadataCatalogVersionNotFoundError(LookupError): - """Raised when an explicitly selected catalog version is absent.""" +__all__ = [ + "InvalidMetadataPageError", + "InvalidPolicyEngineVersionError", + "MetadataCatalogUnavailableError", + "MetadataCatalogVersionNotFoundError", + "MetadataResourceNotFoundError", + "UnsupportedPreviewCountryError", + "V2MetadataQueryService", + "validate_metadata_page", + "validate_policyengine_version", +] class MetadataResourceNotFoundError(LookupError): @@ -76,16 +71,6 @@ class InvalidMetadataPageError(ValueError): """Raised when collection pagination is outside the documented bounds.""" -@dataclass(frozen=True) -class SelectedCatalog: - """One country catalog selected by its canonical PolicyEngine.py version.""" - - country_id: str - policyengine_version: str - model: TaxBenefitModel - model_version: TaxBenefitModelVersion - - ResourceT = TypeVar("ResourceT") @@ -200,44 +185,6 @@ def _escaped_like(value: str) -> str: return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") -def _child_path_expression(column: object, prefix: str, dialect: str) -> object: - remainder = sa.func.substr(column, len(prefix) + 1) - dot_position = ( - sa.func.instr(remainder, ".") - if dialect == "sqlite" - else sa.func.strpos(remainder, ".") - ) - segment = sa.case( - (dot_position > 0, sa.func.substr(remainder, 1, dot_position - 1)), - else_=remainder, - ) - return sa.literal(prefix) + segment - - -def validate_policyengine_version(value: str) -> str: - """Return one bounded canonical PEP 440 version string.""" - - if not isinstance(value, str) or not value or value != value.strip(): - raise InvalidPolicyEngineVersionError( - "policyengine_version must be a non-empty canonical version" - ) - if len(value) > 128: - raise InvalidPolicyEngineVersionError( - "policyengine_version must be at most 128 characters" - ) - try: - parsed = Version(value) - except InvalidVersion as error: - raise InvalidPolicyEngineVersionError( - "policyengine_version must be a canonical PEP 440 version" - ) from error - if str(parsed) != value or parsed == Version("0.0.0"): - raise InvalidPolicyEngineVersionError( - "policyengine_version must be a canonical non-placeholder version" - ) - return value - - class V2MetadataQueryService: """Assemble preview metadata using only an injected v2 read session.""" @@ -259,47 +206,11 @@ def select_catalog( ) -> SelectedCatalog: """Select exactly one initialized country catalog.""" - if country_id not in SUPPORTED_PREVIEW_COUNTRIES: - raise UnsupportedPreviewCountryError(country_id) - explicit_version = policyengine_version is not None - selected_version = ( - validate_policyengine_version(policyengine_version) - if explicit_version - else self._running_policyengine_version - ) - try: - row = self._session.exec( - select(TaxBenefitModel, TaxBenefitModelVersion) - .join( - TaxBenefitModelVersion, - TaxBenefitModelVersion.model_id == TaxBenefitModel.id, - ) - .where( - TaxBenefitModel.name == f"policyengine-{country_id}", - TaxBenefitModelVersion.version == selected_version, - ) - ).one_or_none() - except SQLAlchemyError as error: - raise MetadataCatalogUnavailableError( - "the v2 metadata catalog cannot be queried" - ) from error - - if row is None: - if explicit_version: - raise MetadataCatalogVersionNotFoundError( - f"PolicyEngine.py {selected_version} is not published " - f"for {country_id}" - ) - raise MetadataCatalogUnavailableError( - f"the running PolicyEngine.py {selected_version} catalog " - f"is absent for {country_id}" - ) - model, model_version = row - return SelectedCatalog( + return select_catalog( + self._session, country_id=country_id, - policyengine_version=selected_version, - model=model, - model_version=model_version, + running_policyengine_version=self._running_policyengine_version, + policyengine_version=policyengine_version, ) def _resource_rows(self, statement: object) -> list: @@ -497,95 +408,21 @@ def list_parameter_children( ) -> MetadataPageResult[MetadataParameterChild]: validate_metadata_page(offset, limit) selected = self.select_catalog(country_id, policyengine_version) - version_id = selected.model_version.id - prefix = f"{parent_path}." if parent_path else "" - escaped_prefix = _escaped_like(prefix) - dialect = self._session.get_bind().dialect.name - node_child_path = _child_path_expression(ParameterNode.name, prefix, dialect) - parameter_child_path = _child_path_expression( - Parameter.name, - prefix, - dialect, - ) - paths = sa.union( - select(node_child_path.label("path")).where( - ParameterNode.tax_benefit_model_version_id == version_id, - ParameterNode.name.like(f"{escaped_prefix}%", escape="\\"), - ), - select(parameter_child_path.label("path")).where( - Parameter.tax_benefit_model_version_id == version_id, - Parameter.name.like(f"{escaped_prefix}%", escape="\\"), - ), - ).subquery() - descendant_count = ( - select(sa.func.count(Parameter.id)) - .where( - Parameter.tax_benefit_model_version_id == version_id, - sa.func.substr( - Parameter.name, - 1, - sa.func.length(paths.c.path) + 1, - ) - == paths.c.path + ".", + rows = self._resource_rows( + parameter_children_query( + model_version_id=selected.model_version.id, + parent_path=parent_path, + dialect=self._session.get_bind().dialect.name, + offset=offset, + limit=limit, ) - .correlate(paths) - .scalar_subquery() ) - is_node = sa.or_(descendant_count > 0, Parameter.id.is_(None)) - statement = ( - select( - paths.c.path, - sa.func.coalesce(ParameterNode.label, Parameter.label).label("label"), - sa.case((is_node, "node"), else_="parameter").label("type"), - sa.case((is_node, descendant_count), else_=None).label("child_count"), - Parameter.id.label("parameter_id"), - Parameter.label.label("parameter_label"), - Parameter.description.label("parameter_description"), - Parameter.data_type.label("parameter_data_type"), - Parameter.unit.label("parameter_unit"), - ) - .select_from( - paths.outerjoin( - ParameterNode, - sa.and_( - ParameterNode.name == paths.c.path, - ParameterNode.tax_benefit_model_version_id == version_id, - ), - ).outerjoin( - Parameter, - sa.and_( - Parameter.name == paths.c.path, - Parameter.tax_benefit_model_version_id == version_id, - ), - ) - ) - .order_by(paths.c.path) - .offset(offset) - .limit(limit + 1) + return _page( + selected, + parameter_children_from_rows(rows), + offset=offset, + limit=limit, ) - rows = self._resource_rows(statement) - items = [] - for row in rows: - parameter = None - if row.type == "parameter": - parameter = MetadataParameterSummary( - id=row.parameter_id, - name=row.path, - label=row.parameter_label, - description=row.parameter_description, - data_type=row.parameter_data_type, - unit=row.parameter_unit, - ) - items.append( - MetadataParameterChild( - path=row.path, - label=row.label or row.path.rsplit(".", 1)[-1], - type=row.type, - child_count=row.child_count, - parameter=parameter, - ) - ) - return _page(selected, items, offset=offset, limit=limit) def list_parameter_values( self, @@ -842,258 +679,3 @@ def get_economy_options( ) ], ) - - def get_metadata( - self, - country_id: str, - policyengine_version: str | None = None, - ) -> MetadataResult: - if country_id not in SUPPORTED_PREVIEW_COUNTRIES: - raise UnsupportedPreviewCountryError(country_id) - explicit_version = policyengine_version is not None - selected_version = ( - validate_policyengine_version(policyengine_version) - if explicit_version - else self._running_policyengine_version - ) - try: - return self._read_metadata( - country_id, - selected_version, - explicit_version=explicit_version, - ) - except ( - MetadataCatalogUnavailableError, - MetadataCatalogVersionNotFoundError, - ): - raise - except SQLAlchemyError as error: - raise MetadataCatalogUnavailableError( - "the v2 metadata catalog cannot be queried" - ) from error - - def _read_metadata( - self, - country_id: str, - policyengine_version: str, - *, - explicit_version: bool, - ) -> MetadataResult: - model = self._session.exec( - select(TaxBenefitModel).where( - TaxBenefitModel.name == f"policyengine-{country_id}" - ) - ).one_or_none() - if model is None: - if explicit_version: - raise MetadataCatalogVersionNotFoundError( - f"PolicyEngine.py {policyengine_version} is not published " - f"for {country_id}" - ) - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 metadata catalog is not initialized" - ) - - model_version = self._session.exec( - select(TaxBenefitModelVersion).where( - TaxBenefitModelVersion.model_id == model.id, - TaxBenefitModelVersion.version == policyengine_version, - ) - ).one_or_none() - if model_version is None: - if explicit_version: - raise MetadataCatalogVersionNotFoundError( - f"PolicyEngine.py {policyengine_version} is not published " - f"for {country_id}" - ) - raise MetadataCatalogUnavailableError( - f"the running PolicyEngine.py {policyengine_version} catalog " - f"is absent for {country_id}" - ) - - variables = self._session.exec( - select(Variable) - .where(Variable.tax_benefit_model_version_id == model_version.id) - .order_by(Variable.name) - ).all() - nodes = self._session.exec( - select(ParameterNode) - .where(ParameterNode.tax_benefit_model_version_id == model_version.id) - .order_by(ParameterNode.name) - ).all() - parameters = self._session.exec( - select(Parameter) - .where(Parameter.tax_benefit_model_version_id == model_version.id) - .order_by(Parameter.name) - ).all() - parameter_values = self._session.exec( - select(ParameterValue) - .join(Parameter, Parameter.id == ParameterValue.parameter_id) - .where( - Parameter.tax_benefit_model_version_id == model_version.id, - ParameterValue.policy_id.is_(None), - ParameterValue.dynamic_id.is_(None), - ) - .order_by(Parameter.name, ParameterValue.start_date) - ).all() - regions = self._session.exec( - select(Region) - .where(Region.tax_benefit_model_version_id == model_version.id) - .order_by(Region.code) - ).all() - - if not variables or not nodes or not parameters or not regions: - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 metadata catalog is incomplete" - ) - if {value.parameter_id for value in parameter_values} != { - parameter.id for parameter in parameters - }: - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 parameter values are incomplete" - ) - - dataset_ids = {region.default_dataset_id for region in regions} - datasets = self._session.exec( - select(Dataset) - .where( - Dataset.id.in_(dataset_ids), - Dataset.tax_benefit_model_version_id == model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ) - .order_by(Dataset.name) - ).all() - if {dataset.id for dataset in datasets} != dataset_ids: - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 region datasets are incomplete" - ) - - values_by_parameter = defaultdict(list) - for value in parameter_values: - values_by_parameter[value.parameter_id].append( - MetadataParameterValue( - id=value.id, - value=value.value_json, - start_date=value.start_date, - end_date=value.end_date, - ) - ) - - national_region = next( - (region for region in regions if region.code == country_id), - None, - ) - if national_region is None: - raise MetadataCatalogUnavailableError( - f"the {country_id} national v2 region is absent" - ) - datasets_by_id = {dataset.id: dataset for dataset in datasets} - national_dataset = datasets_by_id[national_region.default_dataset_id] - time_periods = model_version.metadata_time_periods - if ( - not isinstance(model_version.current_law_id, int) - or not isinstance(time_periods, list) - or not time_periods - or any(not isinstance(year, int) for year in time_periods) - ): - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 model-version options are incomplete" - ) - - return MetadataResult( - current_law_id=model_version.current_law_id, - model=MetadataModel( - id=model.id, - name=model.name, - description=model_version.description, - ), - model_version=MetadataModelVersion( - id=model_version.id, - model_id=model.id, - version=model_version.version, - description=model_version.description, - ), - variables=[ - MetadataVariable( - id=variable.id, - name=variable.name, - label=variable.label, - entity=variable.entity, - description=variable.description, - data_type=variable.data_type, - possible_values=variable.possible_values, - default_value=variable.default_value, - adds=variable.adds, - subtracts=variable.subtracts, - ) - for variable in variables - ], - parameter_nodes=[ - MetadataParameterNode( - id=node.id, - name=node.name, - label=node.label, - description=node.description, - ) - for node in nodes - ], - parameters=[ - MetadataParameter( - id=parameter.id, - name=parameter.name, - label=parameter.label, - description=parameter.description, - data_type=parameter.data_type, - unit=parameter.unit, - values=values_by_parameter[parameter.id], - ) - for parameter in parameters - ], - datasets=[ - MetadataDataset( - id=dataset.id, - name=dataset.name, - description=dataset.description, - year=dataset.year, - ) - for dataset in datasets - ], - regions=[ - MetadataRegion( - id=region.id, - code=region.code, - label=region.label, - region_type=region.region_type.value, - requires_filter=region.requires_filter, - filter_field=region.filter_field, - filter_value=region.filter_value, - filter_strategy=region.filter_strategy, - parent_code=region.parent_code, - state_code=region.state_code, - state_name=region.state_name, - default_dataset_id=region.default_dataset_id, - ) - for region in regions - ], - economy_options=MetadataEconomyOptions( - region=[ - MetadataRegionOption( - name=region.code, - label=region.label, - type=region.region_type.value, - ) - for region in regions - ], - time_period=[ - MetadataTimePeriodOption(name=year, label=str(year)) - for year in time_periods - ], - datasets=[ - MetadataDatasetOption( - name=national_dataset.name, - label=get_dataset_display_label(national_dataset.name), - ) - ], - ), - ) diff --git a/policyengine_api/data/v2/catalog/schemas.py b/policyengine_api/data/v2/catalog/schemas.py index ec4cc3ee8..ddd01fb8e 100644 --- a/policyengine_api/data/v2/catalog/schemas.py +++ b/policyengine_api/data/v2/catalog/schemas.py @@ -7,7 +7,7 @@ from typing import Annotated, Generic, Literal, TypeVar from uuid import UUID -from pydantic import BaseModel, ConfigDict, Field, JsonValue, StringConstraints +from pydantic import BaseModel, ConfigDict, JsonValue, StringConstraints class StrictResponseModel(BaseModel): @@ -51,30 +51,6 @@ class MetadataVariable(StrictResponseModel): subtracts: list[str] | None -class MetadataParameterNode(StrictResponseModel): - id: UUID - name: str - label: str | None - description: str | None - - -class MetadataParameterValue(StrictResponseModel): - id: UUID - value: JsonValue - start_date: datetime - end_date: datetime | None - - -class MetadataParameter(StrictResponseModel): - id: UUID - name: str - label: str | None - description: str | None - data_type: str | None - unit: str | None - values: list[MetadataParameterValue] - - class MetadataParameterSummary(StrictResponseModel): id: UUID name: str @@ -133,12 +109,6 @@ class MetadataDatasetOption(StrictResponseModel): default: Literal[True] = True -class MetadataEconomyOptions(StrictResponseModel): - region: list[MetadataRegionOption] - time_period: list[MetadataTimePeriodOption] - datasets: list[MetadataDatasetOption] - - class MetadataModelVersionDetail(MetadataModelVersion): current_law_id: int metadata_time_periods: list[int] @@ -292,30 +262,6 @@ class MetadataEconomyOptionsResponse( pass -class MetadataResult(StrictResponseModel): - current_law_id: int - model: MetadataModel - model_version: MetadataModelVersion - variables: list[MetadataVariable] - parameter_nodes: list[MetadataParameterNode] - parameters: list[MetadataParameter] - datasets: list[MetadataDataset] - regions: list[MetadataRegion] - economy_options: MetadataEconomyOptions - - -class MetadataSuccessResponse(StrictResponseModel): - status: Literal["ok"] = "ok" - message: None = None - result: MetadataResult - - class MetadataErrorResponse(StrictResponseModel): status: Literal["error"] = "error" message: Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)] - - -MetadataPreviewResponse = Annotated[ - MetadataSuccessResponse | MetadataErrorResponse, - Field(discriminator="status"), -] diff --git a/policyengine_api/fastapi_routes/dependencies.py b/policyengine_api/fastapi_routes/dependencies.py index 5b87e8465..4b3166689 100644 --- a/policyengine_api/fastapi_routes/dependencies.py +++ b/policyengine_api/fastapi_routes/dependencies.py @@ -8,7 +8,6 @@ from importlib import metadata as importlib_metadata from typing import Protocol -from policyengine_api.data.v2.catalog.schemas import MetadataResult from policyengine_api.json_types import JSONObject @@ -18,14 +17,8 @@ class MetadataReader(Protocol): def get_metadata(self, country_id: str) -> JSONObject: ... -class V2MetadataReader(Protocol): - """Read one already-initialized v2 metadata catalog.""" - - def get_metadata( - self, - country_id: str, - policyengine_version: str | None = None, - ) -> MetadataResult: ... +class V2MetadataResourceReader(Protocol): + """Own the database session used by one v2 metadata resource request.""" def close(self) -> None: ... @@ -65,7 +58,7 @@ def _running_policyengine_version() -> str: return importlib_metadata.version("policyengine") -def _default_v2_metadata_reader_factory() -> V2MetadataReader: +def _default_v2_metadata_reader_factory() -> V2MetadataResourceReader: from policyengine_api.data.v2.catalog.query import V2MetadataQueryService from policyengine_api.data.v2.database import get_v2_session_factory @@ -83,7 +76,7 @@ class NativeRouteDependencies: gateway_client_factory: Callable[[], SimulationGatewayProbe] metadata_reader_factory: Callable[[], MetadataReader] specification_provider: Callable[[], JSONObject] - v2_metadata_reader_factory: Callable[[], V2MetadataReader] | None = None + v2_metadata_reader_factory: Callable[[], V2MetadataResourceReader] | None = None @classmethod def defaults(cls) -> "NativeRouteDependencies": diff --git a/policyengine_api/fastapi_routes/v2_metadata.py b/policyengine_api/fastapi_routes/v2_metadata.py index 75b290a6e..f0ff4f7fc 100644 --- a/policyengine_api/fastapi_routes/v2_metadata.py +++ b/policyengine_api/fastapi_routes/v2_metadata.py @@ -1,20 +1,11 @@ -"""Dormant, read-only API v2 metadata preview routes.""" +"""Dormant, read-only API v2 metadata resource routes.""" from __future__ import annotations from fastapi import APIRouter, Request from starlette.responses import JSONResponse -from policyengine_api.data.v2.catalog.query import ( - InvalidPolicyEngineVersionError, - MetadataCatalogUnavailableError, - MetadataCatalogVersionNotFoundError, -) -from policyengine_api.data.v2.catalog.schemas import ( - MetadataErrorResponse, - MetadataSuccessResponse, -) -from policyengine_api.data.v2.settings import V2ConfigurationError +from policyengine_api.data.v2.catalog.schemas import MetadataErrorResponse from policyengine_api.fastapi_routes.dependencies import NativeRouteDependencies from policyengine_api.fastapi_routes.v2_metadata_geography import ( build_v2_metadata_geography_router, @@ -27,42 +18,10 @@ ) -ERROR_RESPONSES = { - 400: { - "model": MetadataErrorResponse, - "description": "The requested PolicyEngine.py version is invalid.", - }, - 404: { - "model": MetadataErrorResponse, - "description": "The requested PolicyEngine.py catalog is absent.", - }, - 405: { - "model": MetadataErrorResponse, - "description": "The preview supports GET only.", - }, - 503: { - "model": MetadataErrorResponse, - "description": "The initialized v2 catalog is unavailable.", - }, - 500: { - "model": MetadataErrorResponse, - "description": "The preview query failed internally.", - }, -} - - -def _error_response(status_code: int, message: str) -> JSONResponse: - error = MetadataErrorResponse(message=message) - return JSONResponse( - status_code=status_code, - content=error.model_dump(mode="json"), - ) - - def build_v2_metadata_router( dependencies: NativeRouteDependencies, ) -> APIRouter: - """Build isolated preview routes without loading v2 configuration.""" + """Build isolated resource routes without loading v2 configuration.""" router = APIRouter() router.include_router(build_v2_metadata_model_router(dependencies)) @@ -72,7 +31,7 @@ def build_v2_metadata_router( @router.get( "/v2/openapi.json", include_in_schema=False, - summary="OpenAPI document for dormant v2 preview routes", + summary="OpenAPI document for dormant v2 metadata resources", ) def v2_preview_openapi(request: Request) -> JSONResponse: schema = request.app.openapi() @@ -86,81 +45,27 @@ def v2_preview_openapi(request: Request) -> JSONResponse: } return JSONResponse(preview_schema) - def read( - country_id: str, - policyengine_version: str | None, - ) -> MetadataSuccessResponse | JSONResponse: - reader = None - try: - factory = dependencies.v2_metadata_reader_factory - if factory is None: - from policyengine_api.fastapi_routes.dependencies import ( - _default_v2_metadata_reader_factory, - ) - - factory = _default_v2_metadata_reader_factory - reader = factory() - result = reader.get_metadata(country_id, policyengine_version) - return MetadataSuccessResponse(result=result) - except InvalidPolicyEngineVersionError as error: - return _error_response(400, str(error)) - except MetadataCatalogVersionNotFoundError as error: - return _error_response(404, str(error)) - except (V2ConfigurationError, MetadataCatalogUnavailableError): - return _error_response(503, "V2 metadata catalog is unavailable") - except Exception: # noqa: BLE001 - preview must return typed errors - return _error_response(500, "V2 metadata query failed") - finally: - if reader is not None: - try: - reader.close() - except Exception: - pass - - @router.get( - "/v2/us/metadata", - response_model=MetadataSuccessResponse, - responses=ERROR_RESPONSES, - summary="Preview US metadata from the v2 catalog", - ) - def us_metadata_preview( - policyengine_version: str | None = None, - ) -> MetadataSuccessResponse | JSONResponse: - return read("us", policyengine_version) - @router.get( - "/v2/uk/metadata", - response_model=MetadataSuccessResponse, - responses=ERROR_RESPONSES, - summary="Preview UK metadata from the v2 catalog", - ) - def uk_metadata_preview( - policyengine_version: str | None = None, - ) -> MetadataSuccessResponse | JSONResponse: - return read("uk", policyengine_version) - - @router.get( - "/v2/{country_id}/metadata", + "/v2/{resource_path:path}", response_model=MetadataErrorResponse, status_code=404, - responses={500: ERROR_RESPONSES[500]}, - summary="Reject an unsupported v2 metadata preview country", + include_in_schema=False, ) - def unsupported_country(country_id: str) -> MetadataErrorResponse: + def unsupported_resource(resource_path: str) -> MetadataErrorResponse: return MetadataErrorResponse( - message=f"V2 metadata is not available for country {country_id}" + message=f"V2 metadata resource {resource_path!r} was not found" ) @router.api_route( - "/v2/{country_id}/metadata", + "/v2/{resource_path:path}", methods=["POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"], response_model=MetadataErrorResponse, status_code=405, include_in_schema=False, ) - def unsupported_method(country_id: str) -> MetadataErrorResponse: + def unsupported_method(resource_path: str) -> MetadataErrorResponse: return MetadataErrorResponse( - message=f"V2 metadata for country {country_id} supports GET only" + message=f"V2 metadata resource {resource_path!r} supports GET only" ) return router diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index b42af29b2..8d20bf6dc 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -1,4 +1,4 @@ -"""Typed read-only query coverage for the v2 metadata preview.""" +"""Read-only query coverage for v2 metadata resources.""" from __future__ import annotations @@ -7,9 +7,8 @@ from pathlib import Path from uuid import uuid4 -from pydantic import TypeAdapter, ValidationError import pytest -from sqlalchemy import delete, event +from sqlalchemy import event from sqlalchemy.pool import StaticPool from sqlmodel import Session, create_engine, select @@ -22,11 +21,6 @@ UnsupportedPreviewCountryError, V2MetadataQueryService, ) -from policyengine_api.data.v2.catalog.schemas import ( - MetadataErrorResponse, - MetadataPreviewResponse, - MetadataSuccessResponse, -) from policyengine_api.data.v2.models import ( Dataset, Parameter, @@ -42,141 +36,38 @@ from tests.fixtures.v2_catalog import POLICYENGINE_VERSION, normalized_catalog -@pytest.fixture -def catalog_session() -> Session: - engine = create_engine( - "sqlite://", - connect_args={"check_same_thread": False}, - poolclass=StaticPool, - ) - V2_METADATA.create_all(engine) - catalog = normalized_catalog() - with Session(engine) as session: - for country in catalog.countries: - session.add( - TaxBenefitModel( - id=country.model.id, - name=country.model.name, - description=country.model.description, - ) - ) - session.add( - TaxBenefitModelVersion( - id=country.model_version.id, - model_id=country.model.id, - version=country.model_version.version, - description=country.model_version.description, - current_law_id=country.model_version.current_law_id, - metadata_time_periods=list( - country.model_version.metadata_time_periods - ), - ) - ) - session.add_all( - Variable( - id=record.id, - tax_benefit_model_version_id=country.model_version.id, - name=record.name, - label=record.label, - entity=record.entity, - description=record.description, - data_type=record.data_type, - possible_values=record.possible_values, - default_value=record.default_value, - adds=record.adds, - subtracts=record.subtracts, - ) - for record in country.variables - ) - session.add_all( - ParameterNode( - id=record.id, - tax_benefit_model_version_id=country.model_version.id, - name=record.name, - label=record.label, - description=record.description, - ) - for record in country.parameter_nodes - ) - session.add_all( - Parameter( - id=record.id, - tax_benefit_model_version_id=country.model_version.id, - name=record.name, - label=record.label, - description=record.description, - data_type=record.data_type, - unit=record.unit, - ) - for record in country.parameters - ) - session.add_all( - ParameterValue( - id=value.id, - parameter_id=value.parameter_id, - value_json=value.value_json, - start_date=value.start_date, - end_date=value.end_date, - ) - for parameter in country.parameters - for value in parameter.values - ) - session.add_all( - Dataset( - id=record.id, - tax_benefit_model_version_id=country.model_version.id, - name=record.name, - description=record.description, - year=record.year, - ) - for record in country.datasets - ) - session.add_all( - Region( - id=record.id, - tax_benefit_model_version_id=country.model_version.id, - default_dataset_id=record.default_dataset_id, - code=record.code, - label=record.label, - region_type=RegionType(record.region_type), - requires_filter=record.requires_filter, - filter_field=record.filter_field, - filter_value=record.filter_value, - filter_strategy=record.filter_strategy, - parent_code=record.parent_code, - state_code=record.state_code, - state_name=record.state_name, - ) - for record in country.regions - ) - session.commit() - - session = Session(engine) - try: - yield session - finally: - session.close() - engine.dispose() - - -def _add_country_version( +def _insert_country( session: Session, + country, *, - policyengine_version: str, - current_law_id: int, - time_periods: list[int], + include_model: bool, + current_law_id: int | None = None, + time_periods: list[int] | None = None, ) -> None: - country = normalized_catalog(policyengine_version=policyengine_version).country( - "us" - ) + if include_model: + session.add( + TaxBenefitModel( + id=country.model.id, + name=country.model.name, + description=country.model.description, + ) + ) session.add( TaxBenefitModelVersion( id=country.model_version.id, model_id=country.model.id, - version=policyengine_version, - description=f"US model for {policyengine_version}", - current_law_id=current_law_id, - metadata_time_periods=time_periods, + version=country.model_version.version, + description=country.model_version.description, + current_law_id=( + country.model_version.current_law_id + if current_law_id is None + else current_law_id + ), + metadata_time_periods=( + list(country.model_version.metadata_time_periods) + if time_periods is None + else time_periods + ), ) ) session.add_all( @@ -233,7 +124,7 @@ def _add_country_version( id=record.id, tax_benefit_model_version_id=country.model_version.id, name=record.name, - description=f"{record.description} for {policyengine_version}", + description=record.description, year=record.year, ) for record in country.datasets @@ -244,7 +135,7 @@ def _add_country_version( tax_benefit_model_version_id=country.model_version.id, default_dataset_id=record.default_dataset_id, code=record.code, - label=f"{record.label} {policyengine_version}", + label=record.label, region_type=RegionType(record.region_type), requires_filter=record.requires_filter, filter_field=record.filter_field, @@ -256,276 +147,51 @@ def _add_country_version( ) for record in country.regions ) - session.commit() - - -def test_query_serializes_complete_typed_metadata_without_writes( - catalog_session: Session, -) -> None: - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) - model_classes = ( - TaxBenefitModel, - TaxBenefitModelVersion, - Variable, - ParameterNode, - Parameter, - ParameterValue, - Dataset, - Region, - ) - before = { - model_class.__tablename__: len(catalog_session.exec(select(model_class)).all()) - for model_class in model_classes - } - - result = service.get_metadata("us") - - assert result.current_law_id == 2 - assert result.model.name == "policyengine-us" - assert result.model_version.version == POLICYENGINE_VERSION - assert [variable.name for variable in result.variables] == ["employment_income"] - assert [parameter.name for parameter in result.parameters] == ["gov.example.rate"] - assert [value.value for value in result.parameters[0].values] == [0.1, 0.2] - assert {dataset.name for dataset in result.datasets} == { - "populace_us_2024", - "populace_us_ca_2024", - } - assert all(not dataset.is_output_dataset for dataset in result.datasets) - assert all(dataset.storage_path is None for dataset in result.datasets) - assert result.economy_options.region[0].name == "place/CA-44000" - assert result.economy_options.time_period[0].name == 2035 - assert result.economy_options.time_period[-1].name == 2022 - assert [option.name for option in result.economy_options.datasets] == [ - "populace_us_2024" - ] - assert [option.label for option in result.economy_options.datasets] == ["Microcosm"] - after = { - model_class.__tablename__: len(catalog_session.exec(select(model_class)).all()) - for model_class in model_classes - } - assert after == before - - -def test_query_excludes_output_datasets(catalog_session: Session) -> None: - model = catalog_session.exec( - select(TaxBenefitModel).where(TaxBenefitModel.name == "policyengine-us") - ).one() - catalog_session.add( - Dataset( - id=uuid4(), - tax_benefit_model_version_id=catalog_session.exec( - select(TaxBenefitModelVersion).where( - TaxBenefitModelVersion.model_id == model.id, - TaxBenefitModelVersion.version == POLICYENGINE_VERSION, - ) - ) - .one() - .id, - name="simulation-output", - description="Generated result", - storage_path="private-output-reference", - year=2026, - is_output_dataset=True, - ) - ) - catalog_session.commit() - result = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ).get_metadata("us") - - assert "simulation-output" not in {dataset.name for dataset in result.datasets} - -def test_query_defaults_to_running_version_and_allows_exact_override( - catalog_session: Session, -) -> None: - _add_country_version( - catalog_session, - policyengine_version="5.0.5", - current_law_id=22, - time_periods=[2041, 2040], - ) - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) - - default = service.get_metadata("us") - selected = service.get_metadata("us", "5.0.5") - - assert default.model_version.version == POLICYENGINE_VERSION - assert default.current_law_id == 2 - assert default.economy_options.time_period[0].name == 2035 - assert selected.model_version.version == "5.0.5" - assert selected.model.description == "US model for 5.0.5" - assert selected.current_law_id == 22 - assert [option.name for option in selected.economy_options.time_period] == [ - 2041, - 2040, - ] - assert {dataset.id for dataset in default.datasets}.isdisjoint( - dataset.id for dataset in selected.datasets - ) - assert {region.id for region in default.regions}.isdisjoint( - region.id for region in selected.regions - ) - - -def test_query_rejects_invalid_or_absent_selected_versions( - catalog_session: Session, -) -> None: - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, +@pytest.fixture +def catalog_session() -> Session: + engine = create_engine( + "sqlite://", + connect_args={"check_same_thread": False}, + poolclass=StaticPool, ) - - for invalid in ( - "", - f" {POLICYENGINE_VERSION}", - "not a version", - "0.0.0", - "1" * 129, - ): - with pytest.raises(InvalidPolicyEngineVersionError): - service.get_metadata("us", invalid) - with pytest.raises(MetadataCatalogVersionNotFoundError): - service.get_metadata("us", "4.99.0") - with pytest.raises(MetadataCatalogUnavailableError): - V2MetadataQueryService( - catalog_session, - running_policyengine_version="4.99.0", - ).get_metadata("us") - - -def test_query_distinguishes_an_uninitialized_catalog_from_an_absent_version() -> None: - engine = create_engine("sqlite://") V2_METADATA.create_all(engine) - try: - with Session(engine) as session: - service = V2MetadataQueryService( - session, - running_policyengine_version=POLICYENGINE_VERSION, - ) + with Session(engine) as session: + for country in normalized_catalog().countries: + _insert_country(session, country, include_model=True) + session.commit() - with pytest.raises( - MetadataCatalogUnavailableError, - match="not initialized", - ): - service.get_metadata("us") - with pytest.raises( - MetadataCatalogVersionNotFoundError, - match=( - f"PolicyEngine.py {POLICYENGINE_VERSION} is not published for us" - ), - ): - service.get_metadata("us", POLICYENGINE_VERSION) + session = Session(engine) + try: + yield session finally: + session.close() engine.dispose() -def test_query_rejects_a_region_whose_dataset_is_absent( - catalog_session: Session, -) -> None: - national_region = catalog_session.exec( - select(Region).where(Region.code == "us") - ).one() - catalog_session.exec( - delete(Dataset).where(Dataset.id == national_region.default_dataset_id) - ) - catalog_session.commit() - - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) - with pytest.raises( - MetadataCatalogUnavailableError, - match="region datasets are incomplete", - ): - service.get_metadata("us") - - -def test_query_requires_a_national_region(catalog_session: Session) -> None: - national_region = catalog_session.exec( - select(Region).where(Region.code == "us") - ).one() - catalog_session.delete(national_region) - catalog_session.commit() - - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) - with pytest.raises( - MetadataCatalogUnavailableError, - match="national v2 region is absent", - ): - service.get_metadata("us") - - -def test_query_requires_nonempty_integer_time_periods( - catalog_session: Session, +def _add_country_version( + session: Session, + *, + policyengine_version: str, + current_law_id: int, + time_periods: list[int], ) -> None: - us_version = catalog_session.exec( - select(TaxBenefitModelVersion) - .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) - .where( - TaxBenefitModel.name == "policyengine-us", - TaxBenefitModelVersion.version == POLICYENGINE_VERSION, - ) - ).one() - us_version.metadata_time_periods = [] - catalog_session.add(us_version) - catalog_session.commit() - - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, + country = normalized_catalog(policyengine_version=policyengine_version).country( + "us" ) - with pytest.raises( - MetadataCatalogUnavailableError, - match="model-version options are incomplete", - ): - service.get_metadata("us") - - -def test_query_rejects_unsupported_and_incomplete_catalogs( - catalog_session: Session, -) -> None: - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, + _insert_country( + session, + country, + include_model=False, + current_law_id=current_law_id, + time_periods=time_periods, ) - with pytest.raises(UnsupportedPreviewCountryError): - service.get_metadata("ca") - - uk_model = catalog_session.exec( - select(TaxBenefitModel).where(TaxBenefitModel.name == "policyengine-uk") - ).one() - uk_model_version = catalog_session.exec( - select(TaxBenefitModelVersion).where( - TaxBenefitModelVersion.model_id == uk_model.id, - TaxBenefitModelVersion.version == POLICYENGINE_VERSION, - ) - ).one() - for region in catalog_session.exec( - select(Region).where(Region.tax_benefit_model_version_id == uk_model_version.id) - ).all(): - catalog_session.delete(region) - catalog_session.commit() - with pytest.raises(MetadataCatalogUnavailableError, match="incomplete"): - service.get_metadata("uk") + session.commit() -def test_query_rejects_incomplete_parameter_values( - catalog_session: Session, -) -> None: - us_model_version = catalog_session.exec( +def _us_model_version(session: Session) -> TaxBenefitModelVersion: + return session.exec( select(TaxBenefitModelVersion) .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) .where( @@ -533,110 +199,19 @@ def test_query_rejects_incomplete_parameter_values( TaxBenefitModelVersion.version == POLICYENGINE_VERSION, ) ).one() - us_parameter_ids = set( - catalog_session.exec( - select(Parameter.id).where( - Parameter.tax_benefit_model_version_id == us_model_version.id - ) - ).all() - ) - for value in catalog_session.exec( - select(ParameterValue).where(ParameterValue.parameter_id.in_(us_parameter_ids)) - ).all(): - catalog_session.delete(value) - catalog_session.commit() - - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) - with pytest.raises( - MetadataCatalogUnavailableError, - match="parameter values are incomplete", - ): - service.get_metadata("us") -def test_response_outcomes_are_discriminated_and_strict( - catalog_session: Session, -) -> None: - result = V2MetadataQueryService( - catalog_session, +def _service(session: Session) -> V2MetadataQueryService: + return V2MetadataQueryService( + session, running_policyengine_version=POLICYENGINE_VERSION, - ).get_metadata("uk") - adapter = TypeAdapter(MetadataPreviewResponse) - - success = adapter.validate_python( - MetadataSuccessResponse(result=result).model_dump() - ) - error = adapter.validate_python( - MetadataErrorResponse(message="Catalog unavailable").model_dump() - ) - assert success.status == "ok" - assert error.status == "error" - - with pytest.raises(ValidationError): - adapter.validate_python({"status": "ok", "message": None}) - with pytest.raises(ValidationError): - adapter.validate_python({"status": "error", "message": ""}) - with pytest.raises(ValidationError): - adapter.validate_python({"status": "error", "message": " "}) - with pytest.raises(ValidationError): - adapter.validate_python( - { - "status": "error", - "message": "Failure", - "result": result.model_dump(), - } - ) - with pytest.raises(ValidationError): - adapter.validate_python({"status": "pending", "message": "Wait"}) - - -def test_query_module_imports_no_policyengine_or_v1_metadata_source() -> None: - source_path = ( - Path(__file__).parents[3] - / "policyengine_api" - / "data" - / "v2" - / "catalog" - / "query.py" - ) - tree = ast.parse(source_path.read_text(encoding="utf-8")) - imported = { - alias.name - for node in ast.walk(tree) - if isinstance(node, ast.Import) - for alias in node.names - } | { - node.module or "" for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) - } - assert not any( - module.startswith( - ( - "policyengine_core", - "policyengine_us", - "policyengine_uk", - "policyengine_api.country", - "policyengine_api.services.metadata_service", - "policyengine_api.data.v2.catalog.extraction", - ) - ) - for module in imported ) def test_resource_collection_uses_bounded_pagination_without_counting( catalog_session: Session, ) -> None: - model_version = catalog_session.exec( - select(TaxBenefitModelVersion) - .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) - .where( - TaxBenefitModel.name == "policyengine-us", - TaxBenefitModelVersion.version == POLICYENGINE_VERSION, - ) - ).one() + model_version = _us_model_version(catalog_session) catalog_session.add( Variable( id=uuid4(), @@ -661,10 +236,7 @@ def record_statement(_connection, _cursor, statement, *_args) -> None: bind = catalog_session.get_bind() event.listen(bind, "before_cursor_execute", record_statement) try: - result = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ).list_variables("us", offset=0, limit=1) + result = _service(catalog_session).list_variables("us", limit=1) finally: event.remove(bind, "before_cursor_execute", record_statement) @@ -696,10 +268,11 @@ def record_statement(_connection, _cursor, statement, *_args) -> None: event.listen(bind, "before_cursor_execute", record_statement) try: with pytest.raises(InvalidMetadataPageError): - V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ).list_variables("us", offset=offset, limit=limit) + _service(catalog_session).list_variables( + "us", + offset=offset, + limit=limit, + ) finally: event.remove(bind, "before_cursor_execute", record_statement) @@ -717,10 +290,7 @@ def record_statement(_connection, _cursor, statement, *_args) -> None: bind = catalog_session.get_bind() event.listen(bind, "before_cursor_execute", record_statement) try: - result = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ).list_parameters("us") + result = _service(catalog_session).list_parameters("us") finally: event.remove(bind, "before_cursor_execute", record_statement) @@ -733,18 +303,13 @@ def record_statement(_connection, _cursor, statement, *_args) -> None: def test_parameter_values_are_separate_canonical_resources( catalog_session: Session, ) -> None: + model_version = _us_model_version(catalog_session) parameter = catalog_session.exec( - select(Parameter) - .join( - TaxBenefitModelVersion, - TaxBenefitModelVersion.id == Parameter.tax_benefit_model_version_id, - ) - .where( - TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + select(Parameter).where( + Parameter.tax_benefit_model_version_id == model_version.id, Parameter.name == "gov.example.rate", ) - ).first() - override_id = uuid4() + ).one() catalog_session.add( ParameterValue( id=uuid4(), @@ -752,20 +317,16 @@ def test_parameter_values_are_separate_canonical_resources( value_json=0.9, start_date=datetime(2026, 1, 1, tzinfo=timezone.utc), end_date=None, - policy_id=override_id, + policy_id=uuid4(), ) ) catalog_session.commit() - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) - all_values = service.list_parameter_values( + all_values = _service(catalog_session).list_parameter_values( "us", parameter_id=parameter.id, ) - current_value = service.list_parameter_values( + current_value = _service(catalog_session).list_parameter_values( "us", parameter_id=parameter.id, current=True, @@ -780,10 +341,7 @@ def test_parameter_values_are_separate_canonical_resources( def test_parameter_children_are_loaded_one_level_at_a_time( catalog_session: Session, ) -> None: - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) + service = _service(catalog_session) root = service.list_parameter_children("us") government = service.list_parameter_children("us", parent_path="gov") @@ -798,16 +356,12 @@ def test_parameter_children_are_loaded_one_level_at_a_time( ] assert root.items[0].child_count == 1 assert example.items[0].parameter is not None - assert example.items[0].parameter.name == "gov.example.rate" def test_resource_filters_and_details_remain_inside_selected_catalog( catalog_session: Session, ) -> None: - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) + service = _service(catalog_session) variables = service.list_variables("us", search="employment") parameters = service.list_parameters("us", search="example rate") states = service.list_regions("us", region_type="state") @@ -821,6 +375,30 @@ def test_resource_filters_and_details_remain_inside_selected_catalog( service.get_region("us", uuid4()) +def test_all_resource_families_support_collection_and_detail_reads( + catalog_session: Session, +) -> None: + service = _service(catalog_session) + selection = service.get_model_by_country("us") + model = service.list_models("us").items[0] + model_version = service.list_model_versions("us").items[0] + variable = service.list_variables("us").items[0] + parameter = service.list_parameters("us").items[0] + value = service.list_parameter_values("us", parameter_id=parameter.id).items[0] + dataset = service.list_datasets("us").items[0] + region = service.list_regions("us").items[0] + + assert service.get_model("us", model.id).item == model + assert service.get_model_version("us", model_version.id).item == model_version + assert service.get_variable("us", variable.id).item == variable + assert service.get_parameter("us", parameter.id).item == parameter + assert service.get_parameter_value("us", value.id).item == value + assert service.get_dataset("us", dataset.id).item == dataset + assert service.get_region("us", region.id).item == region + assert selection.model == model + assert selection.model_version == model_version + + def test_each_resource_result_identifies_an_exact_selected_version( catalog_session: Session, ) -> None: @@ -830,10 +408,7 @@ def test_each_resource_result_identifies_an_exact_selected_version( current_law_id=22, time_periods=[2041, 2040], ) - service = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ) + service = _service(catalog_session) default_variables = service.list_variables("us") selected_variables = service.list_variables("us", "5.0.5") @@ -846,7 +421,52 @@ def test_each_resource_result_identifies_an_exact_selected_version( assert [item.name for item in selected_options.time_period] == [2041, 2040] -def test_economy_options_reads_only_regions_and_the_national_dataset( +def test_version_selection_rejects_invalid_absent_and_unsupported_requests( + catalog_session: Session, +) -> None: + service = _service(catalog_session) + + for invalid in ("", f" {POLICYENGINE_VERSION}", "not a version", "0.0.0"): + with pytest.raises(InvalidPolicyEngineVersionError): + service.list_variables("us", invalid) + with pytest.raises(MetadataCatalogVersionNotFoundError): + service.list_variables("us", "4.99.0") + with pytest.raises(UnsupportedPreviewCountryError): + service.list_variables("ca") + with pytest.raises(MetadataCatalogUnavailableError): + V2MetadataQueryService( + catalog_session, + running_policyengine_version="4.99.0", + ).list_variables("us") + + +def test_dataset_collection_excludes_outputs_and_storage_references( + catalog_session: Session, +) -> None: + model_version = _us_model_version(catalog_session) + output_id = uuid4() + catalog_session.add( + Dataset( + id=output_id, + tax_benefit_model_version_id=model_version.id, + name="simulation-output", + description="Generated result", + storage_path="private-output-reference", + year=2026, + is_output_dataset=True, + ) + ) + catalog_session.commit() + + service = _service(catalog_session) + result = service.list_datasets("us") + + assert "simulation-output" not in {dataset.name for dataset in result.items} + with pytest.raises(MetadataResourceNotFoundError): + service.get_dataset("us", output_id) + + +def test_economy_options_read_only_regions_and_the_national_dataset( catalog_session: Session, ) -> None: statements: list[str] = [] @@ -857,10 +477,7 @@ def record_statement(_connection, _cursor, statement, *_args) -> None: bind = catalog_session.get_bind() event.listen(bind, "before_cursor_execute", record_statement) try: - result = V2MetadataQueryService( - catalog_session, - running_policyengine_version=POLICYENGINE_VERSION, - ).get_economy_options("us") + result = _service(catalog_session).get_economy_options("us") finally: event.remove(bind, "before_cursor_execute", record_statement) @@ -868,3 +485,51 @@ def record_statement(_connection, _cursor, statement, *_args) -> None: assert all("variables" not in statement for statement in statements) assert all("parameters" not in statement for statement in statements) assert [dataset.label for dataset in result.datasets] == ["Microcosm"] + + +def test_economy_options_require_a_national_region_and_dataset( + catalog_session: Session, +) -> None: + national_region = catalog_session.exec( + select(Region).where(Region.code == "us") + ).one() + catalog_session.delete(national_region) + catalog_session.commit() + + with pytest.raises(MetadataCatalogUnavailableError, match="national v2 region"): + _service(catalog_session).get_economy_options("us") + + +def test_query_modules_import_no_policyengine_or_v1_metadata_source() -> None: + source_directory = ( + Path(__file__).parents[3] / "policyengine_api" / "data" / "v2" / "catalog" + ) + modules = ("query.py", "catalog_selection.py", "parameter_tree_query.py") + imported = set() + for module in modules: + tree = ast.parse((source_directory / module).read_text(encoding="utf-8")) + imported.update( + alias.name + for node in ast.walk(tree) + if isinstance(node, ast.Import) + for alias in node.names + ) + imported.update( + node.module or "" + for node in ast.walk(tree) + if isinstance(node, ast.ImportFrom) + ) + + assert not any( + module.startswith( + ( + "policyengine_core", + "policyengine_us", + "policyengine_uk", + "policyengine_api.country", + "policyengine_api.services.metadata_service", + "policyengine_api.data.v2.catalog.extraction", + ) + ) + for module in imported + ) diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py index ea6b6fc8d..75283c69a 100644 --- a/tests/unit/v2/test_metadata_routes.py +++ b/tests/unit/v2/test_metadata_routes.py @@ -1,4 +1,4 @@ -"""Typed route coverage for the dormant v2 metadata preview.""" +"""Typed route coverage for dormant v2 metadata resources.""" from __future__ import annotations @@ -21,22 +21,17 @@ from policyengine_api.data.v2.catalog.schemas import ( MetadataCanonicalParameterValue, MetadataDataset, + MetadataDatasetOption, MetadataDetailResult, - MetadataEconomyOptions, MetadataEconomyOptionsResult, MetadataModel, MetadataModelSelectionResult, - MetadataModelVersion, MetadataModelVersionDetail, MetadataPageResult, - MetadataParameter, MetadataParameterChild, - MetadataParameterNode, MetadataParameterSummary, - MetadataParameterValue, MetadataRegion, MetadataRegionOption, - MetadataResult, MetadataTimePeriodOption, MetadataVariable, ) @@ -49,44 +44,17 @@ ) -class Reader: - def __init__( - self, - result: MetadataResult, - error: Exception | None = None, - close_error: Exception | None = None, - ): - self.result = result - self.error = error - self.close_error = close_error - self.calls = [] - self.closed = False - - def get_metadata( - self, - country_id: str, - policyengine_version: str | None = None, - ) -> MetadataResult: - self.calls.append((country_id, policyengine_version)) - if self.error is not None: - raise self.error - return self.result - - def close(self) -> None: - self.closed = True - if self.close_error is not None: - raise self.close_error - - class ResourceReader: def __init__( self, results: dict[str, object], *, error: Exception | None = None, + close_error: Exception | None = None, ): self.results = results self.error = error + self.close_error = close_error self.calls: list[tuple[str, tuple, dict]] = [] self.closed = False @@ -104,120 +72,78 @@ def read(*args, **kwargs): def close(self) -> None: self.closed = True + if self.close_error is not None: + raise self.close_error -def _result(country_id: str = "us") -> MetadataResult: +def _resource_results() -> tuple[dict[str, object], dict[str, object]]: + version = "4.20.3" model_id = uuid4() - version_id = uuid4() + model_version_id = uuid4() + variable_id = uuid4() + parameter_id = uuid4() + parameter_value_id = uuid4() dataset_id = uuid4() region_id = uuid4() - parameter_id = uuid4() - return MetadataResult( - current_law_id=2 if country_id == "us" else 1, - model=MetadataModel( - id=model_id, - name=f"policyengine-{country_id}", - description="Model", - ), - model_version=MetadataModelVersion( - id=version_id, - model_id=model_id, - version="4.20.3", - description="PolicyEngine.py catalog", - ), - variables=[ - MetadataVariable( - id=uuid4(), - name="employment_income", - label="Employment income", - entity="person", - description=None, - data_type="float", - possible_values=None, - default_value=0, - adds=None, - subtracts=None, - ) - ], - parameter_nodes=[ - MetadataParameterNode( - id=uuid4(), - name="gov.example", - label="Example", - description=None, - ) - ], - parameters=[ - MetadataParameter( - id=parameter_id, - name="gov.example.rate", - label="Rate", - description=None, - data_type="float", - unit="/1", - values=[ - MetadataParameterValue( - id=uuid4(), - value=0.1, - start_date=datetime(2025, 1, 1, tzinfo=timezone.utc), - end_date=None, - ) - ], - ) - ], - datasets=[ - MetadataDataset( - id=dataset_id, - name=f"populace_{country_id}_2024", - description="Populace", - year=2024, - ) - ], - regions=[ - MetadataRegion( - id=region_id, - code=country_id, - label=country_id.upper(), - region_type="national", - requires_filter=False, - filter_field=None, - filter_value=None, - filter_strategy=None, - parent_code=None, - state_code=None, - state_name=None, - default_dataset_id=dataset_id, - ) - ], - economy_options=MetadataEconomyOptions( - region=[ - MetadataRegionOption( - name=country_id, - label=country_id.upper(), - type="national", - ) - ], - time_period=[MetadataTimePeriodOption(name=2026, label="2026")], - datasets=[], - ), + model = MetadataModel( + id=model_id, + name="policyengine-us", + description="US model", ) - - -def _resource_results(country_id: str = "us") -> dict[str, object]: - combined = _result(country_id) - version = combined.model_version.version model_version = MetadataModelVersionDetail( - **combined.model_version.model_dump(), - current_law_id=combined.current_law_id, + id=model_version_id, + model_id=model_id, + version=version, + description="PolicyEngine.py catalog", + current_law_id=2, metadata_time_periods=[2026], ) - parameter = combined.parameters[0] - parameter_summary = MetadataParameterSummary( - **parameter.model_dump(exclude={"values"}) + variable = MetadataVariable( + id=variable_id, + name="employment_income", + label="Employment income", + entity="person", + description=None, + data_type="float", + possible_values=None, + default_value=0, + adds=None, + subtracts=None, + ) + parameter = MetadataParameterSummary( + id=parameter_id, + name="gov.example.rate", + label="Rate", + description=None, + data_type="float", + unit="/1", ) parameter_value = MetadataCanonicalParameterValue( - parameter_id=parameter.id, - **parameter.values[0].model_dump(), + id=parameter_value_id, + parameter_id=parameter_id, + value=0.1, + start_date=datetime(2025, 1, 1, tzinfo=timezone.utc), + end_date=None, + ) + dataset = MetadataDataset( + id=dataset_id, + name="populace_us_2024", + description="National input dataset", + year=2024, + ) + region = MetadataRegion( + id=region_id, + code="us", + label="United States", + region_type="national", + requires_filter=False, + filter_field=None, + filter_value=None, + filter_strategy=None, + parent_code=None, + state_code=None, + state_name=None, + default_dataset_id=dataset_id, ) def page(items: list[object]) -> MetadataPageResult: @@ -232,43 +158,64 @@ def page(items: list[object]) -> MetadataPageResult: def detail(item: object) -> MetadataDetailResult: return MetadataDetailResult(policyengine_version=version, item=item) - return { - "list_models": page([combined.model]), - "get_model": detail(combined.model), - "get_model_by_country": MetadataModelSelectionResult( - policyengine_version=version, - model=combined.model, - model_version=model_version, - ), - "list_model_versions": page([model_version]), - "get_model_version": detail(model_version), - "list_variables": page(combined.variables), - "get_variable": detail(combined.variables[0]), - "list_parameters": page([parameter_summary]), - "list_parameter_children": page( - [ - MetadataParameterChild( - path=parameter_summary.name, - label=parameter_summary.label or parameter_summary.name, - type="parameter", - parameter=parameter_summary, - ) - ] - ), - "get_parameter": detail(parameter_summary), - "list_parameter_values": page([parameter_value]), - "get_parameter_value": detail(parameter_value), - "list_datasets": page(combined.datasets), - "get_dataset": detail(combined.datasets[0]), - "list_regions": page(combined.regions), - "get_region_by_code": detail(combined.regions[0]), - "get_region": detail(combined.regions[0]), - "get_economy_options": MetadataEconomyOptionsResult( - policyengine_version=version, - current_law_id=combined.current_law_id, - **combined.economy_options.model_dump(), - ), - } + return ( + { + "list_models": page([model]), + "get_model": detail(model), + "get_model_by_country": MetadataModelSelectionResult( + policyengine_version=version, + model=model, + model_version=model_version, + ), + "list_model_versions": page([model_version]), + "get_model_version": detail(model_version), + "list_variables": page([variable]), + "get_variable": detail(variable), + "list_parameters": page([parameter]), + "list_parameter_children": page( + [ + MetadataParameterChild( + path=parameter.name, + label=parameter.label or parameter.name, + type="parameter", + parameter=parameter, + ) + ] + ), + "get_parameter": detail(parameter), + "list_parameter_values": page([parameter_value]), + "get_parameter_value": detail(parameter_value), + "list_datasets": page([dataset]), + "get_dataset": detail(dataset), + "list_regions": page([region]), + "get_region_by_code": detail(region), + "get_region": detail(region), + "get_economy_options": MetadataEconomyOptionsResult( + policyengine_version=version, + current_law_id=2, + region=[ + MetadataRegionOption( + name="us", + label="United States", + type="national", + ) + ], + time_period=[MetadataTimePeriodOption(name=2026, label="2026")], + datasets=[ + MetadataDatasetOption(name="populace_us_2024", label="Microcosm") + ], + ), + }, + { + "model_id": model_id, + "model_version_id": model_version_id, + "variable_id": variable_id, + "parameter_id": parameter_id, + "parameter_value_id": parameter_value_id, + "dataset_id": dataset_id, + "region_id": region_id, + }, + ) def _client(factory) -> TestClient: @@ -306,275 +253,16 @@ def v1_metadata(country_id: str): ) -@pytest.mark.parametrize("country_id", ["us", "uk"]) -def test_preview_get_returns_typed_catalog_response(country_id: str) -> None: - readers = [] - - def factory(): - reader = Reader(_result(country_id)) - readers.append(reader) - return reader - - response = _client(factory).get(f"/v2/{country_id}/metadata") - - assert response.status_code == 200 - assert response.headers["content-type"].startswith("application/json") - payload = response.json() - assert payload["status"] == "ok" - assert payload["message"] is None - assert payload["result"]["current_law_id"] == (2 if country_id == "us" else 1) - assert payload["result"]["model_version"]["version"] == "4.20.3" - assert payload["result"]["economy_options"]["region"][0]["name"] == country_id - assert isinstance( - payload["result"]["economy_options"]["time_period"][0]["name"], - int, - ) - assert readers[0].calls == [(country_id, None)] - assert readers[0].closed - - -def test_unsupported_country_and_methods_return_typed_client_errors() -> None: - calls = [] - - def factory(): - calls.append("called") - return Reader(_result()) - - client = _client(factory) - country_response = client.get("/v2/ca/metadata") - assert country_response.status_code == 404 - assert country_response.json()["status"] == "error" - assert country_response.json()["message"] - - for method in ("POST", "PUT", "PATCH", "DELETE", "OPTIONS"): - response = client.request(method, "/v2/us/metadata") - assert response.status_code == 405 - assert response.json()["status"] == "error" - assert response.json()["message"] - assert calls == [] - - -@pytest.mark.parametrize( - ("error", "expected_status"), - [ - (MetadataCatalogUnavailableError("missing"), 503), - (V2ConfigurationError("missing URL"), 503), - (RuntimeError("private database detail"), 500), - ], -) -def test_preview_failures_are_typed_and_hide_internal_details( - error: Exception, - expected_status: int, -) -> None: - reader = Reader(_result(), error=error) - response = _client(lambda: reader).get("/v2/us/metadata") - - assert response.status_code == expected_status - assert response.json()["status"] == "error" - assert response.json()["message"] - assert "private database detail" not in response.text - assert reader.closed - - -@pytest.mark.parametrize( - ("version", "error", "expected_status"), - [ - ( - "not a version", - InvalidPolicyEngineVersionError("invalid PolicyEngine.py version"), - 400, - ), - ( - "4.99.0", - MetadataCatalogVersionNotFoundError( - "PolicyEngine.py 4.99.0 is not published for us" - ), - 404, - ), - ], -) -def test_preview_version_selector_returns_typed_client_errors( - version: str, - error: Exception, - expected_status: int, -) -> None: - reader = Reader(_result(), error=error) - - response = _client(lambda: reader).get( - "/v2/us/metadata", - params={"policyengine_version": version}, - ) - - assert response.status_code == expected_status - assert response.json()["status"] == "error" - assert response.json()["message"] - assert reader.calls == [("us", version)] - assert reader.closed - - -def test_preview_passes_explicit_version_to_reader() -> None: - result = _result() - result.model_version.version = "4.19.0" - reader = Reader(result) - - response = _client(lambda: reader).get( - "/v2/us/metadata", - params={"policyengine_version": "4.19.0"}, - ) - - assert response.status_code == 200 - assert response.json()["result"]["model_version"]["version"] == "4.19.0" - assert reader.calls == [("us", "4.19.0")] - - -def test_preview_uses_default_reader_factory_when_none_is_injected( - monkeypatch: pytest.MonkeyPatch, -) -> None: - reader = Reader(_result()) - monkeypatch.setattr( - route_dependencies, - "_default_v2_metadata_reader_factory", - lambda: reader, - ) - - response = _client(None).get("/v2/us/metadata") - - assert response.status_code == 200 - assert reader.calls == [("us", None)] - assert reader.closed - - -def test_preview_ignores_reader_close_failure() -> None: - reader = Reader(_result(), close_error=RuntimeError("close failed")) - - response = _client(lambda: reader).get("/v2/us/metadata") - - assert response.status_code == 200 - assert response.json()["status"] == "ok" - assert reader.closed - - -def test_default_reader_uses_the_installed_policyengine_version( - monkeypatch: pytest.MonkeyPatch, -) -> None: - from policyengine_api.data.v2 import database - from policyengine_api.data.v2.catalog import query - - session = object() - captured = {} - reader = object() - - def query_service(candidate_session, *, running_policyengine_version): - captured["session"] = candidate_session - captured["version"] = running_policyengine_version - return reader - - monkeypatch.setattr(database, "get_v2_session_factory", lambda: lambda: session) - monkeypatch.setattr(query, "V2MetadataQueryService", query_service) - monkeypatch.setattr( - route_dependencies.importlib_metadata, - "version", - lambda package: "5.2.0" if package == "policyengine" else "unexpected", - ) - route_dependencies._running_policyengine_version.cache_clear() - try: - result = route_dependencies._default_v2_metadata_reader_factory() - finally: - route_dependencies._running_policyengine_version.cache_clear() - - assert result is reader - assert captured == {"session": session, "version": "5.2.0"} - - -def test_preview_reads_repeat_without_changing_v1_routing() -> None: - reader = Reader(_result()) - client = _client(lambda: reader) - - first = client.get("/v2/us/metadata") - second = client.get("/v2/us/metadata") - v1 = client.get("/us/metadata") - - assert first.json() == second.json() - assert reader.calls == [("us", None), ("us", None)] - assert v1.json()["result"] == {"source": "v1", "country_id": "us"} - - -def test_openapi_references_explicit_preview_response_schemas() -> None: - client = _client(lambda: Reader(_result())) - response = client.get("/v2/openapi.json") - - assert response.status_code == 200 - assert response.headers["content-type"].startswith("application/json") - schema = response.json() - assert set(schema["paths"]) == { - "/v2/datasets", - "/v2/datasets/{dataset_id}", - "/v2/economy-options", - "/v2/parameters", - "/v2/parameters/children", - "/v2/parameters/{parameter_id}", - "/v2/parameter-values", - "/v2/parameter-values/{value_id}", - "/v2/regions", - "/v2/regions/by-code/{region_code}", - "/v2/regions/{region_id}", - "/v2/tax-benefit-models", - "/v2/tax-benefit-models/by-country/{country_id}", - "/v2/tax-benefit-models/{model_id}", - "/v2/tax-benefit-model-versions", - "/v2/tax-benefit-model-versions/{version_id}", - "/v2/us/metadata", - "/v2/uk/metadata", - "/v2/variables", - "/v2/variables/{variable_id}", - "/v2/{country_id}/metadata", - } - - for path in ("/v2/us/metadata", "/v2/uk/metadata"): - operation = schema["paths"][path]["get"] - assert set(operation["responses"]) >= { - "200", - "400", - "404", - "405", - "500", - "503", - } - assert ( - operation["responses"]["200"]["content"]["application/json"]["schema"][ - "$ref" - ] - == "#/components/schemas/MetadataSuccessResponse" - ) - for status in ("400", "404", "405", "500", "503"): - assert ( - operation["responses"][status]["content"]["application/json"]["schema"][ - "$ref" - ] - == "#/components/schemas/MetadataErrorResponse" - ) - - unsupported = schema["paths"]["/v2/{country_id}/metadata"] - assert set(unsupported) == {"get"} - assert ( - unsupported["get"]["responses"]["404"]["content"]["application/json"]["schema"][ - "$ref" - ] - == "#/components/schemas/MetadataErrorResponse" - ) - - -def test_each_split_resource_route_returns_its_typed_result() -> None: - reader = ResourceReader(_resource_results()) +def test_each_resource_route_returns_its_typed_result() -> None: + results, ids = _resource_results() + reader = ResourceReader(results) client = _client(lambda: reader) - combined = _result() - parameter = combined.parameters[0] requests = [ ("list_models", "/v2/tax-benefit-models?country_id=us"), ("get_model_by_country", "/v2/tax-benefit-models/by-country/us"), ( "get_model", - f"/v2/tax-benefit-models/{combined.model.id}?country_id=us", + f"/v2/tax-benefit-models/{ids['model_id']}?country_id=us", ), ( "list_model_versions", @@ -582,38 +270,29 @@ def test_each_split_resource_route_returns_its_typed_result() -> None: ), ( "get_model_version", - f"/v2/tax-benefit-model-versions/{combined.model_version.id}?country_id=us", + f"/v2/tax-benefit-model-versions/{ids['model_version_id']}?country_id=us", ), ("list_variables", "/v2/variables?country_id=us"), - ( - "get_variable", - f"/v2/variables/{combined.variables[0].id}?country_id=us", - ), + ("get_variable", f"/v2/variables/{ids['variable_id']}?country_id=us"), ("list_parameters", "/v2/parameters?country_id=us"), ( "list_parameter_children", "/v2/parameters/children?country_id=us&parent_path=gov.example", ), - ("get_parameter", f"/v2/parameters/{parameter.id}?country_id=us"), + ( + "get_parameter", + f"/v2/parameters/{ids['parameter_id']}?country_id=us", + ), ("list_parameter_values", "/v2/parameter-values?country_id=us"), ( "get_parameter_value", - f"/v2/parameter-values/{parameter.values[0].id}?country_id=us", + f"/v2/parameter-values/{ids['parameter_value_id']}?country_id=us", ), ("list_datasets", "/v2/datasets?country_id=us"), - ( - "get_dataset", - f"/v2/datasets/{combined.datasets[0].id}?country_id=us", - ), + ("get_dataset", f"/v2/datasets/{ids['dataset_id']}?country_id=us"), ("list_regions", "/v2/regions?country_id=us"), - ( - "get_region_by_code", - f"/v2/regions/by-code/{combined.regions[0].code}?country_id=us", - ), - ( - "get_region", - f"/v2/regions/{combined.regions[0].id}?country_id=us", - ), + ("get_region_by_code", "/v2/regions/by-code/us?country_id=us"), + ("get_region", f"/v2/regions/{ids['region_id']}?country_id=us"), ("get_economy_options", "/v2/economy-options?country_id=us"), ] @@ -629,8 +308,9 @@ def test_each_split_resource_route_returns_its_typed_result() -> None: assert reader.closed -def test_split_route_forwards_version_filters_and_pagination() -> None: - reader = ResourceReader(_resource_results()) +def test_resource_route_forwards_version_filters_and_pagination() -> None: + results, _ids = _resource_results() + reader = ResourceReader(results) response = _client(lambda: reader).get( "/v2/variables", params={ @@ -662,12 +342,13 @@ def test_split_route_forwards_version_filters_and_pagination() -> None: {"country_id": "us", "limit": 501}, ], ) -def test_split_route_validation_failures_use_the_error_schema(params: dict) -> None: +def test_request_validation_failures_use_the_error_schema(params: dict) -> None: calls = [] def factory(): calls.append("called") - return ResourceReader(_resource_results()) + results, _ids = _resource_results() + return ResourceReader(results) response = _client(factory).get("/v2/variables", params=params) @@ -688,14 +369,16 @@ def factory(): (MetadataResourceNotFoundError("missing variable"), 404), (MetadataCatalogVersionNotFoundError("missing version"), 404), (MetadataCatalogUnavailableError("unavailable"), 503), + (V2ConfigurationError("missing database URL"), 503), (RuntimeError("private query detail"), 500), ], ) -def test_split_route_query_failures_use_documented_error_statuses( +def test_query_failures_use_documented_error_statuses( error: Exception, expected_status: int, ) -> None: - reader = ResourceReader(_resource_results(), error=error) + results, _ids = _resource_results() + reader = ResourceReader(results, error=error) response = _client(lambda: reader).get("/v2/variables?country_id=us") assert response.status_code == expected_status @@ -703,3 +386,140 @@ def test_split_route_query_failures_use_documented_error_statuses( assert response.json()["message"] assert "private query detail" not in response.text assert reader.closed + + +def test_unknown_resources_and_unsupported_methods_use_error_schema() -> None: + calls = [] + + def factory(): + calls.append("called") + results, _ids = _resource_results() + return ResourceReader(results) + + client = _client(factory) + missing = client.get("/v2/not-a-resource") + removed_combined = client.get("/v2/us/metadata") + + assert missing.status_code == 404 + assert missing.json()["status"] == "error" + assert removed_combined.status_code == 404 + assert removed_combined.json()["status"] == "error" + for method in ("POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"): + response = client.request(method, "/v2/variables") + assert response.status_code == 405 + if method != "HEAD": + assert response.json()["status"] == "error" + assert calls == [] + + +def test_default_reader_uses_the_installed_policyengine_version( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from policyengine_api.data.v2 import database + from policyengine_api.data.v2.catalog import query + + session = object() + captured = {} + reader = object() + + def query_service(candidate_session, *, running_policyengine_version): + captured["session"] = candidate_session + captured["version"] = running_policyengine_version + return reader + + monkeypatch.setattr(database, "get_v2_session_factory", lambda: lambda: session) + monkeypatch.setattr(query, "V2MetadataQueryService", query_service) + monkeypatch.setattr( + route_dependencies.importlib_metadata, + "version", + lambda package: "5.2.0" if package == "policyengine" else "unexpected", + ) + route_dependencies._running_policyengine_version.cache_clear() + try: + result = route_dependencies._default_v2_metadata_reader_factory() + finally: + route_dependencies._running_policyengine_version.cache_clear() + + assert result is reader + assert captured == {"session": session, "version": "5.2.0"} + + +def test_default_factory_and_close_failures_preserve_success( + monkeypatch: pytest.MonkeyPatch, +) -> None: + results, _ids = _resource_results() + reader = ResourceReader(results, close_error=RuntimeError("close failed")) + monkeypatch.setattr( + route_dependencies, + "_default_v2_metadata_reader_factory", + lambda: reader, + ) + + response = _client(None).get("/v2/variables?country_id=us") + + assert response.status_code == 200 + assert response.json()["status"] == "ok" + assert reader.closed + + +def test_resource_reads_do_not_change_v1_metadata_routing() -> None: + results, _ids = _resource_results() + reader = ResourceReader(results) + client = _client(lambda: reader) + + first = client.get("/v2/variables?country_id=us") + second = client.get("/v2/variables?country_id=us") + v1 = client.get("/us/metadata") + + assert first.json() == second.json() + assert v1.json()["result"] == {"source": "v1", "country_id": "us"} + + +def test_openapi_references_explicit_resource_response_schemas() -> None: + results, _ids = _resource_results() + response = _client(lambda: ResourceReader(results)).get("/v2/openapi.json") + + assert response.status_code == 200 + schema = response.json() + expected_paths = { + "/v2/datasets", + "/v2/datasets/{dataset_id}", + "/v2/economy-options", + "/v2/parameters", + "/v2/parameters/children", + "/v2/parameters/{parameter_id}", + "/v2/parameter-values", + "/v2/parameter-values/{value_id}", + "/v2/regions", + "/v2/regions/by-code/{region_code}", + "/v2/regions/{region_id}", + "/v2/tax-benefit-models", + "/v2/tax-benefit-models/by-country/{country_id}", + "/v2/tax-benefit-models/{model_id}", + "/v2/tax-benefit-model-versions", + "/v2/tax-benefit-model-versions/{version_id}", + "/v2/variables", + "/v2/variables/{variable_id}", + } + assert set(schema["paths"]) == expected_paths + + for path in expected_paths: + operation = schema["paths"][path]["get"] + assert set(operation["responses"]) >= { + "200", + "400", + "404", + "405", + "422", + "500", + "503", + } + success_schema = operation["responses"]["200"]["content"]["application/json"][ + "schema" + ] + assert success_schema["$ref"].startswith("#/components/schemas/Metadata") + for status in ("400", "404", "405", "422", "500", "503"): + error_schema = operation["responses"][status]["content"][ + "application/json" + ]["schema"] + assert error_schema["$ref"] == "#/components/schemas/MetadataErrorResponse" From 41e66deef2bda37c62a854a8fc4dc6029e2d24c8 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:50:27 +0400 Subject: [PATCH 18/27] Update v2 metadata route contracts --- docs/engineering/migration-contracts.md | 30 ++- docs/generated/migration_contracts.json | 252 +++++++++++++++++- docs/migration/stage-9-v2-metadata.md | 52 ++-- policyengine_api/asgi_factory.py | 5 +- policyengine_api/migration_logging.py | 25 +- policyengine_api/migration_registry.py | 12 +- scripts/guards/migration_contracts.py | 2 +- tests/contract/registry.py | 75 +++++- .../test_app_v2_workflow_contracts.py | 13 +- .../routes/test_migration_context_logging.py | 4 +- .../unit/test_migration_contract_artifacts.py | 4 +- tests/unit/test_migration_flags.py | 4 +- tests/unit/v2/test_metadata_routes.py | 16 ++ 13 files changed, 429 insertions(+), 65 deletions(-) diff --git a/docs/engineering/migration-contracts.md b/docs/engineering/migration-contracts.md index 3f96a9b9f..8cab684a5 100644 --- a/docs/engineering/migration-contracts.md +++ b/docs/engineering/migration-contracts.md @@ -8,7 +8,7 @@ Generated from `policyengine_api/migration_registry.py` and `tests/contract/regi | --- | ---: | | route group count | 9 | | workflow count | 8 | -| request count | 16 | +| request count | 32 | | db entity count | 6 | | sim flow count | 3 | @@ -18,7 +18,7 @@ Generated from `policyengine_api/migration_registry.py` and `tests/contract/regi | --- | --- | --- | --- | | `health` | `health`, `simulation-gateway-check`, `liveness-check`, `readiness-check` | `none` | `none` | | `specification` | `specification` | `none` | `none` | -| `metadata` | `metadata` | `metadata` | `none` | +| `metadata` | `metadata`, `datasets`, `economy-options`, `parameter-values`, `parameters`, `regions`, `tax-benefit-model-versions`, `tax-benefit-models`, `variables` | `metadata` | `none` | | `policy` | `policy`, `policies`, `user-policy` | `policy` | `none` | | `household` | `household`, `calculate`, `calculate-full` | `household` | `household` | | `economy` | `economy` | `simulation` | `economy` | @@ -69,15 +69,31 @@ Generated from `policyengine_api/migration_registry.py` and `tests/contract/regi | `GET` | `/us/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | | `GET` | `/uk/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | -### `region_selection_v2_preview` +### `metadata_resources_v2_preview` -- Current contract: `typed_v2_preview` -- Future owner: Later metadata read cutover and preview-path removal +- Current contract: `typed_v2_resources` +- Future owner: Later metadata read cutover and v2 route-prefix removal | Method | Path | Status | Route group | Stable response fields | | --- | --- | ---: | --- | --- | -| `GET` | `/v2/us/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | -| `GET` | `/v2/uk/metadata` | 200 | `metadata` | `status`, `result.current_law_id`, `result.economy_options.region`, `result.economy_options.time_period` | +| `GET` | `/v2/tax-benefit-models?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/tax-benefit-model-versions?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/variables?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/parameters?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/parameters/children?country_id=us&parent_path=gov` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/parameter-values?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/datasets?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/regions?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.items`, `result.offset`, `result.limit`, `result.has_more` | +| `GET` | `/v2/tax-benefit-models/{model_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/tax-benefit-model-versions/{version_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/variables/{variable_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/parameters/{parameter_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/parameter-values/{value_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/datasets/{dataset_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/regions/{region_id}?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/regions/by-code/state/ca?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.item` | +| `GET` | `/v2/tax-benefit-models/by-country/us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.model`, `result.model_version` | +| `GET` | `/v2/economy-options?country_id=us` | 200 | `metadata` | `status`, `message`, `result.policyengine_version`, `result.current_law_id`, `result.region`, `result.time_period`, `result.datasets` | ### `simulation_submit_poll` diff --git a/docs/generated/migration_contracts.json b/docs/generated/migration_contracts.json index ce2464979..fa25f245c 100644 --- a/docs/generated/migration_contracts.json +++ b/docs/generated/migration_contracts.json @@ -1,7 +1,7 @@ { "metadata": { "db_entity_count": 6, - "request_count": 16, + "request_count": 32, "route_group_count": 9, "sim_flow_count": 3, "workflow_count": 8 @@ -30,7 +30,15 @@ "db_entity": "metadata", "name": "metadata", "path_segments": [ - "metadata" + "metadata", + "datasets", + "economy-options", + "parameter-values", + "parameters", + "regions", + "tax-benefit-model-versions", + "tax-benefit-models", + "variables" ], "sim_flow": null }, @@ -218,32 +226,252 @@ ] }, { - "current_contract": "typed_v2_preview", - "future_owner_pr": "Later metadata read cutover and preview-path removal", - "name": "region_selection_v2_preview", + "current_contract": "typed_v2_resources", + "future_owner_pr": "Later metadata read cutover and v2 route-prefix removal", + "name": "metadata_resources_v2_preview", "requests": [ { "expected_status": 200, "method": "GET", - "path": "/v2/us/metadata", + "path": "/v2/tax-benefit-models?country_id=us", "route_group": "metadata", "stable_response_fields": [ "status", - "result.current_law_id", - "result.economy_options.region", - "result.economy_options.time_period" + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/tax-benefit-model-versions?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/variables?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/parameters?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/parameters/children?country_id=us&parent_path=gov", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/parameter-values?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/datasets?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/regions?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/tax-benefit-models/{model_id}?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" ] }, { "expected_status": 200, "method": "GET", - "path": "/v2/uk/metadata", + "path": "/v2/tax-benefit-model-versions/{version_id}?country_id=us", "route_group": "metadata", "stable_response_fields": [ "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/variables/{variable_id}?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/parameters/{parameter_id}?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/parameter-values/{value_id}?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/datasets/{dataset_id}?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/regions/{region_id}?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/regions/by-code/state/ca?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.item" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/tax-benefit-models/by-country/us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", + "result.model", + "result.model_version" + ] + }, + { + "expected_status": 200, + "method": "GET", + "path": "/v2/economy-options?country_id=us", + "route_group": "metadata", + "stable_response_fields": [ + "status", + "message", + "result.policyengine_version", "result.current_law_id", - "result.economy_options.region", - "result.economy_options.time_period" + "result.region", + "result.time_period", + "result.datasets" ] } ] diff --git a/docs/migration/stage-9-v2-metadata.md b/docs/migration/stage-9-v2-metadata.md index fc3f92859..b3d48a3f0 100644 --- a/docs/migration/stage-9-v2-metadata.md +++ b/docs/migration/stage-9-v2-metadata.md @@ -8,12 +8,18 @@ approved environment inventory and secret-management system. ## Scope Stage 9 populates the dormant v2 US and UK reference catalogs and exposes -read-only preview endpoints from the Cloud Run ASGI application at -`GET /v2/us/metadata` and `GET /v2/uk/metadata`. Their generated OpenAPI -document is available at `GET /v2/openapi.json`. App Engine continues to run -the Flask v1 application and does not expose these routes. Stage 9 does not -change `GET /us/metadata`, `GET /uk/metadata`, their callers, or their v1 data -source. Existing clients must not be redirected to the preview endpoints. +read-only, resource-specific routes from the Cloud Run ASGI application under +`/v2`. The routes separately serve tax-benefit models, model versions, +variables, parameters, direct parameter-tree children, canonical parameter +values, logical input datasets, regions, and compact economy-selection +options. Collection routes require `country_id`, use bounded `offset` and +`limit` pagination, and never assemble the complete catalog into one response. +The generated OpenAPI document is available at `GET /v2/openapi.json`. + +App Engine continues to run the Flask v1 application and does not expose these +routes. Stage 9 does not change `GET /us/metadata`, `GET /uk/metadata`, their +callers, or their v1 data source. Existing clients must not be redirected to +the v2 resource routes. The initializer creates only reusable logical input `Dataset` rows. Each row has `is_output_dataset=false` and a null `storage_path`. It creates no @@ -164,25 +170,31 @@ target. ## Preview verification -After Cloud Run candidate creation, explicitly request both preview GET -endpoints and validate their typed response envelopes. A request without a -`policyengine_version` query parameter selects the exact PolicyEngine.py version -installed in that candidate artifact. It does not select the newest database -row. Also request a known published version with, for example, -`?policyengine_version=5.0.4`, and confirm that the response contains that -exact version's complete snapshot. +After Cloud Run candidate creation, explicitly request US and UK collections +for variables, parameters, parameter values, datasets, and regions, followed +by representative detail routes and `GET /v2/economy-options`. Confirm that +parameter collection responses do not contain parameter values and that direct +parameter-tree-child requests return only one hierarchy level. A request +without a `policyengine_version` query parameter selects the exact +PolicyEngine.py version installed in that candidate artifact. It does not +select the newest database row. Also request a known published version with, +for example, `?policyengine_version=5.0.4`, and confirm that each response +identifies that exact selected version. A successful response has HTTP 200, `status: "ok"`, `message: null`, and a -typed `result`. A malformed or noncanonical explicit version returns a typed -HTTP 400 error. A valid explicit version that has not been published for the -country returns a typed HTTP 404 error. An absent or incomplete catalog for the -candidate's installed default version returns a typed HTTP 503 service error. -None of these outcomes reads or repairs v1. Unsupported countries and methods +typed `result`. Collection results contain `policyengine_version`, `items`, +`offset`, `limit`, and `has_more`; detail results contain +`policyengine_version` and `item`. A malformed or noncanonical explicit +version returns a typed HTTP 400 error. A valid explicit version that has not +been published for the country returns a typed HTTP 404 error. An absent +catalog for the candidate's installed default version returns a typed HTTP 503 +service error. None of these outcomes reads or repairs v1. Unsupported +countries and methods return typed client errors. Request `GET /v2/openapi.json` and confirm that the public document contains -the US, UK, and unsupported-country preview paths and explicit component schema -references for every documented response. +every documented resource path and explicit component schema references for +success and error responses. Also request the unprefixed US and UK metadata endpoints and confirm their responses still come from v1. Do not modify route selectors, internal callers, diff --git a/policyengine_api/asgi_factory.py b/policyengine_api/asgi_factory.py index 9a1a76435..115296141 100644 --- a/policyengine_api/asgi_factory.py +++ b/policyengine_api/asgi_factory.py @@ -142,7 +142,10 @@ def log_native_route(status_code: int) -> None: path=request.url.path, status_code=status_code, started_at=started_at, - country_id=request.path_params.get("country_id"), + country_id=( + request.path_params.get("country_id") + or request.query_params.get("country_id") + ), route_impl=RouteImplementation.FASTAPI_NATIVE, ) except Exception: diff --git a/policyengine_api/migration_logging.py b/policyengine_api/migration_logging.py index d98de5d55..02a54e26d 100644 --- a/policyengine_api/migration_logging.py +++ b/policyengine_api/migration_logging.py @@ -19,14 +19,31 @@ ) -V2_METADATA_PREVIEW_READ_PATHS = frozenset( +V2_METADATA_RESOURCE_SEGMENTS = frozenset( { - "/v2/us/metadata", - "/v2/uk/metadata", + "datasets", + "economy-options", + "parameter-values", + "parameters", + "regions", + "tax-benefit-model-versions", + "tax-benefit-models", + "variables", } ) +def _is_v2_metadata_resource_read(method: str, path: str) -> bool: + if method != "GET": + return False + segments = [segment for segment in path.strip("/").split("/") if segment] + return ( + len(segments) >= 2 + and segments[0] == "v2" + and segments[1] in V2_METADATA_RESOURCE_SEGMENTS + ) + + def register_migration_request_logging(app: flask.Flask) -> None: """Register request IDs, backend headers, and migration logging for Flask.""" @@ -80,7 +97,7 @@ def log_migration_request( elapsed_ms = round((time.time() - started_at) * 1000, 2) route_group = infer_route_group(path) - is_v2_metadata_read = method == "GET" and path in V2_METADATA_PREVIEW_READ_PATHS + is_v2_metadata_read = _is_v2_metadata_resource_read(method, path) migration_context = get_migration_log_context( route_group, route_impl=route_impl, diff --git a/policyengine_api/migration_registry.py b/policyengine_api/migration_registry.py index bc5171400..a19376b8c 100644 --- a/policyengine_api/migration_registry.py +++ b/policyengine_api/migration_registry.py @@ -29,7 +29,17 @@ class RouteGroupConfig: ), RouteGroupConfig( name="metadata", - path_segments=("metadata",), + path_segments=( + "metadata", + "datasets", + "economy-options", + "parameter-values", + "parameters", + "regions", + "tax-benefit-model-versions", + "tax-benefit-models", + "variables", + ), db_entity="metadata", ), RouteGroupConfig( diff --git a/scripts/guards/migration_contracts.py b/scripts/guards/migration_contracts.py index abf50e29b..699991520 100644 --- a/scripts/guards/migration_contracts.py +++ b/scripts/guards/migration_contracts.py @@ -9,7 +9,7 @@ REPO_ROOT = Path(__file__).resolve().parents[2] -ALLOWED_CURRENT_CONTRACTS = frozenset({"api_v1_compatible", "typed_v2_preview"}) +ALLOWED_CURRENT_CONTRACTS = frozenset({"api_v1_compatible", "typed_v2_resources"}) def _check_unique_values( diff --git a/tests/contract/registry.py b/tests/contract/registry.py index 0592c5e2a..c24e9eb41 100644 --- a/tests/contract/registry.py +++ b/tests/contract/registry.py @@ -123,31 +123,86 @@ class WorkflowContract: ), ), WorkflowContract( - name="region_selection_v2_preview", - current_contract="typed_v2_preview", - future_owner_pr="Later metadata read cutover and preview-path removal", + name="metadata_resources_v2_preview", + current_contract="typed_v2_resources", + future_owner_pr="Later metadata read cutover and v2 route-prefix removal", requests=( + *( + ContractRequest( + method="GET", + path=path, + expected_status=200, + stable_response_fields=( + "status", + "message", + "result.policyengine_version", + "result.items", + "result.offset", + "result.limit", + "result.has_more", + ), + route_group="metadata", + ) + for path in ( + "/v2/tax-benefit-models?country_id=us", + "/v2/tax-benefit-model-versions?country_id=us", + "/v2/variables?country_id=us", + "/v2/parameters?country_id=us", + "/v2/parameters/children?country_id=us&parent_path=gov", + "/v2/parameter-values?country_id=us", + "/v2/datasets?country_id=us", + "/v2/regions?country_id=us", + ) + ), + *( + ContractRequest( + method="GET", + path=path, + expected_status=200, + stable_response_fields=( + "status", + "message", + "result.policyengine_version", + "result.item", + ), + route_group="metadata", + ) + for path in ( + "/v2/tax-benefit-models/{model_id}?country_id=us", + "/v2/tax-benefit-model-versions/{version_id}?country_id=us", + "/v2/variables/{variable_id}?country_id=us", + "/v2/parameters/{parameter_id}?country_id=us", + "/v2/parameter-values/{value_id}?country_id=us", + "/v2/datasets/{dataset_id}?country_id=us", + "/v2/regions/{region_id}?country_id=us", + "/v2/regions/by-code/state/ca?country_id=us", + ) + ), ContractRequest( method="GET", - path="/v2/us/metadata", + path="/v2/tax-benefit-models/by-country/us", expected_status=200, stable_response_fields=( "status", - "result.current_law_id", - "result.economy_options.region", - "result.economy_options.time_period", + "message", + "result.policyengine_version", + "result.model", + "result.model_version", ), route_group="metadata", ), ContractRequest( method="GET", - path="/v2/uk/metadata", + path="/v2/economy-options?country_id=us", expected_status=200, stable_response_fields=( "status", + "message", + "result.policyengine_version", "result.current_law_id", - "result.economy_options.region", - "result.economy_options.time_period", + "result.region", + "result.time_period", + "result.datasets", ), route_group="metadata", ), diff --git a/tests/contract/test_app_v2_workflow_contracts.py b/tests/contract/test_app_v2_workflow_contracts.py index 6e0aa49e4..8a1991d84 100644 --- a/tests/contract/test_app_v2_workflow_contracts.py +++ b/tests/contract/test_app_v2_workflow_contracts.py @@ -12,7 +12,7 @@ def test_app_v2_workflow_contract_registry_is_complete(): "household_save_edit_read", "household_calculate", "region_selection", - "region_selection_v2_preview", + "metadata_resources_v2_preview", "simulation_submit_poll", "report_create_poll", "budget_window_submit_poll", @@ -20,8 +20,8 @@ def test_app_v2_workflow_contract_registry_is_complete(): for workflow in APP_V2_WORKFLOW_CONTRACTS: expected_contract = ( - "typed_v2_preview" - if workflow.name == "region_selection_v2_preview" + "typed_v2_resources" + if workflow.name == "metadata_resources_v2_preview" else "api_v1_compatible" ) assert workflow.current_contract == expected_contract @@ -41,4 +41,9 @@ def test_app_v2_workflow_contract_registry_is_complete(): ) assert {request.path for request in APP_V2_ROUTE_CONTRACTS} - { request.path for request in APP_V1_COMPATIBLE_ROUTE_CONTRACTS - } == {"/v2/us/metadata", "/v2/uk/metadata"} + } == { + request.path + for workflow in APP_V2_WORKFLOW_CONTRACTS + if workflow.name == "metadata_resources_v2_preview" + for request in workflow.requests + } diff --git a/tests/unit/routes/test_migration_context_logging.py b/tests/unit/routes/test_migration_context_logging.py index dde356f26..7b5499771 100644 --- a/tests/unit/routes/test_migration_context_logging.py +++ b/tests/unit/routes/test_migration_context_logging.py @@ -254,7 +254,7 @@ def test_native_metadata_logs_country_and_actual_implementation(): assert log_payload["migration"]["route_impl"] == "fastapi_native" -def test_v2_metadata_preview_logs_its_actual_supabase_read_source(monkeypatch): +def test_v2_metadata_resource_logs_its_actual_supabase_read_source(monkeypatch): monkeypatch.setenv("DB_READ_METADATA", "invalid-unprefixed-setting") monkeypatch.setenv("DB_WRITE_METADATA", "invalid-unprefixed-setting") @@ -262,7 +262,7 @@ def test_v2_metadata_preview_logs_its_actual_supabase_read_source(monkeypatch): log_migration_request( request_id="request-123", method="GET", - path="/v2/us/metadata", + path="/v2/variables", status_code=200, started_at=None, country_id="us", diff --git a/tests/unit/test_migration_contract_artifacts.py b/tests/unit/test_migration_contract_artifacts.py index 93f39f193..9517db7bc 100644 --- a/tests/unit/test_migration_contract_artifacts.py +++ b/tests/unit/test_migration_contract_artifacts.py @@ -11,7 +11,7 @@ def test_migration_contract_payload_summarizes_route_contracts(): assert payload["metadata"] == { "route_group_count": 9, "workflow_count": 8, - "request_count": 16, + "request_count": 32, "db_entity_count": 6, "sim_flow_count": 3, } @@ -20,7 +20,7 @@ def test_migration_contract_payload_summarizes_route_contracts(): "household_save_edit_read", "household_calculate", "region_selection", - "region_selection_v2_preview", + "metadata_resources_v2_preview", "simulation_submit_poll", "report_create_poll", "budget_window_submit_poll", diff --git a/tests/unit/test_migration_flags.py b/tests/unit/test_migration_flags.py index b808a0941..e1afa756e 100644 --- a/tests/unit/test_migration_flags.py +++ b/tests/unit/test_migration_flags.py @@ -135,7 +135,9 @@ def test_explicit_migration_context_rejects_invalid_database_sources( ("/readiness-check", "health"), ("/v2/openapi.json", "specification"), ("/us/metadata", "metadata"), - ("/v2/us/metadata", "metadata"), + ("/v2/variables", "metadata"), + ("/v2/parameters/children", "metadata"), + ("/v2/regions/state%2Fca", "metadata"), ("/us/policy/1", "policy"), ("/us/policies", "policy"), ("/us/household/1", "household"), diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py index 75283c69a..1fa8dfa91 100644 --- a/tests/unit/v2/test_metadata_routes.py +++ b/tests/unit/v2/test_metadata_routes.py @@ -3,6 +3,7 @@ from __future__ import annotations from datetime import datetime, timezone +from unittest.mock import patch from uuid import uuid4 from fastapi.testclient import TestClient @@ -475,6 +476,21 @@ def test_resource_reads_do_not_change_v1_metadata_routing() -> None: assert v1.json()["result"] == {"source": "v1", "country_id": "us"} +def test_resource_route_logging_records_the_country_query_parameter() -> None: + results, _ids = _resource_results() + + with patch("policyengine_api.migration_logging.logger") as mock_logger: + response = _client(lambda: ResourceReader(results)).get( + "/v2/variables?country_id=uk" + ) + + assert response.status_code == 200 + payload = mock_logger.log_struct.call_args.args[0] + assert payload["country_id"] == "uk" + assert payload["migration"]["route_group"] == "metadata" + assert payload["migration"]["db_read"] == "supabase" + + def test_openapi_references_explicit_resource_response_schemas() -> None: results, _ids = _resource_results() response = _client(lambda: ResourceReader(results)).get("/v2/openapi.json") From e7090fdfaccc580f3eb53af5567c31665d995e03 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:51:48 +0400 Subject: [PATCH 19/27] Test v2 metadata resources on PostgreSQL --- tests/integration/test_v2_metadata_routes.py | 247 ++++++++++++------- 1 file changed, 162 insertions(+), 85 deletions(-) diff --git a/tests/integration/test_v2_metadata_routes.py b/tests/integration/test_v2_metadata_routes.py index 18d175f17..58e648054 100644 --- a/tests/integration/test_v2_metadata_routes.py +++ b/tests/integration/test_v2_metadata_routes.py @@ -1,4 +1,4 @@ -"""Postgres-backed integration coverage for v2 metadata preview reads.""" +"""PostgreSQL-backed integration coverage for v2 metadata resource reads.""" from __future__ import annotations @@ -48,7 +48,7 @@ def _disposable_url() -> str: "::1", "postgres", }: - pytest.fail("v2 preview route tests require disposable local Postgres") + pytest.fail("v2 metadata route tests require disposable local PostgreSQL") return database_url @@ -97,72 +97,148 @@ def v1_metadata(country_id: str): ) -def test_postgres_preview_returns_complete_us_and_uk_catalogs_without_writes( +def _catalog_counts(engine: Engine) -> tuple[int, ...]: + with engine.connect() as connection: + return tuple( + connection.execute( + text( + """ + SELECT + (SELECT count(*) FROM variables), + (SELECT count(*) FROM parameters), + (SELECT count(*) FROM parameter_values), + (SELECT count(*) FROM datasets), + (SELECT count(*) FROM regions) + """ + ) + ).one() + ) + + +@pytest.mark.parametrize( + ("country_id", "current_law_id", "dataset_label"), + [("us", 2, "Microcosm"), ("uk", 1, "Enhanced FRS")], +) +def test_postgres_resource_collections_are_separate_and_read_only( published_engine: Engine, + country_id: str, + current_law_id: int, + dataset_label: str, ) -> None: client = _client(published_engine) - with published_engine.connect() as connection: - before = connection.execute( - text( - """ - SELECT - (SELECT count(*) FROM variables), - (SELECT count(*) FROM parameters), - (SELECT count(*) FROM parameter_values), - (SELECT count(*) FROM datasets), - (SELECT count(*) FROM regions) - """ - ) - ).one() + before = _catalog_counts(published_engine) + query = {"country_id": country_id} + + models = client.get("/v2/tax-benefit-models", params=query) + model_versions = client.get("/v2/tax-benefit-model-versions", params=query) + variables = client.get("/v2/variables", params=query) + parameters = client.get("/v2/parameters", params=query) + parameter_values = client.get("/v2/parameter-values", params=query) + datasets = client.get("/v2/datasets", params=query) + regions = client.get("/v2/regions", params=query) + options = client.get("/v2/economy-options", params=query) + + responses = ( + models, + model_versions, + variables, + parameters, + parameter_values, + datasets, + regions, + options, + ) + assert all(response.status_code == 200 for response in responses) + for response in responses: + assert response.json()["result"]["policyengine_version"] == ( + POLICYENGINE_VERSION + ) + + assert models.json()["result"]["items"][0]["name"] == (f"policyengine-{country_id}") + assert model_versions.json()["result"]["items"][0]["version"] == ( + POLICYENGINE_VERSION + ) + assert variables.json()["result"]["items"] + parameter_items = parameters.json()["result"]["items"] + assert parameter_items + assert "values" not in parameter_items[0] + assert parameter_values.json()["result"]["items"] + assert datasets.json()["result"]["items"] + assert regions.json()["result"]["items"] + assert all( + not dataset["is_output_dataset"] and dataset["storage_path"] is None + for dataset in datasets.json()["result"]["items"] + ) + option_result = options.json()["result"] + assert option_result["current_law_id"] == current_law_id + assert option_result["datasets"][0]["label"] == dataset_label + assert all( + isinstance(period["name"], int) and isinstance(period["label"], str) + for period in option_result["time_period"] + ) + assert _catalog_counts(published_engine) == before + - us = client.get("/v2/us/metadata") - uk = client.get("/v2/uk/metadata") +def test_postgres_parameter_tree_returns_direct_children_and_leaf_details( + published_engine: Engine, +) -> None: + client = _client(published_engine) + query = {"country_id": "us"} - assert us.status_code == uk.status_code == 200 - assert us.json()["result"]["current_law_id"] == 2 - assert uk.json()["result"]["current_law_id"] == 1 - assert us.json()["result"]["economy_options"]["datasets"][0]["label"] == ( - "Microcosm" + root = client.get("/v2/parameters/children", params=query) + government = client.get( + "/v2/parameters/children", + params={**query, "parent_path": "gov"}, ) - assert uk.json()["result"]["economy_options"]["datasets"][0]["label"] == ( - "Enhanced FRS" + example = client.get( + "/v2/parameters/children", + params={**query, "parent_path": "gov.example"}, ) - for country_id, response in (("us", us), ("uk", uk)): - result = response.json()["result"] - assert result["model"]["name"] == f"policyengine-{country_id}" - assert result["model_version"]["version"] == POLICYENGINE_VERSION - assert result["variables"] - assert result["parameter_nodes"] - assert result["parameters"] - assert result["parameters"][0]["values"] - assert result["datasets"] - assert result["regions"] - assert all( - not dataset["is_output_dataset"] and dataset["storage_path"] is None - for dataset in result["datasets"] - ) - assert all( - isinstance(period["name"], int) and isinstance(period["label"], str) - for period in result["economy_options"]["time_period"] - ) - with published_engine.connect() as connection: - after = connection.execute( - text( - """ - SELECT - (SELECT count(*) FROM variables), - (SELECT count(*) FROM parameters), - (SELECT count(*) FROM parameter_values), - (SELECT count(*) FROM datasets), - (SELECT count(*) FROM regions) - """ - ) - ).one() - assert after == before + assert root.status_code == government.status_code == example.status_code == 200 + assert [ + (item["path"], item["type"]) for item in root.json()["result"]["items"] + ] == [("gov", "node")] + assert [ + (item["path"], item["type"]) for item in government.json()["result"]["items"] + ] == [("gov.example", "node")] + leaf = example.json()["result"]["items"][0] + assert leaf["path"] == "gov.example.rate" + assert leaf["type"] == "parameter" + assert leaf["parameter"]["name"] == "gov.example.rate" + + +def test_postgres_collection_ids_resolve_through_detail_routes( + published_engine: Engine, +) -> None: + client = _client(published_engine) + query = {"country_id": "us"} + resources = ( + ("tax-benefit-models", "model_id"), + ("tax-benefit-model-versions", "version_id"), + ("variables", "variable_id"), + ("parameters", "parameter_id"), + ("parameter-values", "value_id"), + ("datasets", "dataset_id"), + ("regions", "region_id"), + ) + + for resource, _parameter_name in resources: + collection = client.get(f"/v2/{resource}", params=query) + resource_id = collection.json()["result"]["items"][0]["id"] + detail = client.get(f"/v2/{resource}/{resource_id}", params=query) + assert detail.status_code == 200 + assert detail.json()["result"]["item"]["id"] == resource_id + + by_country = client.get("/v2/tax-benefit-models/by-country/us") + by_code = client.get("/v2/regions/by-code/state/ca", params=query) + assert by_country.status_code == 200 + assert by_country.json()["result"]["model"]["name"] == "policyengine-us" + assert by_code.status_code == 200 + assert by_code.json()["result"]["item"]["code"] == "state/ca" -def test_postgres_preview_excludes_an_existing_output_dataset( +def test_postgres_dataset_collection_excludes_an_existing_output( published_engine: Engine, ) -> None: with Session(published_engine) as session: @@ -189,16 +265,19 @@ def test_postgres_preview_excludes_an_existing_output_dataset( ) session.commit() - response = _client(published_engine).get("/v2/us/metadata") + response = _client(published_engine).get( + "/v2/datasets", + params={"country_id": "us"}, + ) assert response.status_code == 200 assert "existing-output" not in { - dataset["name"] for dataset in response.json()["result"]["datasets"] + dataset["name"] for dataset in response.json()["result"]["items"] } @pytest.mark.parametrize("country_id", ["us", "uk"]) -def test_postgres_preview_defaults_to_running_version_and_accepts_exact_version( +def test_postgres_resources_default_to_running_version_and_accept_exact_override( published_engine: Engine, country_id: str, ) -> None: @@ -208,40 +287,37 @@ def test_postgres_preview_defaults_to_running_version_and_accepts_exact_version( ) client = _client(published_engine) - default_response = client.get(f"/v2/{country_id}/metadata") + default_response = client.get( + "/v2/variables", + params={"country_id": country_id}, + ) selected_response = client.get( - f"/v2/{country_id}/metadata", - params={"policyengine_version": "5.0.5"}, + "/v2/variables", + params={"country_id": country_id, "policyengine_version": "5.0.5"}, ) assert default_response.status_code == selected_response.status_code == 200 default_result = default_response.json()["result"] selected_result = selected_response.json()["result"] - assert default_result["model_version"]["version"] == POLICYENGINE_VERSION - assert selected_result["model_version"]["version"] == "5.0.5" - assert ( - default_result["model_version"]["id"] != selected_result["model_version"]["id"] - ) - assert {dataset["id"] for dataset in default_result["datasets"]}.isdisjoint( - dataset["id"] for dataset in selected_result["datasets"] - ) - assert {region["id"] for region in default_result["regions"]}.isdisjoint( - region["id"] for region in selected_result["regions"] + assert default_result["policyengine_version"] == POLICYENGINE_VERSION + assert selected_result["policyengine_version"] == "5.0.5" + assert {item["id"] for item in default_result["items"]}.isdisjoint( + item["id"] for item in selected_result["items"] ) -def test_postgres_preview_distinguishes_invalid_and_absent_versions( +def test_postgres_resources_distinguish_invalid_and_absent_versions( published_engine: Engine, ) -> None: client = _client(published_engine) invalid = client.get( - "/v2/us/metadata", - params={"policyengine_version": "not a version"}, + "/v2/variables", + params={"country_id": "us", "policyengine_version": "not a version"}, ) absent = client.get( - "/v2/us/metadata", - params={"policyengine_version": "4.99.0"}, + "/v2/variables", + params={"country_id": "us", "policyengine_version": "4.99.0"}, ) assert invalid.status_code == 400 @@ -252,7 +328,7 @@ def test_postgres_preview_distinguishes_invalid_and_absent_versions( assert absent.json()["message"] -def test_postgres_preview_returns_typed_error_for_incomplete_parameter_values( +def test_postgres_parameter_collection_remains_available_without_values( published_engine: Engine, ) -> None: with published_engine.begin() as connection: @@ -273,9 +349,10 @@ def test_postgres_preview_returns_typed_error_for_incomplete_parameter_values( ) ) - response = _client(published_engine).get("/v2/us/metadata") + client = _client(published_engine) + parameters = client.get("/v2/parameters", params={"country_id": "us"}) + values = client.get("/v2/parameter-values", params={"country_id": "us"}) - assert response.status_code == 503 - assert response.json()["status"] == "error" - assert response.json()["message"] - assert "result" not in response.json() + assert parameters.status_code == values.status_code == 200 + assert parameters.json()["result"]["items"] + assert values.json()["result"]["items"] == [] From 8fb38b414603609d61cdf44e7e537e6f2c8976cb Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:55:59 +0400 Subject: [PATCH 20/27] Run v2 resource integration coverage --- .github/workflows/v2-integration-check.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/v2-integration-check.yml b/.github/workflows/v2-integration-check.yml index 9bf9accee..b69399a7b 100644 --- a/.github/workflows/v2-integration-check.yml +++ b/.github/workflows/v2-integration-check.yml @@ -55,7 +55,7 @@ jobs: run: uv run coverage run --branch -m pytest -q tests/integration/test_v2_catalog_installed.py env: RUN_V2_CATALOG_COMPATIBILITY: "1" - - name: Test v2 metadata publication and preview routes + - name: Test v2 metadata publication and resource routes run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication.py tests/integration/test_v2_metadata_routes.py - name: Qualify production-scale v2 metadata publication run: uv run coverage run -a --branch -m pytest -q tests/integration/test_v2_catalog_publication_qualification.py From 37c56c26e65c9a216a644637b25a2d9f590b3c84 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 17:59:05 +0400 Subject: [PATCH 21/27] Verify repeated v2 resource reads --- tests/integration/test_v2_metadata_routes.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/integration/test_v2_metadata_routes.py b/tests/integration/test_v2_metadata_routes.py index 58e648054..a7cc3d79d 100644 --- a/tests/integration/test_v2_metadata_routes.py +++ b/tests/integration/test_v2_metadata_routes.py @@ -132,6 +132,7 @@ def test_postgres_resource_collections_are_separate_and_read_only( models = client.get("/v2/tax-benefit-models", params=query) model_versions = client.get("/v2/tax-benefit-model-versions", params=query) variables = client.get("/v2/variables", params=query) + repeated_variables = client.get("/v2/variables", params=query) parameters = client.get("/v2/parameters", params=query) parameter_values = client.get("/v2/parameter-values", params=query) datasets = client.get("/v2/datasets", params=query) @@ -159,6 +160,7 @@ def test_postgres_resource_collections_are_separate_and_read_only( POLICYENGINE_VERSION ) assert variables.json()["result"]["items"] + assert repeated_variables.json() == variables.json() parameter_items = parameters.json()["result"]["items"] assert parameter_items assert "values" not in parameter_items[0] From 942340de679b39629618dbc366f3aac1a91fb121 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 18:18:54 +0400 Subject: [PATCH 22/27] Fix v2 metadata route edge cases --- .../data/v2/catalog/parameter_tree_query.py | 64 +++++++++++++++---- policyengine_api/data/v2/catalog/query.py | 12 +++- .../fastapi_routes/v2_metadata.py | 19 ++++++ tests/unit/v2/test_metadata_query.py | 55 +++++++++++++++- tests/unit/v2/test_metadata_routes.py | 26 ++++++-- 5 files changed, 154 insertions(+), 22 deletions(-) diff --git a/policyengine_api/data/v2/catalog/parameter_tree_query.py b/policyengine_api/data/v2/catalog/parameter_tree_query.py index 5d2c25eed..cc4c856a7 100644 --- a/policyengine_api/data/v2/catalog/parameter_tree_query.py +++ b/policyengine_api/data/v2/catalog/parameter_tree_query.py @@ -18,18 +18,36 @@ def _escaped_like(value: str) -> str: return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") -def _child_path(column: object, prefix: str, dialect: str) -> object: - remainder = sa.func.substr(column, len(prefix) + 1) +def _path_segment(remainder: object, dialect: str) -> object: dot_position = ( sa.func.instr(remainder, ".") if dialect == "sqlite" else sa.func.strpos(remainder, ".") ) - segment = sa.case( + return sa.case( (dot_position > 0, sa.func.substr(remainder, 1, dot_position - 1)), else_=remainder, ) - return sa.literal(prefix) + segment + + +def _child_path(column: object, prefix: str, dialect: str) -> object: + remainder = sa.func.substr(column, len(prefix) + 1) + return sa.literal(prefix) + _path_segment(remainder, dialect) + + +def _direct_child_path( + column: object, + parent_path: object, + dialect: str, +) -> object: + remainder = sa.func.substr(column, sa.func.length(parent_path) + 2) + return parent_path + "." + _path_segment(remainder, dialect) + + +def _has_path_prefix(column: object, parent_path: object) -> object: + return ( + sa.func.substr(column, 1, sa.func.length(parent_path) + 1) == parent_path + "." + ) def parameter_children_query( @@ -56,27 +74,45 @@ def parameter_children_query( Parameter.name.like(f"{escaped_prefix}%", escape="\\"), ), ).subquery() - descendant_count = ( - select(sa.func.count(Parameter.id)) + direct_child_paths = sa.union( + select( + _direct_child_path( + ParameterNode.name, + paths.c.path, + dialect, + ).label("path") + ) .where( - Parameter.tax_benefit_model_version_id == model_version_id, - sa.func.substr( + ParameterNode.tax_benefit_model_version_id == model_version_id, + _has_path_prefix(ParameterNode.name, paths.c.path), + ) + .correlate(paths), + select( + _direct_child_path( Parameter.name, - 1, - sa.func.length(paths.c.path) + 1, - ) - == paths.c.path + ".", + paths.c.path, + dialect, + ).label("path") + ) + .where( + Parameter.tax_benefit_model_version_id == model_version_id, + _has_path_prefix(Parameter.name, paths.c.path), ) + .correlate(paths), + ).subquery() + direct_child_count = ( + select(sa.func.count()) + .select_from(direct_child_paths) .correlate(paths) .scalar_subquery() ) - is_node = sa.or_(descendant_count > 0, Parameter.id.is_(None)) + is_node = sa.or_(direct_child_count > 0, Parameter.id.is_(None)) return ( select( paths.c.path, sa.func.coalesce(ParameterNode.label, Parameter.label).label("label"), sa.case((is_node, "node"), else_="parameter").label("type"), - sa.case((is_node, descendant_count), else_=None).label("child_count"), + sa.case((is_node, direct_child_count), else_=None).label("child_count"), Parameter.id.label("parameter_id"), Parameter.label.label("parameter_label"), Parameter.description.label("parameter_description"), diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index 8e51cbd14..9d0c99608 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -450,11 +450,19 @@ def list_parameter_values( statement = statement.where(ParameterValue.parameter_id == parameter_id) if current: selected_time = now or datetime.now(timezone.utc) + if selected_time.tzinfo is None: + selected_time = selected_time.replace(tzinfo=timezone.utc) + selected_day = selected_time.astimezone(timezone.utc).replace( + hour=0, + minute=0, + second=0, + microsecond=0, + ) statement = statement.where( - ParameterValue.start_date <= selected_time, + ParameterValue.start_date <= selected_day, sa.or_( ParameterValue.end_date.is_(None), - ParameterValue.end_date > selected_time, + ParameterValue.end_date >= selected_day, ), ) rows = self._resource_rows( diff --git a/policyengine_api/fastapi_routes/v2_metadata.py b/policyengine_api/fastapi_routes/v2_metadata.py index f0ff4f7fc..000eddcdd 100644 --- a/policyengine_api/fastapi_routes/v2_metadata.py +++ b/policyengine_api/fastapi_routes/v2_metadata.py @@ -45,6 +45,25 @@ def v2_preview_openapi(request: Request) -> JSONResponse: } return JSONResponse(preview_schema) + @router.get( + "/v2", + response_model=MetadataErrorResponse, + status_code=404, + include_in_schema=False, + ) + def unsupported_v2_root() -> MetadataErrorResponse: + return MetadataErrorResponse(message="V2 metadata resource was not found") + + @router.api_route( + "/v2", + methods=["POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"], + response_model=MetadataErrorResponse, + status_code=405, + include_in_schema=False, + ) + def unsupported_v2_root_method() -> MetadataErrorResponse: + return MetadataErrorResponse(message="V2 metadata resources support GET only") + @router.get( "/v2/{resource_path:path}", response_model=MetadataErrorResponse, diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index 8d20bf6dc..e0ecadb0b 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -338,24 +338,75 @@ def test_parameter_values_are_separate_canonical_resources( assert all(item.parameter_id == parameter.id for item in all_values.items) +@pytest.mark.parametrize( + ("selected_time", "expected_value"), + [ + (datetime(2025, 12, 31, tzinfo=timezone.utc), 0.1), + (datetime(2025, 12, 31, 23, 59, tzinfo=timezone.utc), 0.1), + (datetime(2026, 1, 1, tzinfo=timezone.utc), 0.2), + ], +) +def test_current_parameter_value_uses_inclusive_effective_dates( + catalog_session: Session, + selected_time: datetime, + expected_value: float, +) -> None: + result = _service(catalog_session).list_parameter_values( + "us", + current=True, + now=selected_time, + ) + + assert [item.value for item in result.items] == [expected_value] + + def test_parameter_children_are_loaded_one_level_at_a_time( catalog_session: Session, ) -> None: + model_version = _us_model_version(catalog_session) + nested_node = ParameterNode( + id=uuid4(), + tax_benefit_model_version_id=model_version.id, + name="gov.example.nested", + label="Nested parameters", + description=None, + ) + nested_parameter = Parameter( + id=uuid4(), + tax_benefit_model_version_id=model_version.id, + name="gov.example.nested.amount", + label="Nested amount", + description=None, + data_type="float", + unit="currency-USD", + ) + catalog_session.add_all([nested_node, nested_parameter]) + catalog_session.commit() service = _service(catalog_session) root = service.list_parameter_children("us") government = service.list_parameter_children("us", parent_path="gov") example = service.list_parameter_children("us", parent_path="gov.example") + nested = service.list_parameter_children( + "us", + parent_path="gov.example.nested", + ) assert [(item.path, item.type) for item in root.items] == [("gov", "node")] assert [(item.path, item.type) for item in government.items] == [ ("gov.example", "node") ] assert [(item.path, item.type) for item in example.items] == [ - ("gov.example.rate", "parameter") + ("gov.example.nested", "node"), + ("gov.example.rate", "parameter"), + ] + assert [(item.path, item.type) for item in nested.items] == [ + ("gov.example.nested.amount", "parameter") ] assert root.items[0].child_count == 1 - assert example.items[0].parameter is not None + assert government.items[0].child_count == 2 + assert example.items[0].child_count == 1 + assert example.items[1].parameter is not None def test_resource_filters_and_details_remain_inside_selected_catalog( diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py index 1fa8dfa91..04c960969 100644 --- a/tests/unit/v2/test_metadata_routes.py +++ b/tests/unit/v2/test_metadata_routes.py @@ -334,6 +334,20 @@ def test_resource_route_forwards_version_filters_and_pagination() -> None: assert reader.closed +def test_region_code_route_decodes_percent_encoded_slash() -> None: + results, _ids = _resource_results() + reader = ResourceReader(results) + + response = _client(lambda: reader).get( + "/v2/regions/by-code/state%2Fca", + params={"country_id": "us"}, + ) + + assert response.status_code == 200 + assert reader.calls == [("get_region_by_code", ("us", "state/ca", None), {})] + assert reader.closed + + @pytest.mark.parametrize( "params", [ @@ -398,18 +412,22 @@ def factory(): return ResourceReader(results) client = _client(factory) + missing_root = client.get("/v2") missing = client.get("/v2/not-a-resource") removed_combined = client.get("/v2/us/metadata") + assert missing_root.status_code == 404 + assert missing_root.json()["status"] == "error" assert missing.status_code == 404 assert missing.json()["status"] == "error" assert removed_combined.status_code == 404 assert removed_combined.json()["status"] == "error" for method in ("POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"): - response = client.request(method, "/v2/variables") - assert response.status_code == 405 - if method != "HEAD": - assert response.json()["status"] == "error" + for path in ("/v2", "/v2/variables"): + response = client.request(method, path) + assert response.status_code == 405 + if method != "HEAD": + assert response.json()["status"] == "error" assert calls == [] From 1e99021ea953f8b10583345269a6fea26075182c Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 18:47:17 +0400 Subject: [PATCH 23/27] Split v2 metadata queries by resource --- .../data/v2/catalog/dataset_query.py | 78 +++ .../data/v2/catalog/model_query.py | 89 +++ .../data/v2/catalog/parameter_query.py | 211 +++++++ policyengine_api/data/v2/catalog/query.py | 573 ++++-------------- .../data/v2/catalog/query_support.py | 70 +++ .../data/v2/catalog/region_query.py | 177 ++++++ .../data/v2/catalog/variable_query.py | 89 +++ tests/unit/v2/test_metadata_query.py | 12 +- 8 files changed, 858 insertions(+), 441 deletions(-) create mode 100644 policyengine_api/data/v2/catalog/dataset_query.py create mode 100644 policyengine_api/data/v2/catalog/model_query.py create mode 100644 policyengine_api/data/v2/catalog/parameter_query.py create mode 100644 policyengine_api/data/v2/catalog/query_support.py create mode 100644 policyengine_api/data/v2/catalog/region_query.py create mode 100644 policyengine_api/data/v2/catalog/variable_query.py diff --git a/policyengine_api/data/v2/catalog/dataset_query.py b/policyengine_api/data/v2/catalog/dataset_query.py new file mode 100644 index 000000000..7b648b40d --- /dev/null +++ b/policyengine_api/data/v2/catalog/dataset_query.py @@ -0,0 +1,78 @@ +"""Logical input-dataset metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +from sqlmodel import Session, select + +from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from policyengine_api.data.v2.catalog.query_support import ( + MetadataResourceNotFoundError, + page_result, + query_rows, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDataset, + MetadataDetailResult, + MetadataPageResult, +) +from policyengine_api.data.v2.models import Dataset + + +def _dataset(dataset: Dataset) -> MetadataDataset: + return MetadataDataset( + id=dataset.id, + name=dataset.name, + description=dataset.description, + year=dataset.year, + ) + + +def list_datasets( + session: Session, + selected: SelectedCatalog, + *, + offset: int, + limit: int, +) -> MetadataPageResult[MetadataDataset]: + rows = query_rows( + session, + select(Dataset) + .where( + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ) + .order_by(Dataset.name) + .offset(offset) + .limit(limit + 1), + ) + return page_result( + selected, + [_dataset(row) for row in rows], + offset=offset, + limit=limit, + ) + + +def get_dataset( + session: Session, + selected: SelectedCatalog, + dataset_id: UUID, +) -> MetadataDetailResult[MetadataDataset]: + rows = query_rows( + session, + select(Dataset).where( + Dataset.id == dataset_id, + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"dataset {dataset_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_dataset(rows[0]), + ) diff --git a/policyengine_api/data/v2/catalog/model_query.py b/policyengine_api/data/v2/catalog/model_query.py new file mode 100644 index 000000000..066d140de --- /dev/null +++ b/policyengine_api/data/v2/catalog/model_query.py @@ -0,0 +1,89 @@ +"""Tax-benefit model and model-version metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from policyengine_api.data.v2.catalog.query_support import ( + MetadataResourceNotFoundError, + page_result, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDetailResult, + MetadataModel, + MetadataModelSelectionResult, + MetadataModelVersionDetail, + MetadataPageResult, +) + + +def _model(selected: SelectedCatalog) -> MetadataModel: + return MetadataModel( + id=selected.model.id, + name=selected.model.name, + description=selected.model_version.description, + ) + + +def _model_version(selected: SelectedCatalog) -> MetadataModelVersionDetail: + return MetadataModelVersionDetail( + id=selected.model_version.id, + model_id=selected.model.id, + version=selected.model_version.version, + description=selected.model_version.description, + current_law_id=selected.model_version.current_law_id, + metadata_time_periods=selected.model_version.metadata_time_periods, + ) + + +def list_models( + selected: SelectedCatalog, + *, + offset: int, + limit: int, +) -> MetadataPageResult[MetadataModel]: + rows = [_model(selected)] if offset == 0 else [] + return page_result(selected, rows, offset=offset, limit=limit) + + +def get_model( + selected: SelectedCatalog, + model_id: UUID, +) -> MetadataDetailResult[MetadataModel]: + if selected.model.id != model_id: + raise MetadataResourceNotFoundError(f"model {model_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_model(selected), + ) + + +def get_model_by_country(selected: SelectedCatalog) -> MetadataModelSelectionResult: + return MetadataModelSelectionResult( + policyengine_version=selected.policyengine_version, + model=_model(selected), + model_version=_model_version(selected), + ) + + +def list_model_versions( + selected: SelectedCatalog, + *, + offset: int, + limit: int, +) -> MetadataPageResult[MetadataModelVersionDetail]: + rows = [_model_version(selected)] if offset == 0 else [] + return page_result(selected, rows, offset=offset, limit=limit) + + +def get_model_version( + selected: SelectedCatalog, + version_id: UUID, +) -> MetadataDetailResult[MetadataModelVersionDetail]: + if selected.model_version.id != version_id: + raise MetadataResourceNotFoundError(f"model version {version_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_model_version(selected), + ) diff --git a/policyengine_api/data/v2/catalog/parameter_query.py b/policyengine_api/data/v2/catalog/parameter_query.py new file mode 100644 index 000000000..19554d0b3 --- /dev/null +++ b/policyengine_api/data/v2/catalog/parameter_query.py @@ -0,0 +1,211 @@ +"""Parameter, parameter-tree, and canonical parameter-value queries.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from uuid import UUID + +import sqlalchemy as sa +from sqlmodel import Session, select + +from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from policyengine_api.data.v2.catalog.parameter_tree_query import ( + parameter_children_from_rows, + parameter_children_query, +) +from policyengine_api.data.v2.catalog.query_support import ( + MetadataResourceNotFoundError, + escape_like, + page_result, + query_rows, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataCanonicalParameterValue, + MetadataDetailResult, + MetadataPageResult, + MetadataParameterChild, + MetadataParameterSummary, +) +from policyengine_api.data.v2.models import Parameter, ParameterValue + + +def _parameter(parameter: Parameter) -> MetadataParameterSummary: + return MetadataParameterSummary( + id=parameter.id, + name=parameter.name, + label=parameter.label, + description=parameter.description, + data_type=parameter.data_type, + unit=parameter.unit, + ) + + +def _parameter_value(value: ParameterValue) -> MetadataCanonicalParameterValue: + return MetadataCanonicalParameterValue( + id=value.id, + parameter_id=value.parameter_id, + value=value.value_json, + start_date=value.start_date, + end_date=value.end_date, + ) + + +def _utc_day_start(selected_time: datetime) -> datetime: + if selected_time.tzinfo is None: + selected_time = selected_time.replace(tzinfo=timezone.utc) + return selected_time.astimezone(timezone.utc).replace( + hour=0, + minute=0, + second=0, + microsecond=0, + ) + + +def list_parameters( + session: Session, + selected: SelectedCatalog, + *, + offset: int, + limit: int, + search: str | None, +) -> MetadataPageResult[MetadataParameterSummary]: + statement = select(Parameter).where( + Parameter.tax_benefit_model_version_id == selected.model_version.id + ) + if search: + pattern = f"%{escape_like(search)}%" + statement = statement.where( + sa.or_( + Parameter.name.ilike(pattern, escape="\\"), + Parameter.label.ilike(pattern, escape="\\"), + Parameter.description.ilike(pattern, escape="\\"), + ) + ) + rows = query_rows( + session, + statement.order_by(Parameter.name).offset(offset).limit(limit + 1), + ) + return page_result( + selected, + [_parameter(row) for row in rows], + offset=offset, + limit=limit, + ) + + +def get_parameter( + session: Session, + selected: SelectedCatalog, + parameter_id: UUID, +) -> MetadataDetailResult[MetadataParameterSummary]: + rows = query_rows( + session, + select(Parameter).where( + Parameter.id == parameter_id, + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"parameter {parameter_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_parameter(rows[0]), + ) + + +def list_parameter_children( + session: Session, + selected: SelectedCatalog, + *, + parent_path: str, + offset: int, + limit: int, +) -> MetadataPageResult[MetadataParameterChild]: + rows = query_rows( + session, + parameter_children_query( + model_version_id=selected.model_version.id, + parent_path=parent_path, + dialect=session.get_bind().dialect.name, + offset=offset, + limit=limit, + ), + ) + return page_result( + selected, + parameter_children_from_rows(rows), + offset=offset, + limit=limit, + ) + + +def list_parameter_values( + session: Session, + selected: SelectedCatalog, + *, + parameter_id: UUID | None, + current: bool, + offset: int, + limit: int, + now: datetime | None, +) -> MetadataPageResult[MetadataCanonicalParameterValue]: + statement = ( + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ) + ) + if parameter_id is not None: + statement = statement.where(ParameterValue.parameter_id == parameter_id) + if current: + selected_day = _utc_day_start(now or datetime.now(timezone.utc)) + statement = statement.where( + ParameterValue.start_date <= selected_day, + sa.or_( + ParameterValue.end_date.is_(None), + ParameterValue.end_date >= selected_day, + ), + ) + rows = query_rows( + session, + statement.order_by( + Parameter.name, + ParameterValue.start_date.desc(), + ParameterValue.id, + ) + .offset(offset) + .limit(limit + 1), + ) + return page_result( + selected, + [_parameter_value(row) for row in rows], + offset=offset, + limit=limit, + ) + + +def get_parameter_value( + session: Session, + selected: SelectedCatalog, + value_id: UUID, +) -> MetadataDetailResult[MetadataCanonicalParameterValue]: + rows = query_rows( + session, + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + ParameterValue.id == value_id, + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"parameter value {value_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_parameter_value(rows[0]), + ) diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index 9d0c99608..0dc4f6e07 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -1,16 +1,19 @@ -"""Read-only v2 catalog queries and deterministic response serialization.""" +"""Route-facing service for read-only v2 metadata resource queries.""" from __future__ import annotations -from datetime import datetime, timezone -from typing import TypeVar +from datetime import datetime from uuid import UUID -import sqlalchemy as sa -from sqlalchemy.exc import SQLAlchemyError -from sqlmodel import Session, select +from sqlmodel import Session -from policyengine_api.dataset_display import get_dataset_display_label +from policyengine_api.data.v2.catalog import ( + dataset_query, + model_query, + parameter_query, + region_query, + variable_query, +) from policyengine_api.data.v2.catalog.catalog_selection import ( InvalidPolicyEngineVersionError, MetadataCatalogUnavailableError, @@ -20,14 +23,14 @@ select_catalog, validate_policyengine_version, ) -from policyengine_api.data.v2.catalog.parameter_tree_query import ( - parameter_children_from_rows, - parameter_children_query, +from policyengine_api.data.v2.catalog.query_support import ( + InvalidMetadataPageError, + MetadataResourceNotFoundError, + validate_metadata_page, ) from policyengine_api.data.v2.catalog.schemas import ( MetadataCanonicalParameterValue, MetadataDataset, - MetadataDatasetOption, MetadataDetailResult, MetadataEconomyOptionsResult, MetadataModel, @@ -37,17 +40,8 @@ MetadataParameterChild, MetadataParameterSummary, MetadataRegion, - MetadataRegionOption, - MetadataTimePeriodOption, MetadataVariable, ) -from policyengine_api.data.v2.models import ( - Dataset, - Parameter, - ParameterValue, - Region, - Variable, -) __all__ = [ @@ -63,130 +57,8 @@ ] -class MetadataResourceNotFoundError(LookupError): - """Raised when a selected catalog does not contain a requested resource.""" - - -class InvalidMetadataPageError(ValueError): - """Raised when collection pagination is outside the documented bounds.""" - - -ResourceT = TypeVar("ResourceT") - - -def _page( - selected: SelectedCatalog, - rows: list[ResourceT], - *, - offset: int, - limit: int, -) -> MetadataPageResult[ResourceT]: - return MetadataPageResult( - policyengine_version=selected.policyengine_version, - items=rows[:limit], - offset=offset, - limit=limit, - has_more=len(rows) > limit, - ) - - -def validate_metadata_page(offset: int, limit: int) -> tuple[int, int]: - if offset < 0: - raise InvalidMetadataPageError("offset must be at least 0") - if not 1 <= limit <= 500: - raise InvalidMetadataPageError("limit must be between 1 and 500") - return offset, limit - - -def _metadata_model(selected: SelectedCatalog) -> MetadataModel: - return MetadataModel( - id=selected.model.id, - name=selected.model.name, - description=selected.model_version.description, - ) - - -def _metadata_model_version(selected: SelectedCatalog) -> MetadataModelVersionDetail: - return MetadataModelVersionDetail( - id=selected.model_version.id, - model_id=selected.model.id, - version=selected.model_version.version, - description=selected.model_version.description, - current_law_id=selected.model_version.current_law_id, - metadata_time_periods=selected.model_version.metadata_time_periods, - ) - - -def _metadata_variable(variable: Variable) -> MetadataVariable: - return MetadataVariable( - id=variable.id, - name=variable.name, - label=variable.label, - entity=variable.entity, - description=variable.description, - data_type=variable.data_type, - possible_values=variable.possible_values, - default_value=variable.default_value, - adds=variable.adds, - subtracts=variable.subtracts, - ) - - -def _metadata_parameter(parameter: Parameter) -> MetadataParameterSummary: - return MetadataParameterSummary( - id=parameter.id, - name=parameter.name, - label=parameter.label, - description=parameter.description, - data_type=parameter.data_type, - unit=parameter.unit, - ) - - -def _metadata_parameter_value( - value: ParameterValue, -) -> MetadataCanonicalParameterValue: - return MetadataCanonicalParameterValue( - id=value.id, - parameter_id=value.parameter_id, - value=value.value_json, - start_date=value.start_date, - end_date=value.end_date, - ) - - -def _metadata_dataset(dataset: Dataset) -> MetadataDataset: - return MetadataDataset( - id=dataset.id, - name=dataset.name, - description=dataset.description, - year=dataset.year, - ) - - -def _metadata_region(region: Region) -> MetadataRegion: - return MetadataRegion( - id=region.id, - code=region.code, - label=region.label, - region_type=region.region_type.value, - requires_filter=region.requires_filter, - filter_field=region.filter_field, - filter_value=region.filter_value, - filter_strategy=region.filter_strategy, - parent_code=region.parent_code, - state_code=region.state_code, - state_name=region.state_name, - default_dataset_id=region.default_dataset_id, - ) - - -def _escaped_like(value: str) -> str: - return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") - - class V2MetadataQueryService: - """Assemble preview metadata using only an injected v2 read session.""" + """Select one catalog and delegate each resource to its query module.""" def __init__(self, session: Session, *, running_policyengine_version: str): self._session = session @@ -213,13 +85,16 @@ def select_catalog( policyengine_version=policyengine_version, ) - def _resource_rows(self, statement: object) -> list: - try: - return list(self._session.exec(statement).all()) - except SQLAlchemyError as error: - raise MetadataCatalogUnavailableError( - "the v2 metadata catalog cannot be queried" - ) from error + def _select_paginated_catalog( + self, + country_id: str, + policyengine_version: str | None, + *, + offset: int, + limit: int, + ) -> SelectedCatalog: + validate_metadata_page(offset, limit) + return self.select_catalog(country_id, policyengine_version) def list_models( self, @@ -229,10 +104,16 @@ def list_models( offset: int = 0, limit: int = 100, ) -> MetadataPageResult[MetadataModel]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - rows = [_metadata_model(selected)] if offset == 0 else [] - return _page(selected, rows, offset=offset, limit=limit) + return model_query.list_models( + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + offset=offset, + limit=limit, + ) def get_model( self, @@ -240,12 +121,9 @@ def get_model( model_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataModel]: - selected = self.select_catalog(country_id, policyengine_version) - if selected.model.id != model_id: - raise MetadataResourceNotFoundError(f"model {model_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_model(selected), + return model_query.get_model( + self.select_catalog(country_id, policyengine_version), + model_id, ) def get_model_by_country( @@ -253,11 +131,8 @@ def get_model_by_country( country_id: str, policyengine_version: str | None = None, ) -> MetadataModelSelectionResult: - selected = self.select_catalog(country_id, policyengine_version) - return MetadataModelSelectionResult( - policyengine_version=selected.policyengine_version, - model=_metadata_model(selected), - model_version=_metadata_model_version(selected), + return model_query.get_model_by_country( + self.select_catalog(country_id, policyengine_version) ) def list_model_versions( @@ -268,10 +143,16 @@ def list_model_versions( offset: int = 0, limit: int = 100, ) -> MetadataPageResult[MetadataModelVersionDetail]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - rows = [_metadata_model_version(selected)] if offset == 0 else [] - return _page(selected, rows, offset=offset, limit=limit) + return model_query.list_model_versions( + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + offset=offset, + limit=limit, + ) def get_model_version( self, @@ -279,14 +160,9 @@ def get_model_version( version_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataModelVersionDetail]: - selected = self.select_catalog(country_id, policyengine_version) - if selected.model_version.id != version_id: - raise MetadataResourceNotFoundError( - f"model version {version_id} was not found" - ) - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_model_version(selected), + return model_query.get_model_version( + self.select_catalog(country_id, policyengine_version), + version_id, ) def list_variables( @@ -298,28 +174,17 @@ def list_variables( limit: int = 100, search: str | None = None, ) -> MetadataPageResult[MetadataVariable]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - statement = select(Variable).where( - Variable.tax_benefit_model_version_id == selected.model_version.id - ) - if search: - pattern = f"%{_escaped_like(search)}%" - statement = statement.where( - sa.or_( - Variable.name.ilike(pattern, escape="\\"), - Variable.label.ilike(pattern, escape="\\"), - Variable.description.ilike(pattern, escape="\\"), - ) - ) - rows = self._resource_rows( - statement.order_by(Variable.name).offset(offset).limit(limit + 1) - ) - return _page( - selected, - [_metadata_variable(row) for row in rows], + return variable_query.list_variables( + self._session, + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), offset=offset, limit=limit, + search=search, ) def get_variable( @@ -328,18 +193,10 @@ def get_variable( variable_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataVariable]: - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(Variable).where( - Variable.id == variable_id, - Variable.tax_benefit_model_version_id == selected.model_version.id, - ) - ) - if not rows: - raise MetadataResourceNotFoundError(f"variable {variable_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_variable(rows[0]), + return variable_query.get_variable( + self._session, + self.select_catalog(country_id, policyengine_version), + variable_id, ) def list_parameters( @@ -351,28 +208,17 @@ def list_parameters( limit: int = 100, search: str | None = None, ) -> MetadataPageResult[MetadataParameterSummary]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - statement = select(Parameter).where( - Parameter.tax_benefit_model_version_id == selected.model_version.id - ) - if search: - pattern = f"%{_escaped_like(search)}%" - statement = statement.where( - sa.or_( - Parameter.name.ilike(pattern, escape="\\"), - Parameter.label.ilike(pattern, escape="\\"), - Parameter.description.ilike(pattern, escape="\\"), - ) - ) - rows = self._resource_rows( - statement.order_by(Parameter.name).offset(offset).limit(limit + 1) - ) - return _page( - selected, - [_metadata_parameter(row) for row in rows], + return parameter_query.list_parameters( + self._session, + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), offset=offset, limit=limit, + search=search, ) def get_parameter( @@ -381,20 +227,10 @@ def get_parameter( parameter_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataParameterSummary]: - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(Parameter).where( - Parameter.id == parameter_id, - Parameter.tax_benefit_model_version_id == selected.model_version.id, - ) - ) - if not rows: - raise MetadataResourceNotFoundError( - f"parameter {parameter_id} was not found" - ) - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_parameter(rows[0]), + return parameter_query.get_parameter( + self._session, + self.select_catalog(country_id, policyengine_version), + parameter_id, ) def list_parameter_children( @@ -406,20 +242,15 @@ def list_parameter_children( offset: int = 0, limit: int = 100, ) -> MetadataPageResult[MetadataParameterChild]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - parameter_children_query( - model_version_id=selected.model_version.id, - parent_path=parent_path, - dialect=self._session.get_bind().dialect.name, + return parameter_query.list_parameter_children( + self._session, + self._select_paginated_catalog( + country_id, + policyengine_version, offset=offset, limit=limit, - ) - ) - return _page( - selected, - parameter_children_from_rows(rows), + ), + parent_path=parent_path, offset=offset, limit=limit, ) @@ -435,50 +266,19 @@ def list_parameter_values( limit: int = 100, now: datetime | None = None, ) -> MetadataPageResult[MetadataCanonicalParameterValue]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - statement = ( - select(ParameterValue) - .join(Parameter, Parameter.id == ParameterValue.parameter_id) - .where( - Parameter.tax_benefit_model_version_id == selected.model_version.id, - ParameterValue.policy_id.is_(None), - ParameterValue.dynamic_id.is_(None), - ) - ) - if parameter_id is not None: - statement = statement.where(ParameterValue.parameter_id == parameter_id) - if current: - selected_time = now or datetime.now(timezone.utc) - if selected_time.tzinfo is None: - selected_time = selected_time.replace(tzinfo=timezone.utc) - selected_day = selected_time.astimezone(timezone.utc).replace( - hour=0, - minute=0, - second=0, - microsecond=0, - ) - statement = statement.where( - ParameterValue.start_date <= selected_day, - sa.or_( - ParameterValue.end_date.is_(None), - ParameterValue.end_date >= selected_day, - ), - ) - rows = self._resource_rows( - statement.order_by( - Parameter.name, - ParameterValue.start_date.desc(), - ParameterValue.id, - ) - .offset(offset) - .limit(limit + 1) - ) - return _page( - selected, - [_metadata_parameter_value(row) for row in rows], + return parameter_query.list_parameter_values( + self._session, + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + parameter_id=parameter_id, + current=current, offset=offset, limit=limit, + now=now, ) def get_parameter_value( @@ -487,24 +287,10 @@ def get_parameter_value( value_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataCanonicalParameterValue]: - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(ParameterValue) - .join(Parameter, Parameter.id == ParameterValue.parameter_id) - .where( - ParameterValue.id == value_id, - Parameter.tax_benefit_model_version_id == selected.model_version.id, - ParameterValue.policy_id.is_(None), - ParameterValue.dynamic_id.is_(None), - ) - ) - if not rows: - raise MetadataResourceNotFoundError( - f"parameter value {value_id} was not found" - ) - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_parameter_value(rows[0]), + return parameter_query.get_parameter_value( + self._session, + self.select_catalog(country_id, policyengine_version), + value_id, ) def list_datasets( @@ -515,22 +301,14 @@ def list_datasets( offset: int = 0, limit: int = 100, ) -> MetadataPageResult[MetadataDataset]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(Dataset) - .where( - Dataset.tax_benefit_model_version_id == selected.model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ) - .order_by(Dataset.name) - .offset(offset) - .limit(limit + 1) - ) - return _page( - selected, - [_metadata_dataset(row) for row in rows], + return dataset_query.list_datasets( + self._session, + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), offset=offset, limit=limit, ) @@ -541,20 +319,10 @@ def get_dataset( dataset_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataDataset]: - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(Dataset).where( - Dataset.id == dataset_id, - Dataset.tax_benefit_model_version_id == selected.model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ) - ) - if not rows: - raise MetadataResourceNotFoundError(f"dataset {dataset_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_dataset(rows[0]), + return dataset_query.get_dataset( + self._session, + self.select_catalog(country_id, policyengine_version), + dataset_id, ) def list_regions( @@ -566,19 +334,15 @@ def list_regions( offset: int = 0, limit: int = 100, ) -> MetadataPageResult[MetadataRegion]: - validate_metadata_page(offset, limit) - selected = self.select_catalog(country_id, policyengine_version) - statement = select(Region).where( - Region.tax_benefit_model_version_id == selected.model_version.id - ) - if region_type is not None: - statement = statement.where(Region.region_type == region_type) - rows = self._resource_rows( - statement.order_by(Region.code).offset(offset).limit(limit + 1) - ) - return _page( - selected, - [_metadata_region(row) for row in rows], + return region_query.list_regions( + self._session, + self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ), + region_type=region_type, offset=offset, limit=limit, ) @@ -589,18 +353,10 @@ def get_region( region_id: UUID, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataRegion]: - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(Region).where( - Region.id == region_id, - Region.tax_benefit_model_version_id == selected.model_version.id, - ) - ) - if not rows: - raise MetadataResourceNotFoundError(f"region {region_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_region(rows[0]), + return region_query.get_region( + self._session, + self.select_catalog(country_id, policyengine_version), + region_id, ) def get_region_by_code( @@ -609,18 +365,10 @@ def get_region_by_code( region_code: str, policyengine_version: str | None = None, ) -> MetadataDetailResult[MetadataRegion]: - selected = self.select_catalog(country_id, policyengine_version) - rows = self._resource_rows( - select(Region).where( - Region.code == region_code, - Region.tax_benefit_model_version_id == selected.model_version.id, - ) - ) - if not rows: - raise MetadataResourceNotFoundError(f"region {region_code!r} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_metadata_region(rows[0]), + return region_query.get_region_by_code( + self._session, + self.select_catalog(country_id, policyengine_version), + region_code, ) def get_economy_options( @@ -628,62 +376,7 @@ def get_economy_options( country_id: str, policyengine_version: str | None = None, ) -> MetadataEconomyOptionsResult: - selected = self.select_catalog(country_id, policyengine_version) - regions = self._resource_rows( - select(Region) - .where(Region.tax_benefit_model_version_id == selected.model_version.id) - .order_by(Region.code) - ) - national_region = next( - (region for region in regions if region.code == country_id), - None, - ) - if national_region is None: - raise MetadataCatalogUnavailableError( - f"the {country_id} national v2 region is absent" - ) - datasets = self._resource_rows( - select(Dataset).where( - Dataset.id == national_region.default_dataset_id, - Dataset.tax_benefit_model_version_id == selected.model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ) - ) - if len(datasets) != 1: - raise MetadataCatalogUnavailableError( - f"the {country_id} national v2 dataset is absent" - ) - time_periods = selected.model_version.metadata_time_periods - if ( - not isinstance(selected.model_version.current_law_id, int) - or not isinstance(time_periods, list) - or not time_periods - or any(not isinstance(year, int) for year in time_periods) - ): - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 model-version options are incomplete" - ) - national_dataset = datasets[0] - return MetadataEconomyOptionsResult( - policyengine_version=selected.policyengine_version, - current_law_id=selected.model_version.current_law_id, - region=[ - MetadataRegionOption( - name=region.code, - label=region.label, - type=region.region_type.value, - ) - for region in regions - ], - time_period=[ - MetadataTimePeriodOption(name=year, label=str(year)) - for year in time_periods - ], - datasets=[ - MetadataDatasetOption( - name=national_dataset.name, - label=get_dataset_display_label(national_dataset.name), - ) - ], + return region_query.get_economy_options( + self._session, + self.select_catalog(country_id, policyengine_version), ) diff --git a/policyengine_api/data/v2/catalog/query_support.py b/policyengine_api/data/v2/catalog/query_support.py new file mode 100644 index 000000000..90ee933fc --- /dev/null +++ b/policyengine_api/data/v2/catalog/query_support.py @@ -0,0 +1,70 @@ +"""Shared execution and pagination for v2 metadata resource queries.""" + +from __future__ import annotations + +from typing import TypeVar + +from sqlalchemy.exc import SQLAlchemyError +from sqlmodel import Session + +from policyengine_api.data.v2.catalog.catalog_selection import ( + MetadataCatalogUnavailableError, + SelectedCatalog, +) +from policyengine_api.data.v2.catalog.schemas import MetadataPageResult + + +class MetadataResourceNotFoundError(LookupError): + """Raised when a selected catalog does not contain a requested resource.""" + + +class InvalidMetadataPageError(ValueError): + """Raised when collection pagination is outside the documented bounds.""" + + +ResourceT = TypeVar("ResourceT") + + +def page_result( + selected: SelectedCatalog, + rows: list[ResourceT], + *, + offset: int, + limit: int, +) -> MetadataPageResult[ResourceT]: + """Return one bounded response page from a limit-plus-one query.""" + + return MetadataPageResult( + policyengine_version=selected.policyengine_version, + items=rows[:limit], + offset=offset, + limit=limit, + has_more=len(rows) > limit, + ) + + +def validate_metadata_page(offset: int, limit: int) -> tuple[int, int]: + """Validate the shared v2 metadata collection bounds.""" + + if offset < 0: + raise InvalidMetadataPageError("offset must be at least 0") + if not 1 <= limit <= 500: + raise InvalidMetadataPageError("limit must be between 1 and 500") + return offset, limit + + +def query_rows(session: Session, statement: object) -> list: + """Execute one read statement and translate database failures.""" + + try: + return list(session.exec(statement).all()) + except SQLAlchemyError as error: + raise MetadataCatalogUnavailableError( + "the v2 metadata catalog cannot be queried" + ) from error + + +def escape_like(value: str) -> str: + """Escape SQL LIKE wildcard characters in a literal search value.""" + + return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") diff --git a/policyengine_api/data/v2/catalog/region_query.py b/policyengine_api/data/v2/catalog/region_query.py new file mode 100644 index 000000000..482c0d492 --- /dev/null +++ b/policyengine_api/data/v2/catalog/region_query.py @@ -0,0 +1,177 @@ +"""Region and economy-option metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +from sqlmodel import Session, select + +from policyengine_api.dataset_display import get_dataset_display_label +from policyengine_api.data.v2.catalog.catalog_selection import ( + MetadataCatalogUnavailableError, + SelectedCatalog, +) +from policyengine_api.data.v2.catalog.query_support import ( + MetadataResourceNotFoundError, + page_result, + query_rows, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDatasetOption, + MetadataDetailResult, + MetadataEconomyOptionsResult, + MetadataPageResult, + MetadataRegion, + MetadataRegionOption, + MetadataTimePeriodOption, +) +from policyengine_api.data.v2.models import Dataset, Region + + +def _region(region: Region) -> MetadataRegion: + return MetadataRegion( + id=region.id, + code=region.code, + label=region.label, + region_type=region.region_type.value, + requires_filter=region.requires_filter, + filter_field=region.filter_field, + filter_value=region.filter_value, + filter_strategy=region.filter_strategy, + parent_code=region.parent_code, + state_code=region.state_code, + state_name=region.state_name, + default_dataset_id=region.default_dataset_id, + ) + + +def list_regions( + session: Session, + selected: SelectedCatalog, + *, + region_type: str | None, + offset: int, + limit: int, +) -> MetadataPageResult[MetadataRegion]: + statement = select(Region).where( + Region.tax_benefit_model_version_id == selected.model_version.id + ) + if region_type is not None: + statement = statement.where(Region.region_type == region_type) + rows = query_rows( + session, + statement.order_by(Region.code).offset(offset).limit(limit + 1), + ) + return page_result( + selected, + [_region(row) for row in rows], + offset=offset, + limit=limit, + ) + + +def get_region( + session: Session, + selected: SelectedCatalog, + region_id: UUID, +) -> MetadataDetailResult[MetadataRegion]: + rows = query_rows( + session, + select(Region).where( + Region.id == region_id, + Region.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"region {region_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_region(rows[0]), + ) + + +def get_region_by_code( + session: Session, + selected: SelectedCatalog, + region_code: str, +) -> MetadataDetailResult[MetadataRegion]: + rows = query_rows( + session, + select(Region).where( + Region.code == region_code, + Region.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"region {region_code!r} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_region(rows[0]), + ) + + +def get_economy_options( + session: Session, + selected: SelectedCatalog, +) -> MetadataEconomyOptionsResult: + country_id = selected.country_id + regions = query_rows( + session, + select(Region) + .where(Region.tax_benefit_model_version_id == selected.model_version.id) + .order_by(Region.code), + ) + national_region = next( + (region for region in regions if region.code == country_id), + None, + ) + if national_region is None: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 region is absent" + ) + datasets = query_rows( + session, + select(Dataset).where( + Dataset.id == national_region.default_dataset_id, + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ), + ) + if len(datasets) != 1: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 dataset is absent" + ) + time_periods = selected.model_version.metadata_time_periods + if ( + not isinstance(selected.model_version.current_law_id, int) + or not isinstance(time_periods, list) + or not time_periods + or any(not isinstance(year, int) for year in time_periods) + ): + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 model-version options are incomplete" + ) + national_dataset = datasets[0] + return MetadataEconomyOptionsResult( + policyengine_version=selected.policyengine_version, + current_law_id=selected.model_version.current_law_id, + region=[ + MetadataRegionOption( + name=region.code, + label=region.label, + type=region.region_type.value, + ) + for region in regions + ], + time_period=[ + MetadataTimePeriodOption(name=year, label=str(year)) + for year in time_periods + ], + datasets=[ + MetadataDatasetOption( + name=national_dataset.name, + label=get_dataset_display_label(national_dataset.name), + ) + ], + ) diff --git a/policyengine_api/data/v2/catalog/variable_query.py b/policyengine_api/data/v2/catalog/variable_query.py new file mode 100644 index 000000000..b1b9409e1 --- /dev/null +++ b/policyengine_api/data/v2/catalog/variable_query.py @@ -0,0 +1,89 @@ +"""Variable metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +import sqlalchemy as sa +from sqlmodel import Session, select + +from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from policyengine_api.data.v2.catalog.query_support import ( + MetadataResourceNotFoundError, + escape_like, + page_result, + query_rows, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataDetailResult, + MetadataPageResult, + MetadataVariable, +) +from policyengine_api.data.v2.models import Variable + + +def _variable(variable: Variable) -> MetadataVariable: + return MetadataVariable( + id=variable.id, + name=variable.name, + label=variable.label, + entity=variable.entity, + description=variable.description, + data_type=variable.data_type, + possible_values=variable.possible_values, + default_value=variable.default_value, + adds=variable.adds, + subtracts=variable.subtracts, + ) + + +def list_variables( + session: Session, + selected: SelectedCatalog, + *, + offset: int, + limit: int, + search: str | None, +) -> MetadataPageResult[MetadataVariable]: + statement = select(Variable).where( + Variable.tax_benefit_model_version_id == selected.model_version.id + ) + if search: + pattern = f"%{escape_like(search)}%" + statement = statement.where( + sa.or_( + Variable.name.ilike(pattern, escape="\\"), + Variable.label.ilike(pattern, escape="\\"), + Variable.description.ilike(pattern, escape="\\"), + ) + ) + rows = query_rows( + session, + statement.order_by(Variable.name).offset(offset).limit(limit + 1), + ) + return page_result( + selected, + [_variable(row) for row in rows], + offset=offset, + limit=limit, + ) + + +def get_variable( + session: Session, + selected: SelectedCatalog, + variable_id: UUID, +) -> MetadataDetailResult[MetadataVariable]: + rows = query_rows( + session, + select(Variable).where( + Variable.id == variable_id, + Variable.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"variable {variable_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_variable(rows[0]), + ) diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index e0ecadb0b..f6c17debc 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -555,7 +555,17 @@ def test_query_modules_import_no_policyengine_or_v1_metadata_source() -> None: source_directory = ( Path(__file__).parents[3] / "policyengine_api" / "data" / "v2" / "catalog" ) - modules = ("query.py", "catalog_selection.py", "parameter_tree_query.py") + modules = ( + "catalog_selection.py", + "dataset_query.py", + "model_query.py", + "parameter_query.py", + "parameter_tree_query.py", + "query.py", + "query_support.py", + "region_query.py", + "variable_query.py", + ) imported = set() for module in modules: tree = ast.parse((source_directory / module).read_text(encoding="utf-8")) From a5315eb79f97977dff493ac9babfece48ace2a0c Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 18:56:42 +0400 Subject: [PATCH 24/27] Move metadata service methods into resource modules --- .../data/v2/catalog/dataset_query.py | 106 ++--- .../data/v2/catalog/model_query.py | 129 +++--- .../data/v2/catalog/parameter_query.py | 305 ++++++++------- policyengine_api/data/v2/catalog/query.py | 366 +----------------- .../data/v2/catalog/query_support.py | 42 ++ .../data/v2/catalog/region_query.py | 258 ++++++------ .../data/v2/catalog/variable_query.py | 110 +++--- tests/unit/v2/test_metadata_query.py | 27 ++ 8 files changed, 581 insertions(+), 762 deletions(-) diff --git a/policyengine_api/data/v2/catalog/dataset_query.py b/policyengine_api/data/v2/catalog/dataset_query.py index 7b648b40d..755d5b006 100644 --- a/policyengine_api/data/v2/catalog/dataset_query.py +++ b/policyengine_api/data/v2/catalog/dataset_query.py @@ -4,10 +4,9 @@ from uuid import UUID -from sqlmodel import Session, select - -from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from sqlmodel import select from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, MetadataResourceNotFoundError, page_result, query_rows, @@ -29,50 +28,61 @@ def _dataset(dataset: Dataset) -> MetadataDataset: ) -def list_datasets( - session: Session, - selected: SelectedCatalog, - *, - offset: int, - limit: int, -) -> MetadataPageResult[MetadataDataset]: - rows = query_rows( - session, - select(Dataset) - .where( - Dataset.tax_benefit_model_version_id == selected.model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ) - .order_by(Dataset.name) - .offset(offset) - .limit(limit + 1), - ) - return page_result( - selected, - [_dataset(row) for row in rows], - offset=offset, - limit=limit, - ) +class DatasetQueryMethods(MetadataQueryContext): + """Route-facing logical input-dataset query methods.""" + def list_datasets( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataDataset]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + rows = query_rows( + self._session, + select(Dataset) + .where( + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ) + .order_by(Dataset.name) + .offset(offset) + .limit(limit + 1), + ) + return page_result( + selected, + [_dataset(row) for row in rows], + offset=offset, + limit=limit, + ) -def get_dataset( - session: Session, - selected: SelectedCatalog, - dataset_id: UUID, -) -> MetadataDetailResult[MetadataDataset]: - rows = query_rows( - session, - select(Dataset).where( - Dataset.id == dataset_id, - Dataset.tax_benefit_model_version_id == selected.model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ), - ) - if not rows: - raise MetadataResourceNotFoundError(f"dataset {dataset_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_dataset(rows[0]), - ) + def get_dataset( + self, + country_id: str, + dataset_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataDataset]: + selected = self.select_catalog(country_id, policyengine_version) + rows = query_rows( + self._session, + select(Dataset).where( + Dataset.id == dataset_id, + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"dataset {dataset_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_dataset(rows[0]), + ) diff --git a/policyengine_api/data/v2/catalog/model_query.py b/policyengine_api/data/v2/catalog/model_query.py index 066d140de..17a0e130b 100644 --- a/policyengine_api/data/v2/catalog/model_query.py +++ b/policyengine_api/data/v2/catalog/model_query.py @@ -6,6 +6,7 @@ from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, MetadataResourceNotFoundError, page_result, ) @@ -37,53 +38,81 @@ def _model_version(selected: SelectedCatalog) -> MetadataModelVersionDetail: ) -def list_models( - selected: SelectedCatalog, - *, - offset: int, - limit: int, -) -> MetadataPageResult[MetadataModel]: - rows = [_model(selected)] if offset == 0 else [] - return page_result(selected, rows, offset=offset, limit=limit) - - -def get_model( - selected: SelectedCatalog, - model_id: UUID, -) -> MetadataDetailResult[MetadataModel]: - if selected.model.id != model_id: - raise MetadataResourceNotFoundError(f"model {model_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_model(selected), - ) - - -def get_model_by_country(selected: SelectedCatalog) -> MetadataModelSelectionResult: - return MetadataModelSelectionResult( - policyengine_version=selected.policyengine_version, - model=_model(selected), - model_version=_model_version(selected), - ) - - -def list_model_versions( - selected: SelectedCatalog, - *, - offset: int, - limit: int, -) -> MetadataPageResult[MetadataModelVersionDetail]: - rows = [_model_version(selected)] if offset == 0 else [] - return page_result(selected, rows, offset=offset, limit=limit) - - -def get_model_version( - selected: SelectedCatalog, - version_id: UUID, -) -> MetadataDetailResult[MetadataModelVersionDetail]: - if selected.model_version.id != version_id: - raise MetadataResourceNotFoundError(f"model version {version_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_model_version(selected), - ) +class ModelQueryMethods(MetadataQueryContext): + """Route-facing model and model-version query methods.""" + + def list_models( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataModel]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + rows = [_model(selected)] if offset == 0 else [] + return page_result(selected, rows, offset=offset, limit=limit) + + def get_model( + self, + country_id: str, + model_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataModel]: + selected = self.select_catalog(country_id, policyengine_version) + if selected.model.id != model_id: + raise MetadataResourceNotFoundError(f"model {model_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_model(selected), + ) + + def get_model_by_country( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataModelSelectionResult: + selected = self.select_catalog(country_id, policyengine_version) + return MetadataModelSelectionResult( + policyengine_version=selected.policyengine_version, + model=_model(selected), + model_version=_model_version(selected), + ) + + def list_model_versions( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataModelVersionDetail]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + rows = [_model_version(selected)] if offset == 0 else [] + return page_result(selected, rows, offset=offset, limit=limit) + + def get_model_version( + self, + country_id: str, + version_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataModelVersionDetail]: + selected = self.select_catalog(country_id, policyengine_version) + if selected.model_version.id != version_id: + raise MetadataResourceNotFoundError( + f"model version {version_id} was not found" + ) + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_model_version(selected), + ) diff --git a/policyengine_api/data/v2/catalog/parameter_query.py b/policyengine_api/data/v2/catalog/parameter_query.py index 19554d0b3..c3807f1f2 100644 --- a/policyengine_api/data/v2/catalog/parameter_query.py +++ b/policyengine_api/data/v2/catalog/parameter_query.py @@ -6,14 +6,13 @@ from uuid import UUID import sqlalchemy as sa -from sqlmodel import Session, select - -from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from sqlmodel import select from policyengine_api.data.v2.catalog.parameter_tree_query import ( parameter_children_from_rows, parameter_children_query, ) from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, MetadataResourceNotFoundError, escape_like, page_result, @@ -61,151 +60,179 @@ def _utc_day_start(selected_time: datetime) -> datetime: ) -def list_parameters( - session: Session, - selected: SelectedCatalog, - *, - offset: int, - limit: int, - search: str | None, -) -> MetadataPageResult[MetadataParameterSummary]: - statement = select(Parameter).where( - Parameter.tax_benefit_model_version_id == selected.model_version.id - ) - if search: - pattern = f"%{escape_like(search)}%" - statement = statement.where( - sa.or_( - Parameter.name.ilike(pattern, escape="\\"), - Parameter.label.ilike(pattern, escape="\\"), - Parameter.description.ilike(pattern, escape="\\"), +class ParameterQueryMethods(MetadataQueryContext): + """Route-facing parameter, tree, and canonical-value query methods.""" + + def list_parameters( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + search: str | None = None, + ) -> MetadataPageResult[MetadataParameterSummary]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + statement = select(Parameter).where( + Parameter.tax_benefit_model_version_id == selected.model_version.id + ) + if search: + pattern = f"%{escape_like(search)}%" + statement = statement.where( + sa.or_( + Parameter.name.ilike(pattern, escape="\\"), + Parameter.label.ilike(pattern, escape="\\"), + Parameter.description.ilike(pattern, escape="\\"), + ) ) + rows = query_rows( + self._session, + statement.order_by(Parameter.name).offset(offset).limit(limit + 1), ) - rows = query_rows( - session, - statement.order_by(Parameter.name).offset(offset).limit(limit + 1), - ) - return page_result( - selected, - [_parameter(row) for row in rows], - offset=offset, - limit=limit, - ) - - -def get_parameter( - session: Session, - selected: SelectedCatalog, - parameter_id: UUID, -) -> MetadataDetailResult[MetadataParameterSummary]: - rows = query_rows( - session, - select(Parameter).where( - Parameter.id == parameter_id, - Parameter.tax_benefit_model_version_id == selected.model_version.id, - ), - ) - if not rows: - raise MetadataResourceNotFoundError(f"parameter {parameter_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_parameter(rows[0]), - ) - - -def list_parameter_children( - session: Session, - selected: SelectedCatalog, - *, - parent_path: str, - offset: int, - limit: int, -) -> MetadataPageResult[MetadataParameterChild]: - rows = query_rows( - session, - parameter_children_query( - model_version_id=selected.model_version.id, - parent_path=parent_path, - dialect=session.get_bind().dialect.name, + return page_result( + selected, + [_parameter(row) for row in rows], offset=offset, limit=limit, - ), - ) - return page_result( - selected, - parameter_children_from_rows(rows), - offset=offset, - limit=limit, - ) + ) + def get_parameter( + self, + country_id: str, + parameter_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataParameterSummary]: + selected = self.select_catalog(country_id, policyengine_version) + rows = query_rows( + self._session, + select(Parameter).where( + Parameter.id == parameter_id, + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError( + f"parameter {parameter_id} was not found" + ) + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_parameter(rows[0]), + ) -def list_parameter_values( - session: Session, - selected: SelectedCatalog, - *, - parameter_id: UUID | None, - current: bool, - offset: int, - limit: int, - now: datetime | None, -) -> MetadataPageResult[MetadataCanonicalParameterValue]: - statement = ( - select(ParameterValue) - .join(Parameter, Parameter.id == ParameterValue.parameter_id) - .where( - Parameter.tax_benefit_model_version_id == selected.model_version.id, - ParameterValue.policy_id.is_(None), - ParameterValue.dynamic_id.is_(None), + def list_parameter_children( + self, + country_id: str, + policyengine_version: str | None = None, + *, + parent_path: str = "", + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataParameterChild]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, ) - ) - if parameter_id is not None: - statement = statement.where(ParameterValue.parameter_id == parameter_id) - if current: - selected_day = _utc_day_start(now or datetime.now(timezone.utc)) - statement = statement.where( - ParameterValue.start_date <= selected_day, - sa.or_( - ParameterValue.end_date.is_(None), - ParameterValue.end_date >= selected_day, + rows = query_rows( + self._session, + parameter_children_query( + model_version_id=selected.model_version.id, + parent_path=parent_path, + dialect=self._session.get_bind().dialect.name, + offset=offset, + limit=limit, ), ) - rows = query_rows( - session, - statement.order_by( - Parameter.name, - ParameterValue.start_date.desc(), - ParameterValue.id, + return page_result( + selected, + parameter_children_from_rows(rows), + offset=offset, + limit=limit, ) - .offset(offset) - .limit(limit + 1), - ) - return page_result( - selected, - [_parameter_value(row) for row in rows], - offset=offset, - limit=limit, - ) + def list_parameter_values( + self, + country_id: str, + policyengine_version: str | None = None, + *, + parameter_id: UUID | None = None, + current: bool = False, + offset: int = 0, + limit: int = 100, + now: datetime | None = None, + ) -> MetadataPageResult[MetadataCanonicalParameterValue]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + statement = ( + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ) + ) + if parameter_id is not None: + statement = statement.where(ParameterValue.parameter_id == parameter_id) + if current: + selected_day = _utc_day_start(now or datetime.now(timezone.utc)) + statement = statement.where( + ParameterValue.start_date <= selected_day, + sa.or_( + ParameterValue.end_date.is_(None), + ParameterValue.end_date >= selected_day, + ), + ) + rows = query_rows( + self._session, + statement.order_by( + Parameter.name, + ParameterValue.start_date.desc(), + ParameterValue.id, + ) + .offset(offset) + .limit(limit + 1), + ) + return page_result( + selected, + [_parameter_value(row) for row in rows], + offset=offset, + limit=limit, + ) -def get_parameter_value( - session: Session, - selected: SelectedCatalog, - value_id: UUID, -) -> MetadataDetailResult[MetadataCanonicalParameterValue]: - rows = query_rows( - session, - select(ParameterValue) - .join(Parameter, Parameter.id == ParameterValue.parameter_id) - .where( - ParameterValue.id == value_id, - Parameter.tax_benefit_model_version_id == selected.model_version.id, - ParameterValue.policy_id.is_(None), - ParameterValue.dynamic_id.is_(None), - ), - ) - if not rows: - raise MetadataResourceNotFoundError(f"parameter value {value_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_parameter_value(rows[0]), - ) + def get_parameter_value( + self, + country_id: str, + value_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataCanonicalParameterValue]: + selected = self.select_catalog(country_id, policyengine_version) + rows = query_rows( + self._session, + select(ParameterValue) + .join(Parameter, Parameter.id == ParameterValue.parameter_id) + .where( + ParameterValue.id == value_id, + Parameter.tax_benefit_model_version_id == selected.model_version.id, + ParameterValue.policy_id.is_(None), + ParameterValue.dynamic_id.is_(None), + ), + ) + if not rows: + raise MetadataResourceNotFoundError( + f"parameter value {value_id} was not found" + ) + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_parameter_value(rows[0]), + ) diff --git a/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py index 0dc4f6e07..8e24ffa60 100644 --- a/policyengine_api/data/v2/catalog/query.py +++ b/policyengine_api/data/v2/catalog/query.py @@ -1,47 +1,24 @@ -"""Route-facing service for read-only v2 metadata resource queries.""" +"""Public read-only v2 metadata query service.""" from __future__ import annotations -from datetime import datetime -from uuid import UUID - -from sqlmodel import Session - -from policyengine_api.data.v2.catalog import ( - dataset_query, - model_query, - parameter_query, - region_query, - variable_query, -) from policyengine_api.data.v2.catalog.catalog_selection import ( InvalidPolicyEngineVersionError, MetadataCatalogUnavailableError, MetadataCatalogVersionNotFoundError, - SelectedCatalog, UnsupportedPreviewCountryError, - select_catalog, validate_policyengine_version, ) +from policyengine_api.data.v2.catalog.dataset_query import DatasetQueryMethods +from policyengine_api.data.v2.catalog.model_query import ModelQueryMethods +from policyengine_api.data.v2.catalog.parameter_query import ParameterQueryMethods from policyengine_api.data.v2.catalog.query_support import ( InvalidMetadataPageError, MetadataResourceNotFoundError, validate_metadata_page, ) -from policyengine_api.data.v2.catalog.schemas import ( - MetadataCanonicalParameterValue, - MetadataDataset, - MetadataDetailResult, - MetadataEconomyOptionsResult, - MetadataModel, - MetadataModelSelectionResult, - MetadataModelVersionDetail, - MetadataPageResult, - MetadataParameterChild, - MetadataParameterSummary, - MetadataRegion, - MetadataVariable, -) +from policyengine_api.data.v2.catalog.region_query import RegionQueryMethods +from policyengine_api.data.v2.catalog.variable_query import VariableQueryMethods __all__ = [ @@ -57,326 +34,11 @@ ] -class V2MetadataQueryService: - """Select one catalog and delegate each resource to its query module.""" - - def __init__(self, session: Session, *, running_policyengine_version: str): - self._session = session - self._running_policyengine_version = validate_policyengine_version( - running_policyengine_version - ) - - def close(self) -> None: - """Close the request-owned read session.""" - - self._session.close() - - def select_catalog( - self, - country_id: str, - policyengine_version: str | None = None, - ) -> SelectedCatalog: - """Select exactly one initialized country catalog.""" - - return select_catalog( - self._session, - country_id=country_id, - running_policyengine_version=self._running_policyengine_version, - policyengine_version=policyengine_version, - ) - - def _select_paginated_catalog( - self, - country_id: str, - policyengine_version: str | None, - *, - offset: int, - limit: int, - ) -> SelectedCatalog: - validate_metadata_page(offset, limit) - return self.select_catalog(country_id, policyengine_version) - - def list_models( - self, - country_id: str, - policyengine_version: str | None = None, - *, - offset: int = 0, - limit: int = 100, - ) -> MetadataPageResult[MetadataModel]: - return model_query.list_models( - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - offset=offset, - limit=limit, - ) - - def get_model( - self, - country_id: str, - model_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataModel]: - return model_query.get_model( - self.select_catalog(country_id, policyengine_version), - model_id, - ) - - def get_model_by_country( - self, - country_id: str, - policyengine_version: str | None = None, - ) -> MetadataModelSelectionResult: - return model_query.get_model_by_country( - self.select_catalog(country_id, policyengine_version) - ) - - def list_model_versions( - self, - country_id: str, - policyengine_version: str | None = None, - *, - offset: int = 0, - limit: int = 100, - ) -> MetadataPageResult[MetadataModelVersionDetail]: - return model_query.list_model_versions( - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - offset=offset, - limit=limit, - ) - - def get_model_version( - self, - country_id: str, - version_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataModelVersionDetail]: - return model_query.get_model_version( - self.select_catalog(country_id, policyengine_version), - version_id, - ) - - def list_variables( - self, - country_id: str, - policyengine_version: str | None = None, - *, - offset: int = 0, - limit: int = 100, - search: str | None = None, - ) -> MetadataPageResult[MetadataVariable]: - return variable_query.list_variables( - self._session, - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - offset=offset, - limit=limit, - search=search, - ) - - def get_variable( - self, - country_id: str, - variable_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataVariable]: - return variable_query.get_variable( - self._session, - self.select_catalog(country_id, policyengine_version), - variable_id, - ) - - def list_parameters( - self, - country_id: str, - policyengine_version: str | None = None, - *, - offset: int = 0, - limit: int = 100, - search: str | None = None, - ) -> MetadataPageResult[MetadataParameterSummary]: - return parameter_query.list_parameters( - self._session, - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - offset=offset, - limit=limit, - search=search, - ) - - def get_parameter( - self, - country_id: str, - parameter_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataParameterSummary]: - return parameter_query.get_parameter( - self._session, - self.select_catalog(country_id, policyengine_version), - parameter_id, - ) - - def list_parameter_children( - self, - country_id: str, - policyengine_version: str | None = None, - *, - parent_path: str = "", - offset: int = 0, - limit: int = 100, - ) -> MetadataPageResult[MetadataParameterChild]: - return parameter_query.list_parameter_children( - self._session, - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - parent_path=parent_path, - offset=offset, - limit=limit, - ) - - def list_parameter_values( - self, - country_id: str, - policyengine_version: str | None = None, - *, - parameter_id: UUID | None = None, - current: bool = False, - offset: int = 0, - limit: int = 100, - now: datetime | None = None, - ) -> MetadataPageResult[MetadataCanonicalParameterValue]: - return parameter_query.list_parameter_values( - self._session, - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - parameter_id=parameter_id, - current=current, - offset=offset, - limit=limit, - now=now, - ) - - def get_parameter_value( - self, - country_id: str, - value_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataCanonicalParameterValue]: - return parameter_query.get_parameter_value( - self._session, - self.select_catalog(country_id, policyengine_version), - value_id, - ) - - def list_datasets( - self, - country_id: str, - policyengine_version: str | None = None, - *, - offset: int = 0, - limit: int = 100, - ) -> MetadataPageResult[MetadataDataset]: - return dataset_query.list_datasets( - self._session, - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - offset=offset, - limit=limit, - ) - - def get_dataset( - self, - country_id: str, - dataset_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataDataset]: - return dataset_query.get_dataset( - self._session, - self.select_catalog(country_id, policyengine_version), - dataset_id, - ) - - def list_regions( - self, - country_id: str, - policyengine_version: str | None = None, - *, - region_type: str | None = None, - offset: int = 0, - limit: int = 100, - ) -> MetadataPageResult[MetadataRegion]: - return region_query.list_regions( - self._session, - self._select_paginated_catalog( - country_id, - policyengine_version, - offset=offset, - limit=limit, - ), - region_type=region_type, - offset=offset, - limit=limit, - ) - - def get_region( - self, - country_id: str, - region_id: UUID, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataRegion]: - return region_query.get_region( - self._session, - self.select_catalog(country_id, policyengine_version), - region_id, - ) - - def get_region_by_code( - self, - country_id: str, - region_code: str, - policyengine_version: str | None = None, - ) -> MetadataDetailResult[MetadataRegion]: - return region_query.get_region_by_code( - self._session, - self.select_catalog(country_id, policyengine_version), - region_code, - ) - - def get_economy_options( - self, - country_id: str, - policyengine_version: str | None = None, - ) -> MetadataEconomyOptionsResult: - return region_query.get_economy_options( - self._session, - self.select_catalog(country_id, policyengine_version), - ) +class V2MetadataQueryService( + ModelQueryMethods, + VariableQueryMethods, + ParameterQueryMethods, + DatasetQueryMethods, + RegionQueryMethods, +): + """Combine the resource-specific query methods into the route-facing API.""" diff --git a/policyengine_api/data/v2/catalog/query_support.py b/policyengine_api/data/v2/catalog/query_support.py index 90ee933fc..8dd0b49e5 100644 --- a/policyengine_api/data/v2/catalog/query_support.py +++ b/policyengine_api/data/v2/catalog/query_support.py @@ -10,6 +10,8 @@ from policyengine_api.data.v2.catalog.catalog_selection import ( MetadataCatalogUnavailableError, SelectedCatalog, + select_catalog as select_metadata_catalog, + validate_policyengine_version, ) from policyengine_api.data.v2.catalog.schemas import MetadataPageResult @@ -25,6 +27,46 @@ class InvalidMetadataPageError(ValueError): ResourceT = TypeVar("ResourceT") +class MetadataQueryContext: + """Own the session and exact catalog selection shared by resource queries.""" + + def __init__(self, session: Session, *, running_policyengine_version: str): + self._session = session + self._running_policyengine_version = validate_policyengine_version( + running_policyengine_version + ) + + def close(self) -> None: + """Close the request-owned read session.""" + + self._session.close() + + def select_catalog( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> SelectedCatalog: + """Select exactly one initialized country catalog.""" + + return select_metadata_catalog( + self._session, + country_id=country_id, + running_policyengine_version=self._running_policyengine_version, + policyengine_version=policyengine_version, + ) + + def _select_paginated_catalog( + self, + country_id: str, + policyengine_version: str | None, + *, + offset: int, + limit: int, + ) -> SelectedCatalog: + validate_metadata_page(offset, limit) + return self.select_catalog(country_id, policyengine_version) + + def page_result( selected: SelectedCatalog, rows: list[ResourceT], diff --git a/policyengine_api/data/v2/catalog/region_query.py b/policyengine_api/data/v2/catalog/region_query.py index 482c0d492..8ac0f3417 100644 --- a/policyengine_api/data/v2/catalog/region_query.py +++ b/policyengine_api/data/v2/catalog/region_query.py @@ -4,14 +4,14 @@ from uuid import UUID -from sqlmodel import Session, select +from sqlmodel import select from policyengine_api.dataset_display import get_dataset_display_label from policyengine_api.data.v2.catalog.catalog_selection import ( MetadataCatalogUnavailableError, - SelectedCatalog, ) from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, MetadataResourceNotFoundError, page_result, query_rows, @@ -45,133 +45,145 @@ def _region(region: Region) -> MetadataRegion: ) -def list_regions( - session: Session, - selected: SelectedCatalog, - *, - region_type: str | None, - offset: int, - limit: int, -) -> MetadataPageResult[MetadataRegion]: - statement = select(Region).where( - Region.tax_benefit_model_version_id == selected.model_version.id - ) - if region_type is not None: - statement = statement.where(Region.region_type == region_type) - rows = query_rows( - session, - statement.order_by(Region.code).offset(offset).limit(limit + 1), - ) - return page_result( - selected, - [_region(row) for row in rows], - offset=offset, - limit=limit, - ) - - -def get_region( - session: Session, - selected: SelectedCatalog, - region_id: UUID, -) -> MetadataDetailResult[MetadataRegion]: - rows = query_rows( - session, - select(Region).where( - Region.id == region_id, - Region.tax_benefit_model_version_id == selected.model_version.id, - ), - ) - if not rows: - raise MetadataResourceNotFoundError(f"region {region_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_region(rows[0]), - ) +class RegionQueryMethods(MetadataQueryContext): + """Route-facing region and economy-option query methods.""" + def list_regions( + self, + country_id: str, + policyengine_version: str | None = None, + *, + region_type: str | None = None, + offset: int = 0, + limit: int = 100, + ) -> MetadataPageResult[MetadataRegion]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + statement = select(Region).where( + Region.tax_benefit_model_version_id == selected.model_version.id + ) + if region_type is not None: + statement = statement.where(Region.region_type == region_type) + rows = query_rows( + self._session, + statement.order_by(Region.code).offset(offset).limit(limit + 1), + ) + return page_result( + selected, + [_region(row) for row in rows], + offset=offset, + limit=limit, + ) -def get_region_by_code( - session: Session, - selected: SelectedCatalog, - region_code: str, -) -> MetadataDetailResult[MetadataRegion]: - rows = query_rows( - session, - select(Region).where( - Region.code == region_code, - Region.tax_benefit_model_version_id == selected.model_version.id, - ), - ) - if not rows: - raise MetadataResourceNotFoundError(f"region {region_code!r} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_region(rows[0]), - ) + def get_region( + self, + country_id: str, + region_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataRegion]: + selected = self.select_catalog(country_id, policyengine_version) + rows = query_rows( + self._session, + select(Region).where( + Region.id == region_id, + Region.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"region {region_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_region(rows[0]), + ) + def get_region_by_code( + self, + country_id: str, + region_code: str, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataRegion]: + selected = self.select_catalog(country_id, policyengine_version) + rows = query_rows( + self._session, + select(Region).where( + Region.code == region_code, + Region.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"region {region_code!r} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_region(rows[0]), + ) -def get_economy_options( - session: Session, - selected: SelectedCatalog, -) -> MetadataEconomyOptionsResult: - country_id = selected.country_id - regions = query_rows( - session, - select(Region) - .where(Region.tax_benefit_model_version_id == selected.model_version.id) - .order_by(Region.code), - ) - national_region = next( - (region for region in regions if region.code == country_id), - None, - ) - if national_region is None: - raise MetadataCatalogUnavailableError( - f"the {country_id} national v2 region is absent" + def get_economy_options( + self, + country_id: str, + policyengine_version: str | None = None, + ) -> MetadataEconomyOptionsResult: + selected = self.select_catalog(country_id, policyengine_version) + regions = query_rows( + self._session, + select(Region) + .where(Region.tax_benefit_model_version_id == selected.model_version.id) + .order_by(Region.code), ) - datasets = query_rows( - session, - select(Dataset).where( - Dataset.id == national_region.default_dataset_id, - Dataset.tax_benefit_model_version_id == selected.model_version.id, - Dataset.is_output_dataset.is_(False), - Dataset.storage_path.is_(None), - ), - ) - if len(datasets) != 1: - raise MetadataCatalogUnavailableError( - f"the {country_id} national v2 dataset is absent" + national_region = next( + (region for region in regions if region.code == country_id), + None, ) - time_periods = selected.model_version.metadata_time_periods - if ( - not isinstance(selected.model_version.current_law_id, int) - or not isinstance(time_periods, list) - or not time_periods - or any(not isinstance(year, int) for year in time_periods) - ): - raise MetadataCatalogUnavailableError( - f"the {country_id} v2 model-version options are incomplete" + if national_region is None: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 region is absent" + ) + datasets = query_rows( + self._session, + select(Dataset).where( + Dataset.id == national_region.default_dataset_id, + Dataset.tax_benefit_model_version_id == selected.model_version.id, + Dataset.is_output_dataset.is_(False), + Dataset.storage_path.is_(None), + ), ) - national_dataset = datasets[0] - return MetadataEconomyOptionsResult( - policyengine_version=selected.policyengine_version, - current_law_id=selected.model_version.current_law_id, - region=[ - MetadataRegionOption( - name=region.code, - label=region.label, - type=region.region_type.value, + if len(datasets) != 1: + raise MetadataCatalogUnavailableError( + f"the {country_id} national v2 dataset is absent" ) - for region in regions - ], - time_period=[ - MetadataTimePeriodOption(name=year, label=str(year)) - for year in time_periods - ], - datasets=[ - MetadataDatasetOption( - name=national_dataset.name, - label=get_dataset_display_label(national_dataset.name), + time_periods = selected.model_version.metadata_time_periods + if ( + not isinstance(selected.model_version.current_law_id, int) + or not isinstance(time_periods, list) + or not time_periods + or any(not isinstance(year, int) for year in time_periods) + ): + raise MetadataCatalogUnavailableError( + f"the {country_id} v2 model-version options are incomplete" ) - ], - ) + national_dataset = datasets[0] + return MetadataEconomyOptionsResult( + policyengine_version=selected.policyengine_version, + current_law_id=selected.model_version.current_law_id, + region=[ + MetadataRegionOption( + name=region.code, + label=region.label, + type=region.region_type.value, + ) + for region in regions + ], + time_period=[ + MetadataTimePeriodOption(name=year, label=str(year)) + for year in time_periods + ], + datasets=[ + MetadataDatasetOption( + name=national_dataset.name, + label=get_dataset_display_label(national_dataset.name), + ) + ], + ) diff --git a/policyengine_api/data/v2/catalog/variable_query.py b/policyengine_api/data/v2/catalog/variable_query.py index b1b9409e1..6f6c79dd9 100644 --- a/policyengine_api/data/v2/catalog/variable_query.py +++ b/policyengine_api/data/v2/catalog/variable_query.py @@ -5,10 +5,9 @@ from uuid import UUID import sqlalchemy as sa -from sqlmodel import Session, select - -from policyengine_api.data.v2.catalog.catalog_selection import SelectedCatalog +from sqlmodel import select from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, MetadataResourceNotFoundError, escape_like, page_result, @@ -37,53 +36,64 @@ def _variable(variable: Variable) -> MetadataVariable: ) -def list_variables( - session: Session, - selected: SelectedCatalog, - *, - offset: int, - limit: int, - search: str | None, -) -> MetadataPageResult[MetadataVariable]: - statement = select(Variable).where( - Variable.tax_benefit_model_version_id == selected.model_version.id - ) - if search: - pattern = f"%{escape_like(search)}%" - statement = statement.where( - sa.or_( - Variable.name.ilike(pattern, escape="\\"), - Variable.label.ilike(pattern, escape="\\"), - Variable.description.ilike(pattern, escape="\\"), +class VariableQueryMethods(MetadataQueryContext): + """Route-facing variable query methods.""" + + def list_variables( + self, + country_id: str, + policyengine_version: str | None = None, + *, + offset: int = 0, + limit: int = 100, + search: str | None = None, + ) -> MetadataPageResult[MetadataVariable]: + selected = self._select_paginated_catalog( + country_id, + policyengine_version, + offset=offset, + limit=limit, + ) + statement = select(Variable).where( + Variable.tax_benefit_model_version_id == selected.model_version.id + ) + if search: + pattern = f"%{escape_like(search)}%" + statement = statement.where( + sa.or_( + Variable.name.ilike(pattern, escape="\\"), + Variable.label.ilike(pattern, escape="\\"), + Variable.description.ilike(pattern, escape="\\"), + ) ) + rows = query_rows( + self._session, + statement.order_by(Variable.name).offset(offset).limit(limit + 1), + ) + return page_result( + selected, + [_variable(row) for row in rows], + offset=offset, + limit=limit, ) - rows = query_rows( - session, - statement.order_by(Variable.name).offset(offset).limit(limit + 1), - ) - return page_result( - selected, - [_variable(row) for row in rows], - offset=offset, - limit=limit, - ) - -def get_variable( - session: Session, - selected: SelectedCatalog, - variable_id: UUID, -) -> MetadataDetailResult[MetadataVariable]: - rows = query_rows( - session, - select(Variable).where( - Variable.id == variable_id, - Variable.tax_benefit_model_version_id == selected.model_version.id, - ), - ) - if not rows: - raise MetadataResourceNotFoundError(f"variable {variable_id} was not found") - return MetadataDetailResult( - policyengine_version=selected.policyengine_version, - item=_variable(rows[0]), - ) + def get_variable( + self, + country_id: str, + variable_id: UUID, + policyengine_version: str | None = None, + ) -> MetadataDetailResult[MetadataVariable]: + selected = self.select_catalog(country_id, policyengine_version) + rows = query_rows( + self._session, + select(Variable).where( + Variable.id == variable_id, + Variable.tax_benefit_model_version_id == selected.model_version.id, + ), + ) + if not rows: + raise MetadataResourceNotFoundError(f"variable {variable_id} was not found") + return MetadataDetailResult( + policyengine_version=selected.policyengine_version, + item=_variable(rows[0]), + ) diff --git a/tests/unit/v2/test_metadata_query.py b/tests/unit/v2/test_metadata_query.py index f6c17debc..2b57275eb 100644 --- a/tests/unit/v2/test_metadata_query.py +++ b/tests/unit/v2/test_metadata_query.py @@ -594,3 +594,30 @@ def test_query_modules_import_no_policyengine_or_v1_metadata_source() -> None: ) for module in imported ) + + +def test_resource_service_methods_are_defined_in_their_query_modules() -> None: + expected_modules = { + "list_models": "model_query", + "get_model": "model_query", + "get_model_by_country": "model_query", + "list_model_versions": "model_query", + "get_model_version": "model_query", + "list_variables": "variable_query", + "get_variable": "variable_query", + "list_parameters": "parameter_query", + "get_parameter": "parameter_query", + "list_parameter_children": "parameter_query", + "list_parameter_values": "parameter_query", + "get_parameter_value": "parameter_query", + "list_datasets": "dataset_query", + "get_dataset": "dataset_query", + "list_regions": "region_query", + "get_region": "region_query", + "get_region_by_code": "region_query", + "get_economy_options": "region_query", + } + + for method_name, module_name in expected_modules.items(): + method = getattr(V2MetadataQueryService, method_name) + assert method.__module__.endswith(f".{module_name}") From b345ffcfcb5e8b26ee2956795efbd436945250af Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 19:17:06 +0400 Subject: [PATCH 25/27] Share v2 metadata error responses by router --- policyengine_api/fastapi_routes/v2_metadata_geography.py | 8 +------- policyengine_api/fastapi_routes/v2_metadata_models.py | 9 +-------- .../fastapi_routes/v2_metadata_parameters.py | 7 +------ 3 files changed, 3 insertions(+), 21 deletions(-) diff --git a/policyengine_api/fastapi_routes/v2_metadata_geography.py b/policyengine_api/fastapi_routes/v2_metadata_geography.py index 8114a74c0..fc366b0f6 100644 --- a/policyengine_api/fastapi_routes/v2_metadata_geography.py +++ b/policyengine_api/fastapi_routes/v2_metadata_geography.py @@ -30,12 +30,11 @@ def build_v2_metadata_geography_router( dependencies: NativeRouteDependencies, ) -> APIRouter: - router = APIRouter(prefix="/v2") + router = APIRouter(prefix="/v2", responses=ERROR_RESPONSES) @router.get( "/datasets", response_model=MetadataDatasetPageResponse, - responses=ERROR_RESPONSES, summary="List logical inputs from a selected PolicyEngine.py catalog", ) def list_datasets( @@ -58,7 +57,6 @@ def list_datasets( @router.get( "/datasets/{dataset_id}", response_model=MetadataDatasetDetailResponse, - responses=ERROR_RESPONSES, summary="Get one logical input from a selected PolicyEngine.py catalog", ) def get_dataset( @@ -79,7 +77,6 @@ def get_dataset( @router.get( "/regions", response_model=MetadataRegionPageResponse, - responses=ERROR_RESPONSES, summary="List regions from a selected PolicyEngine.py catalog", ) def list_regions( @@ -104,7 +101,6 @@ def list_regions( @router.get( "/regions/by-code/{region_code:path}", response_model=MetadataRegionDetailResponse, - responses=ERROR_RESPONSES, summary="Get one region by code from a selected catalog", ) def get_region_by_code( @@ -125,7 +121,6 @@ def get_region_by_code( @router.get( "/regions/{region_id}", response_model=MetadataRegionDetailResponse, - responses=ERROR_RESPONSES, summary="Get one region from a selected PolicyEngine.py catalog", ) def get_region( @@ -146,7 +141,6 @@ def get_region( @router.get( "/economy-options", response_model=MetadataEconomyOptionsResponse, - responses=ERROR_RESPONSES, summary="Get compact economy-selection options from a selected catalog", ) def get_economy_options( diff --git a/policyengine_api/fastapi_routes/v2_metadata_models.py b/policyengine_api/fastapi_routes/v2_metadata_models.py index e9ab4d6f8..8a1147f62 100644 --- a/policyengine_api/fastapi_routes/v2_metadata_models.py +++ b/policyengine_api/fastapi_routes/v2_metadata_models.py @@ -31,12 +31,11 @@ def build_v2_metadata_model_router( dependencies: NativeRouteDependencies, ) -> APIRouter: - router = APIRouter(prefix="/v2") + router = APIRouter(prefix="/v2", responses=ERROR_RESPONSES) @router.get( "/tax-benefit-models", response_model=MetadataModelPageResponse, - responses=ERROR_RESPONSES, summary="List models for one PolicyEngine.py catalog", ) def list_models( @@ -59,7 +58,6 @@ def list_models( @router.get( "/tax-benefit-models/by-country/{country_id}", response_model=MetadataModelSelectionResponse, - responses=ERROR_RESPONSES, summary="Get a country model and selected PolicyEngine.py version", ) def get_model_by_country( @@ -78,7 +76,6 @@ def get_model_by_country( @router.get( "/tax-benefit-models/{model_id}", response_model=MetadataModelDetailResponse, - responses=ERROR_RESPONSES, summary="Get one model from a selected PolicyEngine.py catalog", ) def get_model( @@ -99,7 +96,6 @@ def get_model( @router.get( "/tax-benefit-model-versions", response_model=MetadataModelVersionPageResponse, - responses=ERROR_RESPONSES, summary="List selected PolicyEngine.py model versions", ) def list_model_versions( @@ -122,7 +118,6 @@ def list_model_versions( @router.get( "/tax-benefit-model-versions/{version_id}", response_model=MetadataModelVersionDetailResponse, - responses=ERROR_RESPONSES, summary="Get one selected PolicyEngine.py model version", ) def get_model_version( @@ -143,7 +138,6 @@ def get_model_version( @router.get( "/variables", response_model=MetadataVariablePageResponse, - responses=ERROR_RESPONSES, summary="List variables from a selected PolicyEngine.py catalog", ) def list_variables( @@ -168,7 +162,6 @@ def list_variables( @router.get( "/variables/{variable_id}", response_model=MetadataVariableDetailResponse, - responses=ERROR_RESPONSES, summary="Get one variable from a selected PolicyEngine.py catalog", ) def get_variable( diff --git a/policyengine_api/fastapi_routes/v2_metadata_parameters.py b/policyengine_api/fastapi_routes/v2_metadata_parameters.py index d94c52fb4..b4f17cb20 100644 --- a/policyengine_api/fastapi_routes/v2_metadata_parameters.py +++ b/policyengine_api/fastapi_routes/v2_metadata_parameters.py @@ -29,12 +29,11 @@ def build_v2_metadata_parameter_router( dependencies: NativeRouteDependencies, ) -> APIRouter: - router = APIRouter(prefix="/v2") + router = APIRouter(prefix="/v2", responses=ERROR_RESPONSES) @router.get( "/parameters", response_model=MetadataParameterPageResponse, - responses=ERROR_RESPONSES, summary="List parameters from a selected PolicyEngine.py catalog", ) def list_parameters( @@ -59,7 +58,6 @@ def list_parameters( @router.get( "/parameters/children", response_model=MetadataParameterChildPageResponse, - responses=ERROR_RESPONSES, summary="List direct children of one parameter path", ) def list_parameter_children( @@ -84,7 +82,6 @@ def list_parameter_children( @router.get( "/parameters/{parameter_id}", response_model=MetadataParameterDetailResponse, - responses=ERROR_RESPONSES, summary="Get one parameter from a selected PolicyEngine.py catalog", ) def get_parameter( @@ -105,7 +102,6 @@ def get_parameter( @router.get( "/parameter-values", response_model=MetadataParameterValuePageResponse, - responses=ERROR_RESPONSES, summary="List canonical values from a selected PolicyEngine.py catalog", ) def list_parameter_values( @@ -132,7 +128,6 @@ def list_parameter_values( @router.get( "/parameter-values/{value_id}", response_model=MetadataParameterValueDetailResponse, - responses=ERROR_RESPONSES, summary="Get one canonical value from a selected PolicyEngine.py catalog", ) def get_parameter_value( From fecfd75239353c538bd50dc7af1ea267fb59fd99 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Sun, 30 Aug 2026 19:32:48 +0400 Subject: [PATCH 26/27] Rename v2 metadata publication command --- .github/workflows/seed-v2-database.yml | 2 +- docs/migration/stage-9-v2-metadata.md | 2 +- ...itialize_v2_metadata.py => publish_v2_metadata_catalog.py} | 2 +- tests/unit/v2/test_metadata_deployment.py | 4 +++- 4 files changed, 6 insertions(+), 4 deletions(-) rename scripts/{initialize_v2_metadata.py => publish_v2_metadata_catalog.py} (66%) diff --git a/.github/workflows/seed-v2-database.yml b/.github/workflows/seed-v2-database.yml index aeed6d698..52c3673e2 100644 --- a/.github/workflows/seed-v2-database.yml +++ b/.github/workflows/seed-v2-database.yml @@ -39,6 +39,6 @@ jobs: env: V2_MIGRATION_DATABASE_URL: ${{ secrets.V2_MIGRATION_DATABASE_URL }} - name: Seed and validate the v2 metadata catalog - run: uv run python scripts/initialize_v2_metadata.py + run: uv run python scripts/publish_v2_metadata_catalog.py env: V2_DATA_WRITE_DATABASE_URL: ${{ secrets.V2_DATA_WRITE_DATABASE_URL }} diff --git a/docs/migration/stage-9-v2-metadata.md b/docs/migration/stage-9-v2-metadata.md index b3d48a3f0..abb7c0373 100644 --- a/docs/migration/stage-9-v2-metadata.md +++ b/docs/migration/stage-9-v2-metadata.md @@ -66,7 +66,7 @@ must execute in this order: changes. 3. Remove the migration URL from the command environment. 4. Supply only `V2_DATA_WRITE_DATABASE_URL` and the same target identity, then - run `uv run python scripts/initialize_v2_metadata.py`. + run `uv run python scripts/publish_v2_metadata_catalog.py`. 5. Require the command to finish successfully and retain its non-secret JSON evidence before creating the candidate revision. 6. Deploy the Cloud Run candidate with the target identity and diff --git a/scripts/initialize_v2_metadata.py b/scripts/publish_v2_metadata_catalog.py similarity index 66% rename from scripts/initialize_v2_metadata.py rename to scripts/publish_v2_metadata_catalog.py index 079176823..aab5212d4 100644 --- a/scripts/initialize_v2_metadata.py +++ b/scripts/publish_v2_metadata_catalog.py @@ -1,5 +1,5 @@ #!/usr/bin/env python3 -"""Run the explicit one-time API v2 metadata initialization operation.""" +"""Publish and validate the PolicyEngine.py metadata catalog in API v2.""" from policyengine_api.data.v2.catalog.initialization import main diff --git a/tests/unit/v2/test_metadata_deployment.py b/tests/unit/v2/test_metadata_deployment.py index 7b8ee4ade..80ef0a8a2 100644 --- a/tests/unit/v2/test_metadata_deployment.py +++ b/tests/unit/v2/test_metadata_deployment.py @@ -54,7 +54,9 @@ def test_reusable_seeding_workflow_separates_database_credentials() -> None: assert "V2_DATA_WRITE_DATABASE_URL" in publication assert "V2_MIGRATION_DATABASE_URL" not in publication assert "V2_RUNTIME_DATABASE_URL" not in workflow - assert "scripts/initialize_v2_metadata.py" in publication + assert "scripts/publish_v2_metadata_catalog.py" in publication + assert (REPO / "scripts/publish_v2_metadata_catalog.py").is_file() + assert not (REPO / "scripts/initialize_v2_metadata.py").exists() def test_schema_upgrade_precedes_atomic_catalog_publication() -> None: From b8a80dd9a796258211bd83c51b6552f3aa618b22 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Mon, 31 Aug 2026 20:54:46 +0400 Subject: [PATCH 27/27] Harden Stage 9 deployment readiness --- .github/workflows/push.yml | 4 +- docs/migration/stage-9-v2-metadata.md | 13 ++- policyengine_api/data/v2/migration_target.py | 29 ++--- policyengine_api/data/v2/settings.py | 68 +++++++++--- tests/integration/test_live_v2_metadata.py | 107 +++++++++++++++++++ tests/unit/test_cloud_run_deploy_scripts.py | 4 +- tests/unit/v2/test_catalog_initialization.py | 4 +- tests/unit/v2/test_database.py | 8 +- tests/unit/v2/test_settings.py | 63 +++++++++-- 9 files changed, 242 insertions(+), 58 deletions(-) create mode 100644 tests/integration/test_live_v2_metadata.py diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index c54671259..10e40311d 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -432,7 +432,7 @@ jobs: - name: Install staging test dependencies run: pip install pytest httpx - name: Run staging smoke test - run: python -m pytest tests/integration/test_cloud_run_candidate.py tests/integration/test_live_calculate.py tests/integration/test_live_economy.py tests/integration/test_live_budget_window_cache.py -v + run: python -m pytest tests/integration/test_cloud_run_candidate.py tests/integration/test_live_v2_metadata.py tests/integration/test_live_calculate.py tests/integration/test_live_economy.py tests/integration/test_live_budget_window_cache.py -v env: API_BASE_URL: ${{ needs.deploy-cloud-run-staging.outputs.url }} STAGING_API_TEST_PROBE_ID: cloud-run-${{ needs.deploy-cloud-run-staging.outputs.tag }} @@ -732,7 +732,7 @@ jobs: - name: Install Cloud Run smoke test dependencies run: pip install pytest httpx - name: Run Cloud Run candidate smoke tests - run: python -m pytest tests/integration/test_cloud_run_candidate.py -v + run: python -m pytest tests/integration/test_cloud_run_candidate.py tests/integration/test_live_v2_metadata.py -v env: API_BASE_URL: ${{ steps.candidate.outputs.url }} STAGING_API_TEST_PROBE_ID: cloud-run-${{ steps.cloud_run.outputs.revision_tag }} diff --git a/docs/migration/stage-9-v2-metadata.md b/docs/migration/stage-9-v2-metadata.md index abb7c0373..6d846621e 100644 --- a/docs/migration/stage-9-v2-metadata.md +++ b/docs/migration/stage-9-v2-metadata.md @@ -170,16 +170,19 @@ target. ## Preview verification -After Cloud Run candidate creation, explicitly request US and UK collections +Before either Cloud Run candidate is promoted, the release workflow runs +`tests/integration/test_live_v2_metadata.py` against its tagged URL. The test +forces runtime Secret Manager resolution and requests US and UK collections for variables, parameters, parameter values, datasets, and regions, followed -by representative detail routes and `GET /v2/economy-options`. Confirm that +by model resources and `GET /v2/economy-options`. It confirms that parameter collection responses do not contain parameter values and that direct -parameter-tree-child requests return only one hierarchy level. A request +resource requests remain bounded. The existing Cloud Run candidate test in the +same workflow continues to verify the unprefixed v1 metadata routes. A request without a `policyengine_version` query parameter selects the exact PolicyEngine.py version installed in that candidate artifact. It does not select the newest database row. Also request a known published version with, -for example, `?policyengine_version=5.0.4`, and confirm that each response -identifies that exact selected version. +using the version pinned in the candidate's `pyproject.toml`, and confirm that +each response identifies that exact selected version. A successful response has HTTP 200, `status: "ok"`, `message: null`, and a typed `result`. Collection results contain `policyengine_version`, `items`, diff --git a/policyengine_api/data/v2/migration_target.py b/policyengine_api/data/v2/migration_target.py index 0fd95722d..75dc1bae8 100644 --- a/policyengine_api/data/v2/migration_target.py +++ b/policyengine_api/data/v2/migration_target.py @@ -17,6 +17,7 @@ V2ConfigurationError, load_supabase_target_settings, parse_persistent_postgres_url, + validate_supabase_database_identity, ) V2_ALEMBIC_DISPOSABLE_TEST = "V2_ALEMBIC_DISPOSABLE_TEST" @@ -79,28 +80,6 @@ def _validate_disposable_url(url: URL) -> None: ) -def _validate_persistent_url_identity( - url: URL, - target: ConfiguredSupabaseTarget, -) -> None: - direct_host = f"db.{target.project_ref}.supabase.co" - is_direct = url.host == direct_host - is_pooler = bool( - url.host - and url.host.endswith(".pooler.supabase.com") - and url.username - and url.username.endswith(f".{target.project_ref}") - ) - if not (is_direct or is_pooler): - raise V2MigrationTargetError( - "the v2 migration URL does not identify the configured Supabase project" - ) - if url.database != target.database_name: - raise V2MigrationTargetError( - "the v2 migration URL does not identify the configured database" - ) - - def load_v2_alembic_settings( environ: Mapping[str, str] | None = None, ) -> V2AlembicSettings: @@ -124,6 +103,11 @@ def load_v2_alembic_settings( ) try: configured_identity = load_supabase_target_settings(values) + validate_supabase_database_identity( + persistent, + configured_identity, + setting_name=V2_MIGRATION_DATABASE_URL, + ) except V2ConfigurationError as error: raise V2MigrationTargetError(str(error)) from error target = ConfiguredSupabaseTarget( @@ -134,7 +118,6 @@ def load_v2_alembic_settings( freshness_audited_on=date(2026, 8, 13), freshness_audit_passed=True, ) - _validate_persistent_url_identity(persistent.url, target) return V2AlembicSettings( url=persistent.url, disposable_test=False, diff --git a/policyengine_api/data/v2/settings.py b/policyengine_api/data/v2/settings.py index a0f26cde3..a75b83def 100644 --- a/policyengine_api/data/v2/settings.py +++ b/policyengine_api/data/v2/settings.py @@ -31,6 +31,7 @@ PROJECT_REF_PATTERN = re.compile(r"^[a-z0-9]{20}$") ENVIRONMENT_PATTERN = re.compile(r"^[a-z][a-z0-9-]{1,31}$") SECRET_RESOURCE_PATTERN = re.compile(r"^projects/[^/]+/secrets/[^/]+/versions/[^/]+$") +SUPABASE_DATABASE_NAME = "postgres" class V2ConfigurationError(RuntimeError): @@ -190,6 +191,52 @@ def parse_persistent_postgres_url( return PostgresConnectionSettings(url) +def validate_supabase_database_identity( + connection: PostgresConnectionSettings, + target: SupabaseTargetSettings, + *, + setting_name: str, +) -> None: + """Require a persistent URL to identify the configured Supabase project.""" + + url = connection.url + direct_host = f"db.{target.project_ref}.supabase.co" + is_direct = url.host == direct_host + is_pooler = bool( + url.host + and url.host.endswith(".pooler.supabase.com") + and url.username + and url.username.endswith(f".{target.project_ref}") + ) + if not (is_direct or is_pooler): + raise V2ConfigurationError( + f"{setting_name} does not identify the configured Supabase project" + ) + if url.database != SUPABASE_DATABASE_NAME: + raise V2ConfigurationError( + f"{setting_name} does not identify the configured Supabase database" + ) + + +def _database_settings( + raw_url: str, + environ: Mapping[str, str], + *, + setting_name: str, +) -> V2DatabaseSettings: + connection = parse_persistent_postgres_url( + raw_url, + setting_name=setting_name, + ) + target = load_supabase_target_settings(environ) + validate_supabase_database_identity( + connection, + target, + setting_name=setting_name, + ) + return V2DatabaseSettings(connection=connection, target=target) + + def load_v2_runtime_database_settings( environ: Mapping[str, str] | None = None, *, @@ -202,13 +249,10 @@ def load_v2_runtime_database_settings( values, secret_loader=secret_loader or _load_secret_from_secret_manager, ) - connection = parse_persistent_postgres_url( + return _database_settings( raw_url, setting_name=V2_RUNTIME_DATABASE_URL, - ) - return V2DatabaseSettings( - connection=connection, - target=load_supabase_target_settings(values), + environ=values, ) @@ -218,13 +262,10 @@ def load_v2_migration_database_settings( """Load the schema-migration Postgres identity explicitly.""" values = _environment(environ) - connection = parse_persistent_postgres_url( + return _database_settings( _required(values, V2_MIGRATION_DATABASE_URL), setting_name=V2_MIGRATION_DATABASE_URL, - ) - return V2DatabaseSettings( - connection=connection, - target=load_supabase_target_settings(values), + environ=values, ) @@ -234,11 +275,8 @@ def load_v2_data_write_database_settings( """Load the one-time catalog row-write Postgres identity explicitly.""" values = _environment(environ) - connection = parse_persistent_postgres_url( + return _database_settings( _required(values, V2_DATA_WRITE_DATABASE_URL), setting_name=V2_DATA_WRITE_DATABASE_URL, - ) - return V2DatabaseSettings( - connection=connection, - target=load_supabase_target_settings(values), + environ=values, ) diff --git a/tests/integration/test_live_v2_metadata.py b/tests/integration/test_live_v2_metadata.py new file mode 100644 index 000000000..59ab37914 --- /dev/null +++ b/tests/integration/test_live_v2_metadata.py @@ -0,0 +1,107 @@ +"""Read-only checks for deployed Cloud Run v2 metadata preview routes.""" + +from pathlib import Path +import tomllib + +import pytest + + +REPO = Path(__file__).parents[2] +PAGED_RESOURCES = ( + "tax-benefit-models", + "tax-benefit-model-versions", + "variables", + "parameters", + "parameter-values", + "datasets", + "regions", +) + + +def _installed_policyengine_version() -> str: + dependencies = tomllib.loads((REPO / "pyproject.toml").read_text(encoding="utf-8"))[ + "project" + ]["dependencies"] + prefix = "policyengine[models]==" + matches = [ + item.removeprefix(prefix) for item in dependencies if item.startswith(prefix) + ] + assert len(matches) == 1 + return matches[0] + + +def _assert_ok(response, expected_version: str) -> dict: + assert response.status_code == 200, response.text[:500] + payload = response.json() + assert payload["status"] == "ok" + assert payload["message"] is None + assert payload["result"]["policyengine_version"] == expected_version + return payload["result"] + + +def test_live_v2_openapi_describes_every_preview_resource(api_client) -> None: + response = api_client.get("/v2/openapi.json") + + assert response.status_code == 200, response.text[:500] + document = response.json() + expected_paths = {f"/v2/{resource}" for resource in PAGED_RESOURCES} + expected_paths.update({"/v2/parameters/children", "/v2/economy-options"}) + assert expected_paths <= document["paths"].keys() + assert all("get" in document["paths"][path] for path in expected_paths) + assert "MetadataErrorResponse" in document["components"]["schemas"] + + +@pytest.mark.parametrize("country_id", ["us", "uk"]) +def test_live_v2_resources_use_the_deployed_policyengine_version( + api_client, + country_id: str, +) -> None: + expected_version = _installed_policyengine_version() + query = {"country_id": country_id, "limit": 1} + + for resource in PAGED_RESOURCES: + result = _assert_ok( + api_client.get(f"/v2/{resource}", params=query), + expected_version, + ) + assert result["limit"] == 1 + assert result["items"] + if resource == "parameters": + assert "values" not in result["items"][0] + + economy_options = _assert_ok( + api_client.get("/v2/economy-options", params={"country_id": country_id}), + expected_version, + ) + assert economy_options["region"] + assert economy_options["datasets"] + + explicit = _assert_ok( + api_client.get( + "/v2/variables", + params={ + "country_id": country_id, + "policyengine_version": expected_version, + "limit": 1, + }, + ), + expected_version, + ) + repeated = _assert_ok( + api_client.get("/v2/variables", params=query), + expected_version, + ) + assert repeated == explicit + + +def test_live_v2_invalid_version_returns_a_typed_error(api_client) -> None: + response = api_client.get( + "/v2/variables", + params={"country_id": "us", "policyengine_version": "not a version"}, + ) + + assert response.status_code == 400, response.text[:500] + payload = response.json() + assert payload["status"] == "error" + assert payload["message"] + assert "result" not in payload or payload["result"] is None diff --git a/tests/unit/test_cloud_run_deploy_scripts.py b/tests/unit/test_cloud_run_deploy_scripts.py index 97aa1608a..1da89caeb 100644 --- a/tests/unit/test_cloud_run_deploy_scripts.py +++ b/tests/unit/test_cloud_run_deploy_scripts.py @@ -1717,6 +1717,7 @@ def test_push_workflow_tests_app_engine_and_cloud_run_staging_tracks(): ) cloud_run_test_command = ( "python -m pytest tests/integration/test_cloud_run_candidate.py " + "tests/integration/test_live_v2_metadata.py " "tests/integration/test_live_calculate.py " "tests/integration/test_live_economy.py " "tests/integration/test_live_budget_window_cache.py -v" @@ -2069,7 +2070,8 @@ def test_push_workflow_promotes_production_cloud_run_after_candidate_smoke(): workflow = _push_workflow() cloud_run_production = _workflow_job_block(workflow, "deploy-cloud-run-candidate") smoke_index = cloud_run_production.index( - "python -m pytest tests/integration/test_cloud_run_candidate.py -v" + "python -m pytest tests/integration/test_cloud_run_candidate.py " + "tests/integration/test_live_v2_metadata.py -v" ) promote_index = cloud_run_production.index( "bash .github/scripts/set_cloud_run_revision.sh" diff --git a/tests/unit/v2/test_catalog_initialization.py b/tests/unit/v2/test_catalog_initialization.py index 3dbdc85c3..7bf6f6262 100644 --- a/tests/unit/v2/test_catalog_initialization.py +++ b/tests/unit/v2/test_catalog_initialization.py @@ -23,8 +23,8 @@ ENVIRONMENT = { V2_DATA_WRITE_DATABASE_URL: ( - "postgresql+psycopg://data-writer:test-password@db.example.com/" - "postgres?sslmode=require" + "postgresql+psycopg://data-writer:test-password@db." + "abcdefghijklmnopqrst.supabase.co/postgres?sslmode=require" ), V2_SUPABASE_PROJECT_REF: "abcdefghijklmnopqrst", V2_SUPABASE_ENVIRONMENT: "test-foundation", diff --git a/tests/unit/v2/test_database.py b/tests/unit/v2/test_database.py index de94a1d13..e9a153e38 100644 --- a/tests/unit/v2/test_database.py +++ b/tests/unit/v2/test_database.py @@ -22,12 +22,12 @@ def _environment(*, username: str = "runtime") -> dict[str, str]: V2_SUPABASE_PROJECT_REF: "abcdefghijklmnopqrst", V2_SUPABASE_ENVIRONMENT: "test-foundation", V2_RUNTIME_DATABASE_URL: ( - f"postgresql+psycopg://{username}:test-password@db.example.com:5432/" - "postgres?sslmode=require" + f"postgresql+psycopg://{username}:test-password@db." + "abcdefghijklmnopqrst.supabase.co:5432/postgres?sslmode=require" ), V2_MIGRATION_DATABASE_URL: ( - "postgresql+psycopg://migrator:test-password@db.example.com:5432/" - "postgres?sslmode=require" + "postgresql+psycopg://migrator:test-password@db." + "abcdefghijklmnopqrst.supabase.co:5432/postgres?sslmode=require" ), } diff --git a/tests/unit/v2/test_settings.py b/tests/unit/v2/test_settings.py index e06e45b16..020aa4f13 100644 --- a/tests/unit/v2/test_settings.py +++ b/tests/unit/v2/test_settings.py @@ -22,16 +22,16 @@ V2_SUPABASE_ENVIRONMENT: "test-foundation", } RUNTIME_URL = ( - "postgresql+psycopg://runtime:test-runtime-password@db.example.com:5432/" - "postgres?sslmode=require" + f"postgresql+psycopg://runtime:test-runtime-password@db.{PROJECT_REF}." + "supabase.co:5432/postgres?sslmode=require" ) MIGRATION_URL = ( - "postgresql+psycopg://migrator:test-migration-password@db.example.com:5432/" - "postgres?sslmode=verify-full" + f"postgresql+psycopg://migrator:test-migration-password@db.{PROJECT_REF}." + "supabase.co:5432/postgres?sslmode=verify-full" ) DATA_WRITE_URL = ( - "postgresql+psycopg://data-writer:test-data-write-password@db.example.com:5432/" - "postgres?sslmode=verify-ca" + f"postgresql+psycopg://data-writer:test-data-write-password@db.{PROJECT_REF}." + "supabase.co:5432/postgres?sslmode=verify-ca" ) RUNTIME_SECRET_RESOURCE = ( "projects/test-project/secrets/v2-runtime-database-url/versions/latest" @@ -168,6 +168,57 @@ def test_runtime_rejects_non_persistent_postgres_targets(url: str) -> None: ) +@pytest.mark.parametrize( + ("setting_name", "loader"), + [ + (V2_RUNTIME_DATABASE_URL, load_v2_runtime_database_settings), + (V2_MIGRATION_DATABASE_URL, load_v2_migration_database_settings), + (V2_DATA_WRITE_DATABASE_URL, load_v2_data_write_database_settings), + ], +) +def test_every_database_url_must_identify_the_configured_supabase_project( + setting_name: str, + loader, +) -> None: + environment = { + **TARGET_ENVIRONMENT, + setting_name: ( + "postgresql+psycopg://role:do-not-echo@db." + "aaaaaaaaaaaaaaaaaaaa.supabase.co/postgres?sslmode=require" + ), + } + + with pytest.raises(V2ConfigurationError, match="configured Supabase project"): + loader(environment) + + +def test_pooler_username_can_identify_the_configured_supabase_project() -> None: + settings = load_v2_data_write_database_settings( + { + **TARGET_ENVIRONMENT, + V2_DATA_WRITE_DATABASE_URL: ( + f"postgresql+psycopg://data-writer.{PROJECT_REF}:password@" + "aws-0-us-east-2.pooler.supabase.com:5432/" + "postgres?sslmode=require" + ), + } + ) + + assert settings.connection.url.username == f"data-writer.{PROJECT_REF}" + + +def test_supabase_database_name_must_be_postgres() -> None: + with pytest.raises(V2ConfigurationError, match="configured Supabase database"): + load_v2_runtime_database_settings( + { + **TARGET_ENVIRONMENT, + V2_RUNTIME_DATABASE_URL: RUNTIME_URL.replace( + "/postgres?", "/different?" + ), + } + ) + + def test_v1_and_debug_settings_never_supply_missing_v2_configuration() -> None: environment = { **TARGET_ENVIRONMENT,