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..4af38a66f 100644 --- a/.github/workflows/alembic-v2-check.yml +++ b/.github/workflows/alembic-v2-check.yml @@ -1,12 +1,12 @@ -name: Alembic v2 and runtime-cache checks +name: Alembic v2 schema checks on: workflow_call: 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: @@ -22,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 @@ -52,5 +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: Test real Redis cross-instance semantics - run: uv run pytest -q tests/integration/test_runtime_cache_redis.py diff --git a/.github/workflows/pr.yml b/.github/workflows/pr.yml index bd91d0b4a..ffc5d0e0b 100644 --- a/.github/workflows/pr.yml +++ b/.github/workflows/pr.yml @@ -55,8 +55,13 @@ 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: name: Check changelog fragment runs-on: ubuntu-latest diff --git a/.github/workflows/push.yml b/.github/workflows/push.yml index d320bf972..10e40311d 100644 --- a/.github/workflows/push.yml +++ b/.github/workflows/push.yml @@ -34,10 +34,17 @@ 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 }} + ensure-staging-model-version-aligns-with-sim-api: name: Ensure staging model version aligns with simulation API runs-on: ubuntu-latest @@ -58,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') @@ -101,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') @@ -153,6 +165,19 @@ jobs: if: always() run: bash .github/scripts/stop_cloud_sql_proxy.sh + 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/seed-v2-database.yml + with: + deployment_environment: staging + secrets: inherit + deploy-staging: name: Deploy staging App Engine version runs-on: ubuntu-latest @@ -160,6 +185,7 @@ jobs: - ensure-staging-model-version-aligns-with-sim-api - publish-git-tag - migrate-v1-cloud-sql + - seed-v2-staging-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -255,6 +281,7 @@ jobs: - ensure-staging-model-version-aligns-with-sim-api - publish-git-tag - migrate-v1-cloud-sql + - seed-v2-staging-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -272,6 +299,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. @@ -404,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 }} @@ -491,10 +519,21 @@ jobs: - name: Check simulation API supports PolicyEngine bundle run: bash .github/check-policyengine-bundle-supported.sh + 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/seed-v2-database.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: seed-v2-production-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -620,7 +659,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: seed-v2-production-database if: | (github.repository == 'PolicyEngine/policyengine-api') && (github.event.head_commit.message == 'Update PolicyEngine API') @@ -638,6 +677,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 @@ -692,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/.github/workflows/seed-v2-database.yml b/.github/workflows/seed-v2-database.yml new file mode 100644 index 000000000..52c3673e2 --- /dev/null +++ b/.github/workflows/seed-v2-database.yml @@ -0,0 +1,44 @@ +name: Seed v2 database + +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: + seed: + name: Upgrade schema and seed v2 database + 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: Seed and validate the v2 metadata catalog + 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/.github/workflows/v2-integration-check.yml b/.github/workflows/v2-integration-check.yml new file mode 100644 index 000000000..b69399a7b --- /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 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 + 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/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. diff --git a/docs/engineering/migration-contracts.md b/docs/engineering/migration-contracts.md index 648f79d4d..8cab684a5 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 | 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,6 +69,32 @@ 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` | +### `metadata_resources_v2_preview` + +- 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/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` - Current contract: `api_v1_compatible` diff --git a/docs/generated/migration_contracts.json b/docs/generated/migration_contracts.json index 3ae30a010..fa25f245c 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": 32, "route_group_count": 9, "sim_flow_count": 3, - "workflow_count": 7 + "workflow_count": 8 }, "route_groups": [ { @@ -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 }, @@ -217,6 +225,257 @@ } ] }, + { + "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/tax-benefit-models?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-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/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.region", + "result.time_period", + "result.datasets" + ] + } + ] + }, { "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..6d846621e --- /dev/null +++ b/docs/migration/stage-9-v2-metadata.md @@ -0,0 +1,220 @@ +# 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, 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 +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/seed-v2-database.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/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 + `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.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,813 named parameter nodes, 99,006 +parameters, 1,172,130 parameter values, 2 logical input datasets, and 826 +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. + +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 + +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 model resources and `GET /v2/economy-options`. It confirms that +parameter collection responses do not contain parameter values and that direct +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, +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`, +`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 +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, +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..115296141 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 @@ -17,6 +19,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, @@ -108,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() @@ -126,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: @@ -145,6 +164,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/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/__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/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/dataset_query.py b/policyengine_api/data/v2/catalog/dataset_query.py new file mode 100644 index 000000000..755d5b006 --- /dev/null +++ b/policyengine_api/data/v2/catalog/dataset_query.py @@ -0,0 +1,88 @@ +"""Logical input-dataset metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +from sqlmodel import select +from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, + 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, + ) + + +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( + 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/extraction.py b/policyengine_api/data/v2/catalog/extraction.py new file mode 100644 index 000000000..7d2302920 --- /dev/null +++ b/policyengine_api/data/v2/catalog/extraction.py @@ -0,0 +1,731 @@ +"""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]] = [] + 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( + 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 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" + ) + values_by_start[start_date] = (end_date, value_json) + 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) + 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" + ) + 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/model_query.py b/policyengine_api/data/v2/catalog/model_query.py new file mode 100644 index 000000000..17a0e130b --- /dev/null +++ b/policyengine_api/data/v2/catalog/model_query.py @@ -0,0 +1,118 @@ +"""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 ( + MetadataQueryContext, + 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, + ) + + +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 new file mode 100644 index 000000000..c3807f1f2 --- /dev/null +++ b/policyengine_api/data/v2/catalog/parameter_query.py @@ -0,0 +1,238 @@ +"""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 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, + 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, + ) + + +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), + ) + return page_result( + selected, + [_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 = 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_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, + ) + 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, + ), + ) + return page_result( + selected, + parameter_children_from_rows(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( + 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/parameter_tree_query.py b/policyengine_api/data/v2/catalog/parameter_tree_query.py new file mode 100644 index 000000000..cc4c856a7 --- /dev/null +++ b/policyengine_api/data/v2/catalog/parameter_tree_query.py @@ -0,0 +1,167 @@ +"""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 _path_segment(remainder: object, dialect: str) -> object: + dot_position = ( + sa.func.instr(remainder, ".") + if dialect == "sqlite" + else sa.func.strpos(remainder, ".") + ) + return sa.case( + (dot_position > 0, sa.func.substr(remainder, 1, dot_position - 1)), + else_=remainder, + ) + + +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( + *, + 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() + direct_child_paths = sa.union( + select( + _direct_child_path( + ParameterNode.name, + paths.c.path, + dialect, + ).label("path") + ) + .where( + 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, + 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_(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, direct_child_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/publication.py b/policyengine_api/data/v2/catalog/publication.py new file mode 100644 index 000000000..3e59306cf --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication.py @@ -0,0 +1,149 @@ +"""Atomic PostgreSQL publication for a validated PolicyEngine.py catalog.""" + +from __future__ import annotations + +from collections.abc import Callable +import logging +import time + +import sqlalchemy as sa +from sqlalchemy import Connection, Engine + +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 +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__ = [ + "CatalogPublicationError", + "PublicationEvidence", + "publish_catalog", +] + + +def _verify_expected_revision(connection: Connection) -> None: + if connection.dialect.name != "postgresql": + raise CatalogPublicationError("catalog publication requires PostgreSQL") + if not sa.inspect(connection).has_table(ALEMBIC_VERSION.name): + raise CatalogPublicationError("the v2 Alembic revision table is absent") + revisions = set( + connection.execute(sa.select(ALEMBIC_VERSION.c.version_num)).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( + 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(sa.select(sa.func.count()).select_from(table)).scalar_one() + for table in PROTECTED_TABLES + ) + + +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): + 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) + if checkpoint is not None: + checkpoint("after_reconciliation", connection) + + for country in catalog.countries: + assert_country_matches(connection, country) + 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/publication_reconciliation.py b/policyengine_api/data/v2/catalog/publication_reconciliation.py new file mode 100644 index 000000000..8ebac6d6e --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication_reconciliation.py @@ -0,0 +1,646 @@ +"""Set-based reconciliation and validation for a staged v2 catalog.""" + +from __future__ import annotations + +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, +) + + +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 + + +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, + ) + + +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) + ) + + 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) + ) + + 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 assert_country_matches( + connection: Connection, + country: CountryCatalog, +) -> None: + 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" + ) + + +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]) +) + +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), +) + +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), +) + +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), +) + +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), +) + +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), +) + +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), +) + +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), +) + +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: 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 = ( + 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, + ) + .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 new file mode 100644 index 000000000..325d3ac3c --- /dev/null +++ b/policyengine_api/data/v2/catalog/publication_staging.py @@ -0,0 +1,323 @@ +"""Temporary-table staging for v2 catalog publication.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable, Iterator, Sequence + +from psycopg import sql +from psycopg.types.json import Jsonb +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, +) +from policyengine_api.data.v2.catalog.records import ( + CountryCatalog, + NormalizedCatalog, + iter_batches, +) + + +COPY_BATCH_SIZE = 10_000 + +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", +) + +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", + 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: + return None if value is None else Jsonb(value) + + +def _catalog_rows( + country: CountryCatalog, +) -> dict[Table, 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: Table, + rows: Iterable[Sequence[object]], +) -> int: + """Write one bounded source stream through Psycopg COPY.""" + + raw_connection = connection.connection.driver_connection + 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: + for row in rows: + copy.write_row(row) + count += 1 + return count + + +def create_staging_tables(connection: Connection) -> None: + STAGING_METADATA.create_all(connection, checkfirst=False) + + +def stage_catalog( + connection: Connection, + catalog: NormalizedCatalog, + *, + checkpoint: Callable[[str, Connection], None] | None = None, +) -> dict[str, int]: + observed = {table.name: 0 for table in STAGING_TABLES} + for country in catalog.countries: + for table, rows in _catalog_rows(country).items(): + observed[table.name] += copy_rows( + connection, + table=table, + rows=rows, + ) + if checkpoint is not None: + checkpoint("during_copy", connection) + expected = catalog.entity_counts() + expected_by_table = { + 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") + 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/policyengine_api/data/v2/catalog/query.py b/policyengine_api/data/v2/catalog/query.py new file mode 100644 index 000000000..8e24ffa60 --- /dev/null +++ b/policyengine_api/data/v2/catalog/query.py @@ -0,0 +1,44 @@ +"""Public read-only v2 metadata query service.""" + +from __future__ import annotations + +from policyengine_api.data.v2.catalog.catalog_selection import ( + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, + UnsupportedPreviewCountryError, + 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.region_query import RegionQueryMethods +from policyengine_api.data.v2.catalog.variable_query import VariableQueryMethods + + +__all__ = [ + "InvalidMetadataPageError", + "InvalidPolicyEngineVersionError", + "MetadataCatalogUnavailableError", + "MetadataCatalogVersionNotFoundError", + "MetadataResourceNotFoundError", + "UnsupportedPreviewCountryError", + "V2MetadataQueryService", + "validate_metadata_page", + "validate_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 new file mode 100644 index 000000000..8dd0b49e5 --- /dev/null +++ b/policyengine_api/data/v2/catalog/query_support.py @@ -0,0 +1,112 @@ +"""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, + select_catalog as select_metadata_catalog, + validate_policyengine_version, +) +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") + + +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], + *, + 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/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/region_query.py b/policyengine_api/data/v2/catalog/region_query.py new file mode 100644 index 000000000..8ac0f3417 --- /dev/null +++ b/policyengine_api/data/v2/catalog/region_query.py @@ -0,0 +1,189 @@ +"""Region and economy-option metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +from sqlmodel import select + +from policyengine_api.dataset_display import get_dataset_display_label +from policyengine_api.data.v2.catalog.catalog_selection import ( + MetadataCatalogUnavailableError, +) +from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, + 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, + ) + + +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( + 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( + 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), + ) + 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( + 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), + ), + ) + 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/schemas.py b/policyengine_api/data/v2/catalog/schemas.py new file mode 100644 index 000000000..ddd01fb8e --- /dev/null +++ b/policyengine_api/data/v2/catalog/schemas.py @@ -0,0 +1,267 @@ +"""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, Generic, Literal, TypeVar +from uuid import UUID + +from pydantic import BaseModel, ConfigDict, 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 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 + 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 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 MetadataErrorResponse(StrictResponseModel): + status: Literal["error"] = "error" + message: Annotated[str, StringConstraints(strip_whitespace=True, min_length=1)] 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..6f6c79dd9 --- /dev/null +++ b/policyengine_api/data/v2/catalog/variable_query.py @@ -0,0 +1,99 @@ +"""Variable metadata queries.""" + +from __future__ import annotations + +from uuid import UUID + +import sqlalchemy as sa +from sqlmodel import select +from policyengine_api.data.v2.catalog.query_support import ( + MetadataQueryContext, + 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, + ) + + +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, + ) + + 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/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/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..a75b83def 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,8 @@ 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/[^/]+$") +SUPABASE_DATABASE_NAME = "postgres" class V2ConfigurationError(RuntimeError): @@ -81,6 +87,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: @@ -138,19 +191,68 @@ 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, + *, + secret_loader: Callable[[str], str] | None = None, ) -> V2DatabaseSettings: """Load the future ordinary-runtime Postgres identity explicitly.""" values = _environment(environ) - connection = parse_persistent_postgres_url( - _required(values, V2_RUNTIME_DATABASE_URL), - setting_name=V2_RUNTIME_DATABASE_URL, + raw_url = _resolve_runtime_database_url( + values, + secret_loader=secret_loader or _load_secret_from_secret_manager, ) - return V2DatabaseSettings( - connection=connection, - target=load_supabase_target_settings(values), + return _database_settings( + raw_url, + setting_name=V2_RUNTIME_DATABASE_URL, + environ=values, ) @@ -160,11 +262,21 @@ 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, + environ=values, ) - return V2DatabaseSettings( - 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) + return _database_settings( + _required(values, V2_DATA_WRITE_DATABASE_URL), + setting_name=V2_DATA_WRITE_DATABASE_URL, + environ=values, ) 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/policyengine_api/fastapi_routes/dependencies.py b/policyengine_api/fastapi_routes/dependencies.py index b1a6eccf4..4b3166689 100644 --- a/policyengine_api/fastapi_routes/dependencies.py +++ b/policyengine_api/fastapi_routes/dependencies.py @@ -4,6 +4,8 @@ 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.json_types import JSONObject @@ -15,6 +17,12 @@ class MetadataReader(Protocol): def get_metadata(self, country_id: str) -> JSONObject: ... +class V2MetadataResourceReader(Protocol): + """Own the database session used by one v2 metadata resource request.""" + + def close(self) -> None: ... + + class SimulationGatewayProbe(Protocol): """Minimal simulation-entrypoint health-check interface.""" @@ -45,6 +53,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() -> V2MetadataResourceReader: + 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 +76,7 @@ class NativeRouteDependencies: gateway_client_factory: Callable[[], SimulationGatewayProbe] metadata_reader_factory: Callable[[], MetadataReader] specification_provider: Callable[[], JSONObject] + v2_metadata_reader_factory: Callable[[], V2MetadataResourceReader] | None = None @classmethod def defaults(cls) -> "NativeRouteDependencies": @@ -62,4 +86,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..000eddcdd --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata.py @@ -0,0 +1,90 @@ +"""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.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, +) +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, +) + + +def build_v2_metadata_router( + dependencies: NativeRouteDependencies, +) -> APIRouter: + """Build isolated resource 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", + include_in_schema=False, + summary="OpenAPI document for dormant v2 metadata resources", + ) + 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) + + @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, + status_code=404, + include_in_schema=False, + ) + def unsupported_resource(resource_path: str) -> MetadataErrorResponse: + return MetadataErrorResponse( + message=f"V2 metadata resource {resource_path!r} was not found" + ) + + @router.api_route( + "/v2/{resource_path:path}", + methods=["POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"], + response_model=MetadataErrorResponse, + status_code=405, + include_in_schema=False, + ) + def unsupported_method(resource_path: str) -> MetadataErrorResponse: + return MetadataErrorResponse( + message=f"V2 metadata resource {resource_path!r} supports GET only" + ) + + return router 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..fc366b0f6 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_geography.py @@ -0,0 +1,159 @@ +"""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", responses=ERROR_RESPONSES) + + @router.get( + "/datasets", + response_model=MetadataDatasetPageResponse, + 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, + 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, + 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, + 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, + 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, + 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..8a1147f62 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_models.py @@ -0,0 +1,182 @@ +"""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", responses=ERROR_RESPONSES) + + @router.get( + "/tax-benefit-models", + response_model=MetadataModelPageResponse, + 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, + 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, + 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, + 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, + 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, + 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, + 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..b4f17cb20 --- /dev/null +++ b/policyengine_api/fastapi_routes/v2_metadata_parameters.py @@ -0,0 +1,148 @@ +"""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", responses=ERROR_RESPONSES) + + @router.get( + "/parameters", + response_model=MetadataParameterPageResponse, + 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, + 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, + 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, + 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, + 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/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..02a54e26d 100644 --- a/policyengine_api/migration_logging.py +++ b/policyengine_api/migration_logging.py @@ -19,6 +19,31 @@ ) +V2_METADATA_RESOURCE_SEGMENTS = frozenset( + { + "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.""" @@ -72,9 +97,12 @@ def log_migration_request( elapsed_ms = round((time.time() - started_at) * 1000, 2) route_group = infer_route_group(path) + is_v2_metadata_read = _is_v2_metadata_resource_read(method, path) 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/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/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..699991520 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_resources"}) 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/publish_v2_metadata_catalog.py b/scripts/publish_v2_metadata_catalog.py new file mode 100644 index 000000000..aab5212d4 --- /dev/null +++ b/scripts/publish_v2_metadata_catalog.py @@ -0,0 +1,8 @@ +#!/usr/bin/env python3 +"""Publish and validate the PolicyEngine.py metadata catalog in API v2.""" + +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..c24e9eb41 100644 --- a/tests/contract/registry.py +++ b/tests/contract/registry.py @@ -122,6 +122,92 @@ class WorkflowContract: ), ), ), + WorkflowContract( + 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/tax-benefit-models/by-country/us", + expected_status=200, + stable_response_fields=( + "status", + "message", + "result.policyengine_version", + "result.model", + "result.model_version", + ), + route_group="metadata", + ), + ContractRequest( + method="GET", + path="/v2/economy-options?country_id=us", + expected_status=200, + stable_response_fields=( + "status", + "message", + "result.policyengine_version", + "result.current_law_id", + "result.region", + "result.time_period", + "result.datasets", + ), + route_group="metadata", + ), + ), + ), WorkflowContract( name="simulation_submit_poll", current_contract="api_v1_compatible", @@ -201,3 +287,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..8a1991d84 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", + "metadata_resources_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_resources" + if workflow.name == "metadata_resources_v2_preview" + else "api_v1_compatible" + ) + assert workflow.current_contract == expected_contract assert workflow.future_owner_pr assert workflow.requests @@ -24,3 +34,16 @@ 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 + } == { + 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/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_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/integration/test_v2_catalog_installed.py b/tests/integration/test_v2_catalog_installed.py new file mode 100644 index 000000000..b09ce746a --- /dev/null +++ b/tests/integration/test_v2_catalog_installed.py @@ -0,0 +1,83 @@ +"""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_813, + "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 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 + 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..a7cc3d79d --- /dev/null +++ b/tests/integration/test_v2_metadata_routes.py @@ -0,0 +1,360 @@ +"""PostgreSQL-backed integration coverage for v2 metadata resource 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 metadata route tests require disposable local PostgreSQL") + 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 _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) + 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) + 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) + 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"] + assert repeated_variables.json() == variables.json() + 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 + + +def test_postgres_parameter_tree_returns_direct_children_and_leaf_details( + published_engine: Engine, +) -> None: + client = _client(published_engine) + query = {"country_id": "us"} + + root = client.get("/v2/parameters/children", params=query) + government = client.get( + "/v2/parameters/children", + params={**query, "parent_path": "gov"}, + ) + example = client.get( + "/v2/parameters/children", + params={**query, "parent_path": "gov.example"}, + ) + + 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_dataset_collection_excludes_an_existing_output( + 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/datasets", + params={"country_id": "us"}, + ) + + assert response.status_code == 200 + assert "existing-output" not in { + dataset["name"] for dataset in response.json()["result"]["items"] + } + + +@pytest.mark.parametrize("country_id", ["us", "uk"]) +def test_postgres_resources_default_to_running_version_and_accept_exact_override( + 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( + "/v2/variables", + params={"country_id": country_id}, + ) + selected_response = client.get( + "/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["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_resources_distinguish_invalid_and_absent_versions( + published_engine: Engine, +) -> None: + client = _client(published_engine) + + invalid = client.get( + "/v2/variables", + params={"country_id": "us", "policyengine_version": "not a version"}, + ) + absent = client.get( + "/v2/variables", + params={"country_id": "us", "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"] + + +def test_postgres_parameter_collection_remains_available_without_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' + ) + """ + ) + ) + + client = _client(published_engine) + parameters = client.get("/v2/parameters", params={"country_id": "us"}) + values = client.get("/v2/parameter-values", params={"country_id": "us"}) + + assert parameters.status_code == values.status_code == 200 + assert parameters.json()["result"]["items"] + assert values.json()["result"]["items"] == [] diff --git a/tests/unit/routes/test_migration_context_logging.py b/tests/unit/routes/test_migration_context_logging.py index 38e107fb0..7b5499771 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_resource_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/variables", + 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..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,20 +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 "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 "needs: [lint, alembic-v1-check, alembic-v2-check]" 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 + 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 @@ -116,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" @@ -125,14 +147,46 @@ 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 + 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 + 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/test_cloud_run_deploy_scripts.py b/tests/unit/test_cloud_run_deploy_scripts.py index 6257f0f3f..1da89caeb 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(): @@ -1675,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" @@ -1743,10 +1786,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: seed-v2-production-database" + 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 +1802,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 @@ -2029,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/test_migration_contract_artifacts.py b/tests/unit/test_migration_contract_artifacts.py index 014e3f8c1..9517db7bc 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": 32, "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", + "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 8368b0462..e1afa756e 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"), [ @@ -115,7 +133,11 @@ def test_invalid_migration_flag_raises(monkeypatch): ("/health", "health"), ("/simulation-gateway-check", "health"), ("/readiness-check", "health"), + ("/v2/openapi.json", "specification"), ("/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_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..4e07ede36 --- /dev/null +++ b/tests/unit/v2/test_catalog_extraction.py @@ -0,0 +1,415 @@ +"""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, + ) + + +@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( + 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_collapses_equal_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=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="inconsistent intervals"): + 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..7bf6f6262 --- /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." + "abcdefghijklmnopqrst.supabase.co/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..6d09cafb5 --- /dev/null +++ b/tests/unit/v2/test_catalog_publication.py @@ -0,0 +1,291 @@ +"""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 +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] + + +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, revisions=()): + self.dialect = type("Dialect", (), {"name": dialect})() + self.result = _ScalarResult(values=revisions) + + def execute(self, _statement): + 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 ( + 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 ( + patch.object(publication.sa, "inspect", return_value=_Inspector(True)), + pytest.raises(publication.CatalogPublicationError, match="expected"), + ): + publication._verify_expected_revision( + _RevisionConnection( + dialect="postgresql", + 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)) + 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=copy_table, + rows=rows, + ) + + cursor = connection.connection.driver_connection.selected_cursor + assert count == 3 + assert cursor.statement.as_string() == ( + '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_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__, + ) + staging_ddl = tuple( + str(CreateTable(table).compile(dialect=dialect)) + for table in publication_staging.STAGING_TABLES + ) + assert all( + "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( + isinstance(statement, sa.sql.Select) and not isinstance(statement, TextClause) + for pair in comparisons + for statement in pair + ) + + class Result: + def scalar_one(self): + return None + + class Connection: + statement = None + + def execute(self, statement): + self.statement = statement + return Result() + + connection = Connection() + publication._acquire_publication_lock(connection) + 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: + 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_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_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..80ef0a8a2 --- /dev/null +++ b/tests/unit/v2/test_metadata_deployment.py @@ -0,0 +1,99 @@ +"""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_seeding_workflow_separates_database_credentials() -> None: + workflow = _read(".github/workflows/seed-v2-database.yml") + migration = _step( + workflow, + "Upgrade and verify the v2 schema", + "Seed and validate the v2 metadata catalog", + ) + publication = _step( + workflow, + "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 + 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/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: + 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_seeding_success_is_required_before_candidate_creation() -> None: + workflow = _read(".github/workflows/push.yml") + staging_seed = _job(workflow, "seed-v2-staging-database") + production_seed = _job(workflow, "seed-v2-production-database") + + 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 "seed-v2-staging-database" in _job(workflow, job_name) + + assert ( + "needs: ensure-production-model-version-aligns-with-sim-api" in production_seed + ) + assert "deployment_environment: production" in production_seed + for job_name in ("deploy-production-candidate", "deploy-cloud-run-candidate"): + assert "needs: seed-v2-production-database" 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..2b57275eb --- /dev/null +++ b/tests/unit/v2/test_metadata_query.py @@ -0,0 +1,623 @@ +"""Read-only query coverage for v2 metadata resources.""" + +from __future__ import annotations + +import ast +from datetime import datetime, timezone +from pathlib import Path +from uuid import uuid4 + +import pytest +from sqlalchemy import 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, +) +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 + + +def _insert_country( + session: Session, + country, + *, + include_model: bool, + current_law_id: int | None = None, + time_periods: list[int] | None = None, +) -> None: + 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=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( + 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 + ) + + +@pytest.fixture +def catalog_session() -> Session: + engine = create_engine( + "sqlite://", + connect_args={"check_same_thread": False}, + poolclass=StaticPool, + ) + V2_METADATA.create_all(engine) + with Session(engine) as session: + for country in normalized_catalog().countries: + _insert_country(session, country, include_model=True) + 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" + ) + _insert_country( + session, + country, + include_model=False, + current_law_id=current_law_id, + time_periods=time_periods, + ) + session.commit() + + +def _us_model_version(session: Session) -> TaxBenefitModelVersion: + return session.exec( + select(TaxBenefitModelVersion) + .join(TaxBenefitModel, TaxBenefitModel.id == TaxBenefitModelVersion.model_id) + .where( + TaxBenefitModel.name == "policyengine-us", + TaxBenefitModelVersion.version == POLICYENGINE_VERSION, + ) + ).one() + + +def _service(session: Session) -> V2MetadataQueryService: + return V2MetadataQueryService( + session, + running_policyengine_version=POLICYENGINE_VERSION, + ) + + +def test_resource_collection_uses_bounded_pagination_without_counting( + catalog_session: Session, +) -> None: + model_version = _us_model_version(catalog_session) + 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 = _service(catalog_session).list_variables("us", 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): + _service(catalog_session).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 = _service(catalog_session).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: + model_version = _us_model_version(catalog_session) + parameter = catalog_session.exec( + select(Parameter).where( + Parameter.tax_benefit_model_version_id == model_version.id, + Parameter.name == "gov.example.rate", + ) + ).one() + 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=uuid4(), + ) + ) + catalog_session.commit() + + all_values = _service(catalog_session).list_parameter_values( + "us", + parameter_id=parameter.id, + ) + current_value = _service(catalog_session).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) + + +@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.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 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( + catalog_session: Session, +) -> None: + 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") + 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_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: + _add_country_version( + catalog_session, + policyengine_version="5.0.5", + current_law_id=22, + time_periods=[2041, 2040], + ) + service = _service(catalog_session) + + 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_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] = [] + + 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 = _service(catalog_session).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"] + + +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 = ( + "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")) + 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 + ) + + +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}") diff --git a/tests/unit/v2/test_metadata_routes.py b/tests/unit/v2/test_metadata_routes.py new file mode 100644 index 000000000..04c960969 --- /dev/null +++ b/tests/unit/v2/test_metadata_routes.py @@ -0,0 +1,559 @@ +"""Typed route coverage for dormant v2 metadata resources.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from unittest.mock import patch +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 ( + InvalidMetadataPageError, + InvalidPolicyEngineVersionError, + MetadataCatalogUnavailableError, + MetadataCatalogVersionNotFoundError, + MetadataResourceNotFoundError, + UnsupportedPreviewCountryError, +) +from policyengine_api.data.v2.catalog.schemas import ( + MetadataCanonicalParameterValue, + MetadataDataset, + MetadataDatasetOption, + MetadataDetailResult, + MetadataEconomyOptionsResult, + MetadataModel, + MetadataModelSelectionResult, + MetadataModelVersionDetail, + MetadataPageResult, + MetadataParameterChild, + MetadataParameterSummary, + MetadataRegion, + MetadataRegionOption, + MetadataTimePeriodOption, + 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, + RouteImplementationSettings, +) + + +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 + + 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 + if self.close_error is not None: + raise self.close_error + + +def _resource_results() -> tuple[dict[str, object], dict[str, object]]: + version = "4.20.3" + model_id = uuid4() + model_version_id = uuid4() + variable_id = uuid4() + parameter_id = uuid4() + parameter_value_id = uuid4() + dataset_id = uuid4() + region_id = uuid4() + model = MetadataModel( + id=model_id, + name="policyengine-us", + description="US model", + ) + model_version = MetadataModelVersionDetail( + id=model_version_id, + model_id=model_id, + version=version, + description="PolicyEngine.py catalog", + current_law_id=2, + metadata_time_periods=[2026], + ) + 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( + 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: + 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([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: + 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, + ) + + +def test_each_resource_route_returns_its_typed_result() -> None: + results, ids = _resource_results() + reader = ResourceReader(results) + client = _client(lambda: reader) + 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/{ids['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/{ids['model_version_id']}?country_id=us", + ), + ("list_variables", "/v2/variables?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/{ids['parameter_id']}?country_id=us", + ), + ("list_parameter_values", "/v2/parameter-values?country_id=us"), + ( + "get_parameter_value", + f"/v2/parameter-values/{ids['parameter_value_id']}?country_id=us", + ), + ("list_datasets", "/v2/datasets?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", "/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"), + ] + + 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_resource_route_forwards_version_filters_and_pagination() -> None: + results, _ids = _resource_results() + reader = ResourceReader(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 + + +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", + [ + {}, + {"country_id": "us", "offset": -1}, + {"country_id": "us", "limit": 0}, + {"country_id": "us", "limit": 501}, + ], +) +def test_request_validation_failures_use_the_error_schema(params: dict) -> None: + calls = [] + + def factory(): + calls.append("called") + results, _ids = _resource_results() + return ResourceReader(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), + (V2ConfigurationError("missing database URL"), 503), + (RuntimeError("private query detail"), 500), + ], +) +def test_query_failures_use_documented_error_statuses( + error: Exception, + expected_status: int, +) -> None: + 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 + assert response.json()["status"] == "error" + 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_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"): + 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 == [] + + +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_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") + + 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" 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..020aa4f13 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, ) @@ -19,28 +22,38 @@ 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 = ( + 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" ) -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,90 @@ 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) + + +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", [ @@ -71,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, @@ -83,6 +231,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" },