diff --git a/.github/workflows/pr-image-smoke.yml b/.github/workflows/pr-image-smoke.yml index 23d3a875f..0728b8c00 100644 --- a/.github/workflows/pr-image-smoke.yml +++ b/.github/workflows/pr-image-smoke.yml @@ -30,9 +30,10 @@ jobs: name: Image smoke (staging) runs-on: ubuntu-latest if: github.event.pull_request.head.repo.full_name == github.repository - # Warm cache: ~2 min. After a relock the executor smoke pays the - # certified dataset download layer (~15-20 min). - timeout-minutes: 45 + # Image validation plus real US regional preparation/calculation. Prepared + # years are never reused; the 30-minute Modal calculation timeout and image + # construction must both fit inside this job's hard timeout. + timeout-minutes: 60 steps: - uses: actions/checkout@v6 diff --git a/docs/us-regional-dataset-pr-validation.md b/docs/us-regional-dataset-pr-validation.md new file mode 100644 index 000000000..08859a8c0 --- /dev/null +++ b/docs/us-regional-dataset-pr-validation.md @@ -0,0 +1,56 @@ +# US regional dataset validation before merge + +The `PR image smoke` workflow checks more than package installation and the +national dataset. Its executor command now also calls an ephemeral Modal +function that: + +1. Resolves `state/ut` through the real executor region and dataset resolvers. +2. Requires a non-default certified source with its revision and SHA-256 pin. +3. Calls the same `ensure_datasets` API as production workers for 2026, using + a fresh temporary directory. No prepared year or baseline output is reused. +4. Builds and runs a Utah baseline through the executor's simulation builder, + requesting WIC participation, WIC benefits, state FIPS, and net income. +5. Requires complete boolean participation in prepared and calculated inputs, + nonempty Utah outputs, and finite net income and WIC benefits. +6. Prints the installed `.py` version, selected source revision and hash, and + row counts. Calculation errors propagate and fail the PR check. + +The country year-preparation API currently prepares the shared ACS file before +the simulation filters Utah. This is not just a small household-fixture test: +the approximately 9.8-GB certified source must be downloaded and loaded. The +function uses 8 CPUs, 64 GiB of memory, one container, and a hard 30-minute +execution timeout. The workflow allows 60 minutes including image construction +and its other checks. Timeouts and missing credentials are failures, not reasons +to skip this validation. + +The function runs in an ephemeral Modal app using the same runtime image as +the existing import check. It creates no scheduled process, public endpoint, +artifact bucket, persistent volume, deployment, or manifest. Prepared files +stay in its temporary directory; it does not write to production GCS or SQL. +The existing US national artifact precompute remains unchanged. + +The source is the installed `.py` bundle's public US dataset. Existing GitHub +Modal credentials launch the function; existing UK credentials continue to be +used to construct and validate the shared image. No new environment variable, +secret, permission, or provisioning step is introduced. + +The acceptance-condition unit tests use small DataFrames. They do not stand +in for the real loading check. The real check is invoked by +`.github/scripts/modal-image-smoke.sh` through `src/modal/smoke_app.py`, not by +the local Docker integration command that excludes `beta_only` calculations. + +With the affected `.py` 6.2.2 release, real preparation is expected to fail on +missing ACS WIC participation. Do not catch, skip, or mark that failure as +expected in CI. The executor's two `.py` pins now target 6.2.5, including the +narrowly scoped temporary compatibility fix in PolicyEngine/policyengine.py#566. + +## Pending publication and lockfile refresh + +Publish `.py` 6.2.5 before regenerating the executor's `uv.lock` with `uv lock`. +Commit that generated lockfile, run `uv lock --check`, and rerun the real ACS +validation before merging or deploying this change. Do not invent package +URLs, hashes, or release metadata before publication. + +Until that refresh, the checked-in lockfile retains main's published 6.2.3 +release. Frozen installs therefore use 6.2.3, not the new target; this state +is intentionally incomplete and must not be merged or deployed. diff --git a/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_bundle.py b/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_bundle.py index d20fa1ccb..480daa50c 100644 --- a/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_bundle.py +++ b/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_bundle.py @@ -36,12 +36,13 @@ class Stage12CountryBundle(StrictBundleModel): default_dataset: NonEmptyText default_dataset_uri: NonEmptyText datasets: tuple[Stage12Dataset, ...] + region_dataset_identities: dict[NonEmptyText, NonEmptyText] class Stage12BundleManifest(StrictBundleModel): """Stage 12's normalized subset of a PolicyEngine.py bundle.""" - schema_version: Literal[1] = 1 + schema_version: Literal[2] = 2 source_bundle_schema_version: Literal[2] = 2 policyengine_version: NonEmptyText policyengine_requirement: NonEmptyText @@ -54,3 +55,43 @@ class Stage12BundleManifest(StrictBundleModel): class ResolvedStage12Bundle(StrictBundleModel): bundle: Stage12BundleManifest bundle_manifest_sha256: Sha256Digest + + +def select_stage12_dataset( + country: Stage12CountryBundle, + region: str, +) -> Stage12Dataset: + """Select the certified dataset declared for a concrete region type.""" + + normalized_region = region.strip().lower() + region_type = "national" + if normalized_region != country.country: + if ( + country.country == "us" + and len(normalized_region) == 2 + and normalized_region.isalpha() + ): + region_type = "state" + else: + requested_region_type, separator, _ = normalized_region.partition("/") + if separator: + region_type = requested_region_type + dataset_identity = country.region_dataset_identities.get(region_type) + if dataset_identity is None: + raise ValueError( + f"No certified dataset is declared for region type {region_type!r}." + ) + selected = next( + ( + dataset + for dataset in country.datasets + if dataset.identity == dataset_identity + ), + None, + ) + if selected is None: + raise ValueError( + f"Certified region type {region_type!r} names absent dataset " + f"{dataset_identity!r}." + ) + return selected diff --git a/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py b/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py index 627a43f5b..82be05af8 100644 --- a/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py +++ b/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py @@ -11,6 +11,7 @@ Stage12BundleManifest, Stage12CountryBundle, Stage12Dataset, + select_stage12_dataset, ) from policyengine_simulation_contract.stage12_manifest import ( ACTIVE_MANIFEST_KEY, @@ -41,13 +42,37 @@ def __setitem__(self, key: str, value: object) -> None: self.values[key] = value +def _dataset_fixture(identity: str, sha256: str) -> Stage12Dataset: + """Build synthetic provenance without reading deployed dataset metadata.""" + revision = f"{identity}-revision" + return Stage12Dataset( + identity=identity, + uri=f"hf://policyengine/data/{identity}.h5@{revision}", + artifact_revision=revision, + sha256=sha256, + repo_type="dataset", + ) + + def _resolved_bundle(version: str = "5.2.0") -> ResolvedStage12Bundle: - countries = [] + countries: list[Stage12CountryBundle] = [] for country, package, package_version, dataset in ( ("us", "policyengine-us", "1.764.6", "populace_us_2024"), ("uk", "policyengine-uk", "2.90.2", "populace_uk_2023"), ): - revision = f"{dataset}-revision" + national_dataset = _dataset_fixture(dataset, "a" * 64) + datasets = [national_dataset] + region_dataset_identities = {"national": dataset} + if country == "us": + regional_dataset = _dataset_fixture("populace_us_2024_acs_local", "c" * 64) + datasets.append(regional_dataset) + region_dataset_identities.update( + { + "state": regional_dataset.identity, + "congressional_district": regional_dataset.identity, + } + ) + revision = national_dataset.artifact_revision countries.append( Stage12CountryBundle( country=country, @@ -59,16 +84,9 @@ def _resolved_bundle(version: str = "5.2.0") -> ResolvedStage12Bundle: data_release_version=revision, data_artifact_revision=revision, default_dataset=dataset, - default_dataset_uri=f"hf://policyengine/data/{dataset}.h5@{revision}", - datasets=( - Stage12Dataset( - identity=dataset, - uri=f"hf://policyengine/data/{dataset}.h5@{revision}", - artifact_revision=revision, - sha256="a" * 64, - repo_type="dataset", - ), - ), + default_dataset_uri=national_dataset.uri, + datasets=tuple(datasets), + region_dataset_identities=region_dataset_identities, ) ) bundle = Stage12BundleManifest( @@ -131,6 +149,43 @@ def test_publish_writes_one_complete_v2_document() -> None: assert manifest.versions["5.2.0"].validation.validated is True +@pytest.mark.parametrize( + "region, expected_dataset", + [ + ("us", "populace_us_2024"), + ("CA", "populace_us_2024_acs_local"), + ("DC", "populace_us_2024_acs_local"), + ("state/ca", "populace_us_2024_acs_local"), + ("state/DC", "populace_us_2024_acs_local"), + ("congressional_district/CA-01", "populace_us_2024_acs_local"), + ("congressional_district/DC-01", "populace_us_2024_acs_local"), + ], +) +def test_manifest_round_trip_preserves_regional_dataset_selection( + region: str, expected_dataset: str +) -> None: + store = FakeStore() + worker = _worker() + publish_v2_manifest(store=store, worker=worker) + + _, loaded_worker, _ = V2ManifestLoader(store).resolve() + us = next( + country for country in loaded_worker.bundle.countries if country.country == "us" + ) + selected = select_stage12_dataset(us, region) + + assert selected.identity == expected_dataset + original_us = next( + country for country in worker.bundle.countries if country.country == "us" + ) + assert selected == next( + dataset + for dataset in original_us.datasets + if dataset.identity == expected_dataset + ) + assert us.region_dataset_identities == original_us.region_dataset_identities + + def test_loader_retains_its_last_valid_v2_document() -> None: store = FakeStore() expected = publish_v2_manifest(store=store, worker=_worker()) diff --git a/projects/policyengine-simulation-entry/src/policyengine_simulation_entry/stage12_adapter.py b/projects/policyengine-simulation-entry/src/policyengine_simulation_entry/stage12_adapter.py index c51e7aa2e..83b5c4e8e 100644 --- a/projects/policyengine-simulation-entry/src/policyengine_simulation_entry/stage12_adapter.py +++ b/projects/policyengine-simulation-entry/src/policyengine_simulation_entry/stage12_adapter.py @@ -7,7 +7,10 @@ from typing import Any, cast from uuid import UUID, uuid5 -from policyengine_simulation_contract.stage12_bundle import CountryId +from policyengine_simulation_contract.stage12_bundle import ( + CountryId, + select_stage12_dataset, +) from policyengine_simulation_contract.stage12_execution import ( BundleProvenance, DatasetArtifactMediaType, @@ -134,11 +137,7 @@ def adapt_annual_comparison( if not isinstance(region, str) or not region: return _skip(ComparisonSkipReason.UNSUPPORTED_REQUEST_SHAPE) - selected_dataset = next( - item - for item in bundle_country.datasets - if item.identity == bundle_country.default_dataset - ) + selected_dataset = select_stage12_dataset(bundle_country, region) requested_dataset = payload.get("data") accepted_dataset_values = { None, diff --git a/projects/policyengine-simulation-entry/tests/stage12_fixtures.py b/projects/policyengine-simulation-entry/tests/stage12_fixtures.py index 4133ac0c5..2e43715c1 100644 --- a/projects/policyengine-simulation-entry/tests/stage12_fixtures.py +++ b/projects/policyengine-simulation-entry/tests/stage12_fixtures.py @@ -36,6 +36,33 @@ def worker() -> V2WorkerVersion: ("uk", "policyengine-uk", "2.90.2", "populace_uk_2023"), ): revision = f"{dataset}-revision" + datasets = [ + Stage12Dataset( + identity=dataset, + uri=f"hf://policyengine/data/{dataset}.h5@{revision}", + artifact_revision=revision, + sha256="a" * 64, + repo_type="dataset", + ) + ] + region_dataset_identities = {"national": dataset} + if country == "us": + local_dataset = "populace_us_2024_acs_local" + local_revision = f"{local_dataset}-revision" + datasets.append( + Stage12Dataset( + identity=local_dataset, + uri=(f"hf://policyengine/data/{local_dataset}.h5@{local_revision}"), + artifact_revision=local_revision, + sha256="c" * 64, + repo_type="dataset", + ) + ) + region_dataset_identities = { + "national": dataset, + "state": local_dataset, + "congressional_district": local_dataset, + } countries.append( Stage12CountryBundle( country=country, @@ -48,15 +75,8 @@ def worker() -> V2WorkerVersion: data_artifact_revision=revision, default_dataset=dataset, default_dataset_uri=f"hf://policyengine/data/{dataset}.h5@{revision}", - datasets=( - Stage12Dataset( - identity=dataset, - uri=f"hf://policyengine/data/{dataset}.h5@{revision}", - artifact_revision=revision, - sha256="a" * 64, - repo_type="dataset", - ), - ), + datasets=tuple(datasets), + region_dataset_identities=region_dataset_identities, ) ) country_workers.append( diff --git a/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py b/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py index a8b2cce35..0b787b29f 100644 --- a/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py +++ b/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py @@ -56,6 +56,26 @@ def test_absent_data_resolves_exact_bundle_default() -> None: assert result.report.baseline.bundle.dataset.artifact_revision.endswith("-revision") +@pytest.mark.parametrize( + "region", + ["DC", "state/DC", "congressional_district/DC-01"], +) +def test_regional_request_resolves_declared_acs_local_dataset(region: str) -> None: + payload = {**eligible_payload(), "region": region} + + result = adapt_annual_comparison( + payload, + evaluation_id=EVALUATION_ID, + worker=worker(), + ) + + assert result.report is not None + baseline = result.report.baseline + assert baseline.bundle.dataset.identity == "populace_us_2024_acs_local" + assert baseline.population.artifact.uri == baseline.bundle.dataset.uri + assert baseline.population.artifact.content_sha256 == "c" * 64 + + def test_absent_package_versions_resolve_from_the_selected_bundle() -> None: payload = eligible_payload() payload.pop("version") diff --git a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py index d260dd5ad..a028b1a93 100644 --- a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py +++ b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py @@ -44,6 +44,9 @@ def install_identity_stubs(monkeypatch): dataset_uris={ "populace_cps": "hf://org/repo/populace_cps.h5@rev-abc", }, + dataset_data_versions={"populace_cps": "1.2.3"}, + dataset_revisions={"populace_cps": "rev-abc"}, + dataset_sha256s={"populace_cps": "feedbead" * 8}, ) receipt_entry = { "country": "us", diff --git a/projects/policyengine-simulation-executor/pyproject.toml b/projects/policyengine-simulation-executor/pyproject.toml index 2f2803f64..a1a63aa82 100644 --- a/projects/policyengine-simulation-executor/pyproject.toml +++ b/projects/policyengine-simulation-executor/pyproject.toml @@ -21,7 +21,7 @@ dependencies = [ "policyengine-stage12-persistence", # PolicyEngine.py's models extra is the only selector for core, country # models, and the SPM calculator. Do not duplicate those versions here. - "policyengine[models]==6.2.3", + "policyengine[models]==6.2.5", "tables>=3.10.2", "modal>=0.73.0", "policyengine-observability[fastapi,google,otlp-grpc]>=3.0.2,<4", @@ -46,7 +46,7 @@ packages = ["src/policyengine_simulation_executor"] # certified datasets and writes their receipt. [dependency-groups] modal-simulation-image = [ - "policyengine[models]==6.2.3", + "policyengine[models]==6.2.5", "fastapi>=0.115.0", "tables>=3.10.2", "policyengine-observability[fastapi,google,otlp-grpc]>=3.0.2,<4", diff --git a/projects/policyengine-simulation-executor/src/modal/app.py b/projects/policyengine-simulation-executor/src/modal/app.py index 4e96b2276..7e72dbb7a 100644 --- a/projects/policyengine-simulation-executor/src/modal/app.py +++ b/projects/policyengine-simulation-executor/src/modal/app.py @@ -261,7 +261,7 @@ def _set_modal_call_attributes(runtime) -> None: @app.function( image=simulation_image, cpu=8.0, - memory=32768, + memory=65536, timeout=3600, retries=0, max_containers=100, diff --git a/projects/policyengine-simulation-executor/src/modal/smoke_app.py b/projects/policyengine-simulation-executor/src/modal/smoke_app.py index 10b7673d0..b131ff7ed 100644 --- a/projects/policyengine-simulation-executor/src/modal/smoke_app.py +++ b/projects/policyengine-simulation-executor/src/modal/smoke_app.py @@ -11,6 +11,10 @@ check. That compares all selected package versions with the bundle manifest, reads the dataset-install receipt, and hashes both installed country datasets. +A separate 64-GiB function prepares the certified US regional dataset from +scratch and calculates Utah. This is real data/model coverage; package imports +and national-dataset hashes alone cannot establish ACS compatibility. + Runs the imports the deployed workers perform lazily at request time — ``run_simulation_impl``, the budget-window batch, and both shared libraries. @@ -21,15 +25,14 @@ from pathlib import Path import modal -from src.modal.app import build_runtime_simulation_image -from src.modal.static_runtime_files import add_static_runtime_files - from policyengine_simulation_executor.release_bundle import ( resolve_local_bundle_dataset_path, ) from policyengine_simulation_executor.uk_local_authority_metadata import ( detect_uk_local_authority_metadata_from_hdf, ) +from src.modal.app import build_runtime_simulation_image +from src.modal.static_runtime_files import add_static_runtime_files app = modal.App("policyengine-simulation-executor-smoke") @@ -165,4 +168,16 @@ def smoke_import_executor() -> dict: def main(): report = smoke_import_executor.remote() print(report) + regional_report = smoke_us_regional_dataset.remote() + print(regional_report) + print("US regional dataset preparation and calculation OK") print("executor image smoke OK") + + +@app.function(image=smoke_image, timeout=1800, cpu=8.0, memory=65536, max_containers=1) +def smoke_us_regional_dataset() -> dict[str, str | int]: + """Ephemeral PR check; do not publish outputs or change deployed workers.""" + from src.modal.us_regional_dataset_check import check_certified_us_regional_dataset + + report = check_certified_us_regional_dataset() + return report.model_dump() diff --git a/projects/policyengine-simulation-executor/src/modal/us_regional_dataset_check.py b/projects/policyengine-simulation-executor/src/modal/us_regional_dataset_check.py new file mode 100644 index 000000000..86f06d753 --- /dev/null +++ b/projects/policyengine-simulation-executor/src/modal/us_regional_dataset_check.py @@ -0,0 +1,150 @@ +"""PR-only preparation and calculation using the installed certified ACS source. + +This runs inside an ephemeral Modal function, not a deployed service. It uses +the same region resolver and .py preparation API as request workers, with no +mock loaders, unmanaged inputs, or precomputed year/baseline outputs. +""" + +from importlib.metadata import version +from tempfile import TemporaryDirectory +from typing import Literal + +import numpy as np +import pandas as pd +from pydantic import BaseModel, ConfigDict, Field + +from policyengine_simulation_executor.simulation_runtime import DatasetSelection + + +class USRegionalDatasetCheck(BaseModel): + """Identity and row counts for a successfully prepared Utah calculation.""" + + model_config = ConfigDict(extra="forbid", frozen=True) + + country: Literal["us"] = "us" + region: Literal["state/ut"] = "state/ut" + year: int + policyengine_version: str + dataset: str + artifact_revision: str + source_sha256: str = Field(pattern=r"^[0-9a-f]{64}$") + prepared_person_count: int = Field(gt=0) + household_count: int = Field(gt=0) + person_count: int = Field(gt=0) + + +def _require_certified_regional_selection( + selection: DatasetSelection, +) -> tuple[str, str]: + if selection.is_default or not selection.sha256 or not selection.artifact_revision: + raise ValueError("Utah must select a non-default certified dataset with pins") + return selection.sha256, selection.artifact_revision + + +def summarize_us_regional_calculation( + selection: DatasetSelection, + prepared_person: pd.DataFrame, + output_household: pd.DataFrame, + output_person: pd.DataFrame, + *, + year: int, + policyengine_version: str, +) -> USRegionalDatasetCheck: + """Reject missing participation, wrong geography, or invalid calculation outputs.""" + sha256, revision = _require_certified_regional_selection(selection) + for person in (prepared_person, output_person): + column = "takes_up_wic_if_eligible" + if ( + person.empty + or column not in person + or bool(person[column].isna().any()) + or not pd.api.types.is_bool_dtype(person[column].dtype) + ): + raise ValueError("WIC participation must be present and complete booleans") + state_fips = output_household["state_fips"] + if ( + output_household.empty + or bool(state_fips.isna().any()) + or not bool(state_fips.eq(49).all()) + ): + raise ValueError("Calculated households must be nonempty and scoped to Utah") + for frame, column in ( + (output_household, "household_net_income"), + (output_person, "wic"), + ): + values = frame[column] + if ( + not pd.api.types.is_numeric_dtype(values.dtype) + or not np.isfinite(values.to_numpy()).all() + ): + raise ValueError(f"Calculated {column} must contain only finite numbers") + return USRegionalDatasetCheck( + year=year, + policyengine_version=policyengine_version, + dataset=selection.name, + artifact_revision=revision, + source_sha256=sha256, + prepared_person_count=len(prepared_person), + household_count=len(output_household), + person_count=len(output_person), + ) + + +def check_certified_us_regional_dataset() -> USRegionalDatasetCheck: + """Prepare real ACS inputs and calculate Utah with the installed .py release.""" + from policyengine.tax_benefit_models.us.datasets import ( + PolicyEngineUSDataset, + USYearData, + ) + + from policyengine_simulation_executor.simulation_runtime import ( + DEFAULT_YEAR, + _build_simulation, + _country_module, + _resolve_dataset_selection, + _resolve_region, + ) + + params = {"country": "us", "region": "state/ut", "time_period": str(DEFAULT_YEAR)} + country = _country_module("us") + region = _resolve_region(country_module=country, country="us", params=params) + selection = _resolve_dataset_selection(params, region_resolution=region) + _require_certified_regional_selection(selection) + + # The production loader calls this exact API. Use an isolated directory + # here so a previously prepared year cannot hide a broken source/loader. + # The .py materializer downloads and verifies the bundle-pinned ACS file. + with TemporaryDirectory(prefix="policyengine-us-regional-check-") as data_folder: + datasets = country.ensure_datasets( + datasets=[selection.name], years=[DEFAULT_YEAR], data_folder=data_folder + ) + dataset = datasets[f"{selection.name}_{DEFAULT_YEAR}"] + if not isinstance(dataset, PolicyEngineUSDataset) or dataset.data is None: + raise RuntimeError("ACS preparation produced no US input data") + prepared_person = dataset.data.person + simulation = _build_simulation( + params, + dataset=dataset, + dataset_selection=selection, + policy=None, + scoping_strategy=region.scoping_strategy, + region_code=region.code, + ) + simulation.extra_variables = { + "household": ["state_fips", "household_net_income"], + "person": ["takes_up_wic_if_eligible", "wic"], + } + simulation.ensure() + if simulation.output_dataset is None: + raise RuntimeError("Utah calculation produced no output dataset") + output_data = simulation.output_dataset.data + if not isinstance(output_data, USYearData): + raise TypeError("Utah calculation produced no US entity tables") + return summarize_us_regional_calculation( + selection, + prepared_person, + output_data.household, + output_data.person, + year=DEFAULT_YEAR, + policyengine_version=version("policyengine"), + ) diff --git a/projects/policyengine-simulation-executor/src/modal/utils/update_version_registry.py b/projects/policyengine-simulation-executor/src/modal/utils/update_version_registry.py index 649a652a6..6ab2e88a5 100644 --- a/projects/policyengine-simulation-executor/src/modal/utils/update_version_registry.py +++ b/projects/policyengine-simulation-executor/src/modal/utils/update_version_registry.py @@ -31,6 +31,10 @@ class CountryBundleMetadata(TypedDict): default_dataset_uri: str dataset_uris: dict[str, str] dataset_repo_types: dict[str, str] + dataset_data_versions: dict[str, str] + dataset_revisions: dict[str, str] + dataset_sha256s: dict[str, str] + region_dataset_identities: dict[str, str] class BundleManifestMetadata(TypedDict): @@ -97,6 +101,10 @@ def _country_bundle_metadata(country: str) -> CountryBundleMetadata: "default_dataset_uri": bundle.default_dataset_uri, "dataset_uris": dict(bundle.dataset_uris), "dataset_repo_types": dict(bundle.dataset_repo_types), + "dataset_data_versions": dict(bundle.dataset_data_versions), + "dataset_revisions": dict(bundle.dataset_revisions), + "dataset_sha256s": dict(bundle.dataset_sha256s), + "region_dataset_identities": dict(bundle.region_dataset_identities), } diff --git a/projects/policyengine-simulation-executor/src/modal/v2_app.py b/projects/policyengine-simulation-executor/src/modal/v2_app.py index 357670a2a..60cbee22f 100644 --- a/projects/policyengine-simulation-executor/src/modal/v2_app.py +++ b/projects/policyengine-simulation-executor/src/modal/v2_app.py @@ -179,7 +179,7 @@ def validate_worker_uk() -> dict: @app.function( image=us_worker_image, cpu=8.0, - memory=32768, + memory=65536, timeout=3000, retries=0, max_containers=10, diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py index 9a2f5c511..ad615cc2c 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py @@ -56,6 +56,8 @@ logger = logging.getLogger(__name__) +# Precompute only the national default dataset and its baseline partitions. +# Regional datasets (including ACS-local) are prepared by the request workers. PRECOMPUTE_COUNTRY = "us" # The user-facing priority order (2026 first); with parallel waves it only # orders spawn submission, but keep it intentional. diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/release_bundle.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/release_bundle.py index 430dd6c3f..fdde1c4c8 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/release_bundle.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/release_bundle.py @@ -40,6 +40,10 @@ class CountryReleaseBundle: default_dataset_uri: str dataset_uris: Mapping[str, str] dataset_repo_types: Mapping[str, str] + dataset_data_versions: Mapping[str, str] + dataset_revisions: Mapping[str, str] + dataset_sha256s: Mapping[str, str] + region_dataset_identities: Mapping[str, str] def _normalise_country(country: str) -> str: @@ -118,6 +122,21 @@ def _dataset_uris_from_manifest(manifest) -> tuple[dict[str, str], dict[str, str return dataset_uris, dataset_repo_types +def _dataset_provenance_from_manifest( + manifest, +) -> tuple[dict[str, str], dict[str, str]]: + revisions: dict[str, str] = {} + sha256s: dict[str, str] = {} + for name, reference in manifest.datasets.items(): + revision = _reference_value(reference, "revision") + sha256 = _reference_value(reference, "sha256") + if isinstance(revision, str) and revision: + revisions[name] = revision + if isinstance(sha256, str) and sha256: + sha256s[name] = sha256 + return revisions, sha256s + + def _dataset_uris_from_release( *, data_release: Mapping, @@ -160,6 +179,61 @@ def _dataset_uris_from_release( return dataset_uris, dataset_repo_types +def _dataset_provenance_from_release( + data_release: Mapping, +) -> tuple[dict[str, str], dict[str, str]]: + revisions: dict[str, str] = {} + sha256s: dict[str, str] = {} + for name, reference in _mapping(data_release.get("datasets")).items(): + if not isinstance(reference, Mapping): + continue + revision = reference.get("revision") + sha256 = reference.get("sha256") + if isinstance(revision, str) and revision: + revisions[str(name)] = revision + if isinstance(sha256, str) and sha256: + sha256s[str(name)] = sha256 + return revisions, sha256s + + +def _region_dataset_identities( + *, + datasets: Mapping[str, object], + templates: Mapping[str, object], +) -> dict[str, str]: + dataset_paths = { + name: _reference_value(reference, "path") + for name, reference in datasets.items() + } + region_dataset_identities: dict[str, str] = {} + for region_type, template in templates.items(): + path_template = _reference_value(template, "path_template") + matching_datasets = [ + name for name, path in dataset_paths.items() if path == path_template + ] + if len(matching_datasets) != 1: + raise ValueError( + "PolicyEngine.py bundle region dataset " + f"{region_type!r} must identify exactly one dataset" + ) + region_dataset_identities[region_type] = matching_datasets[0] + return region_dataset_identities + + +def _region_dataset_identities_from_manifest(manifest) -> dict[str, str]: + return _region_dataset_identities( + datasets=manifest.datasets, + templates=manifest.region_datasets, + ) + + +def _region_dataset_identities_from_release(data_release: Mapping) -> dict[str, str]: + return _region_dataset_identities( + datasets=_mapping(data_release.get("datasets")), + templates=_mapping(data_release.get("region_datasets")), + ) + + def _current_policyengine_bundle() -> Mapping | None: try: from policyengine.bundle import get_current_bundle @@ -250,6 +324,8 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: default_dataset = manifest.default_dataset default_dataset_uri = manifest.default_dataset_uri dataset_uris, dataset_repo_types = _dataset_uris_from_manifest(manifest) + dataset_revisions, dataset_sha256s = _dataset_provenance_from_manifest(manifest) + region_dataset_identities = _region_dataset_identities_from_manifest(manifest) data_package: Mapping = {} data_release: Mapping = {} if bundle_metadata is not None: @@ -287,6 +363,14 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: ) dataset_uris.update(release_dataset_uris) dataset_repo_types.update(release_dataset_repo_types) + release_dataset_revisions, release_dataset_sha256s = ( + _dataset_provenance_from_release(data_release) + ) + dataset_revisions.update(release_dataset_revisions) + dataset_sha256s.update(release_dataset_sha256s) + region_dataset_identities = _region_dataset_identities_from_release( + data_release + ) data_package_version = _derive_data_package_version( bundled_data_package=data_package, manifest_data_package_version=manifest.data_package.version, @@ -303,6 +387,18 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: else manifest.data_package.repo_type ), ) + dataset_revisions.setdefault(default_dataset, str(data_artifact_revision)) + dataset_data_versions = dict(dataset_revisions) + dataset_data_versions[default_dataset] = str(data_version) + + unknown_regional_datasets = set(region_dataset_identities.values()).difference( + dataset_uris + ) + if unknown_regional_datasets: + raise ValueError( + "PolicyEngine.py bundle region mappings reference unknown datasets: " + + ", ".join(sorted(unknown_regional_datasets)) + ) return CountryReleaseBundle( country=country, @@ -317,6 +413,10 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: default_dataset_uri=str(default_dataset_uri), dataset_uris=dataset_uris, dataset_repo_types=dataset_repo_types, + dataset_data_versions=dataset_data_versions, + dataset_revisions=dataset_revisions, + dataset_sha256s=dataset_sha256s, + region_dataset_identities=region_dataset_identities, ) diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/simulation_runtime.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/simulation_runtime.py index a6cf23a3c..309c439c4 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/simulation_runtime.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/simulation_runtime.py @@ -54,6 +54,9 @@ class DatasetSelection: name: str uri: str is_default: bool + data_version: str | None = None + artifact_revision: str | None = None + sha256: str | None = None def _normalize_credentials_blob(creds_json: str) -> str: @@ -545,6 +548,36 @@ def _nondefault_data_folder(country: str, name: str, uri: str) -> str: return f"/tmp/policyengine-alternate-data/{identity}" +def dataset_data_folder(country: str, selection: DatasetSelection) -> str: + """Return the runtime folder assigned to a selected managed dataset.""" + + return ( + resolve_data_folder() + if selection.is_default + else _nondefault_data_folder(country, selection.name, selection.uri) + ) + + +def bundle_dataset_selection(country: str, name: str) -> DatasetSelection: + """Build the complete runtime selection for one certified dataset name.""" + + bundle = get_country_release_bundle(country) + try: + uri = bundle.dataset_uris[name] + except KeyError as exc: + raise ValueError( + f"Unsupported dataset {name!r} for country {country!r}" + ) from exc + return DatasetSelection( + name=name, + uri=uri, + is_default=name == bundle.default_dataset, + data_version=bundle.dataset_data_versions.get(name), + artifact_revision=bundle.dataset_revisions.get(name), + sha256=bundle.dataset_sha256s.get(name), + ) + + def _resolve_dataset_selection( params: dict[str, Any], *, @@ -563,7 +596,7 @@ def _resolve_dataset_selection( # Region metadata can contain a URI rather than a managed dataset name. for name, uri in bundle.dataset_uris.items(): if dataset_reference == name or dataset_reference == uri: - return DatasetSelection(name, uri, name == bundle.default_dataset) + return bundle_dataset_selection(country, name) for name, uri in bundle.dataset_uris.items(): if dataset_reference == runtime_dataset_uri( uri, @@ -571,7 +604,7 @@ def _resolve_dataset_selection( artifact_revision=bundle.data_artifact_revision, validate_hf=False, ): - return DatasetSelection(name, uri, name == bundle.default_dataset) + return bundle_dataset_selection(country, name) raise ValueError( f"Unsupported dataset {dataset_reference!r} for country {country!r}; " "choose a name in the certified release manifest" @@ -587,11 +620,7 @@ def _load_dataset( country = params.get("country", "us").lower() year = _parse_year(params) country_module = country_module or _country_module(country) - data_folder = ( - resolve_data_folder() - if selection.is_default - else _nondefault_data_folder(country, selection.name, selection.uri) - ) + data_folder = dataset_data_folder(country, selection) start = time.monotonic() load_options = {"years": [year], "data_folder": data_folder} @@ -762,7 +791,7 @@ def _run_simulation_impl_core( dataset=dataset, baseline=baseline, reform=reform, - resolved_data_version=None, + resolved_data_version=dataset_selection.data_version, resolved_region_code=region_resolution.code, runtime=runtime, uk_local_authority_metadata=uk_local_authority_metadata, diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/manifest.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/manifest.py index 93c3bd737..b5eb721cc 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/manifest.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/manifest.py @@ -58,7 +58,6 @@ def _normalize_stage12_bundle( ) countries = require_mapping(bundle.get("countries"), "countries") data_releases = require_mapping(bundle.get("data_releases"), "data_releases") - normalized_countries = tuple( normalize_country_bundle( country=country, diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/parsing.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/parsing.py index bba9d4ee6..c5df530fb 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/parsing.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/parsing.py @@ -204,6 +204,54 @@ def normalize_country_bundle( f"PolicyEngine.py bundle default dataset URI is inconsistent for {country!r}" ) + region_dataset_templates = require_mapping( + release.get("region_datasets"), + f"data_releases.{country}.region_datasets", + ) + dataset_paths = { + identity: require_text( + require_mapping( + dataset_items[identity], + f"data_releases.{country}.datasets.{identity}", + ).get("path"), + f"data_releases.{country}.datasets.{identity}.path", + ) + for identity in dataset_items + } + normalized_region_dataset_identities: dict[str, str] = {} + for region_type, raw_template in region_dataset_templates.items(): + normalized_region_type = require_text( + region_type, + f"data_releases.{country}.region_datasets.", + ) + template = require_mapping( + raw_template, + f"data_releases.{country}.region_datasets.{normalized_region_type}", + ) + path_template = require_text( + template.get("path_template"), + "data_releases." + f"{country}.region_datasets.{normalized_region_type}.path_template", + ) + matching_datasets = [ + identity + for identity, path in dataset_paths.items() + if path == path_template + ] + if len(matching_datasets) != 1: + raise Stage12BundleError( + "PolicyEngine.py bundle region dataset " + f"{country}.{normalized_region_type} must identify exactly one dataset" + ) + normalized_region_dataset_identities[normalized_region_type] = ( + matching_datasets[0] + ) + + if normalized_region_dataset_identities.get("national") != default_dataset: + raise Stage12BundleError( + f"PolicyEngine.py bundle national region dataset is inconsistent for {country!r}" + ) + certified = require_mapping( release.get("certified_data_artifact"), f"data_releases.{country}.certified_data_artifact", @@ -243,4 +291,5 @@ def normalize_country_bundle( default_dataset=default_dataset, default_dataset_uri=default_dataset_uri, datasets=tuple(datasets), + region_dataset_identities=normalized_region_dataset_identities, ) diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/simulation.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/simulation.py index cc483437f..72187878f 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/simulation.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/simulation.py @@ -11,7 +11,11 @@ import pandas as pd from policyengine_observability import ObservabilityRuntime -from policyengine_simulation_contract.stage12_bundle import CountryId +from policyengine_simulation_contract.stage12_bundle import ( + CountryId, + Stage12Dataset, + select_stage12_dataset, +) from policyengine_simulation_contract.stage12_execution import ( ComparisonRunLifecycleStatus, ComparisonSimulationRecord, @@ -76,7 +80,9 @@ def _require_context( raise ValueError("single-simulation output plan names another country") -def _require_installed_bundle(simulation: SimulationExecutionInput) -> None: +def _require_installed_bundle( + simulation: SimulationExecutionInput, +) -> Stage12Dataset: resolved = load_stage12_bundle() if resolved.bundle_manifest_sha256 != simulation.bundle.bundle_manifest_sha256: raise RuntimeError("installed bundle digest differs from the simulation input") @@ -87,12 +93,16 @@ def _require_installed_bundle(simulation: SimulationExecutionInput) -> None: for item in resolved.bundle.countries if item.country == simulation.geography.country ) + selected_dataset = select_stage12_dataset( + country, + simulation.geography.region, + ) expected = { "country_package_name": country.country_package_name, "country_package_version": country.country_package_version, - "dataset_identity": country.default_dataset, - "dataset_uri": country.default_dataset_uri, - "dataset_revision": country.data_artifact_revision, + "dataset_identity": selected_dataset.identity, + "dataset_uri": selected_dataset.uri, + "dataset_revision": selected_dataset.artifact_revision, } actual = { "country_package_name": simulation.bundle.country_package_name, @@ -103,6 +113,33 @@ def _require_installed_bundle(simulation: SimulationExecutionInput) -> None: } if actual != expected: raise RuntimeError("simulation provenance differs from the installed bundle") + return selected_dataset + + +def _require_runtime_dataset_selection( + *, + selection, + declared_dataset: Stage12Dataset, + simulation: SimulationExecutionInput, +) -> None: + actual = { + "identity": selection.name, + "uri": selection.uri, + "artifact_revision": selection.artifact_revision, + "sha256": selection.sha256, + } + expected = { + "identity": declared_dataset.identity, + "uri": declared_dataset.uri, + "artifact_revision": declared_dataset.artifact_revision, + "sha256": declared_dataset.sha256, + } + if actual != expected: + raise RuntimeError("runtime dataset selection differs from the Stage 12 bundle") + if simulation.population.artifact.content_sha256 != declared_dataset.sha256: + raise RuntimeError( + "population artifact digest differs from the Stage 12 bundle" + ) def calculate_simulation_frames( @@ -116,7 +153,7 @@ def calculate_simulation_frames( raise ValueError("Stage 12 society-wide worker requires a dataset population") if simulation.population.artifact.uri != simulation.bundle.dataset.uri: raise ValueError("population artifact differs from bundle dataset provenance") - _require_installed_bundle(simulation) + declared_dataset = _require_installed_bundle(simulation) country = simulation.geography.country params: dict[str, Any] = { "country": country, @@ -170,6 +207,11 @@ def calculate_simulation_frames( params, region_resolution=region, ) + _require_runtime_dataset_selection( + selection=dataset_selection, + declared_dataset=declared_dataset, + simulation=simulation, + ) dataset_span = ( runtime.span(STAGE12_SIMULATION_STAGES.name(Stage.DATASET_LOAD)) if runtime is not None diff --git a/projects/policyengine-simulation-executor/tests/integration/test_image_smoke_modal.py b/projects/policyengine-simulation-executor/tests/integration/test_image_smoke_modal.py index cf7919c9f..96652ef07 100644 --- a/projects/policyengine-simulation-executor/tests/integration/test_image_smoke_modal.py +++ b/projects/policyengine-simulation-executor/tests/integration/test_image_smoke_modal.py @@ -28,9 +28,11 @@ def test_executor_image_smoke(): cwd=PROJECT_ROOT, capture_output=True, text=True, - timeout=2400, + timeout=3600, + check=False, ) assert result.returncode == 0, result.stdout + result.stderr + assert "US regional dataset preparation and calculation OK" in result.stdout assert "executor image smoke OK" in result.stdout @@ -50,6 +52,7 @@ def test_stage12_static_runtime_file_image_smoke(): capture_output=True, text=True, timeout=2400, + check=False, ) assert result.returncode == 0, result.stdout + result.stderr assert "Stage 12 static runtime file smoke OK" in result.stdout diff --git a/projects/policyengine-simulation-executor/tests/test_canonical_spm.py b/projects/policyengine-simulation-executor/tests/test_canonical_spm.py index b091a9a23..151fdbc33 100644 --- a/projects/policyengine-simulation-executor/tests/test_canonical_spm.py +++ b/projects/policyengine-simulation-executor/tests/test_canonical_spm.py @@ -705,7 +705,9 @@ def test_gateway_submission_uses_registry_capability_before_spawn(monkeypatch): bundle = PolicyEngineBundle( model_version="test-only", policyengine_version="test-only", spm=CAPABILITY ) - monkeypatch.setattr(endpoints, "_build_policyengine_bundle", lambda *args: bundle) + monkeypatch.setattr( + endpoints, "_build_policyengine_bundle", lambda *args, **kwargs: bundle + ) spawn = Mock(return_value=SimpleNamespace(object_id="job")) monkeypatch.setattr( endpoints, diff --git a/projects/policyengine-simulation-executor/tests/test_modal_bundle_image.py b/projects/policyengine-simulation-executor/tests/test_modal_bundle_image.py index db42e8e02..bb99677aa 100644 --- a/projects/policyengine-simulation-executor/tests/test_modal_bundle_image.py +++ b/projects/policyengine-simulation-executor/tests/test_modal_bundle_image.py @@ -119,6 +119,9 @@ def test_modal_image_uses_policyengine_bundle_install(monkeypatch): app.gcp_secret, app.hf_secret, ] + function_options = dict(app.app.function_calls) + assert function_options["run_simulation"]["memory"] == 65536 + assert function_options["run_simulation_segment"]["memory"] == 32768 def _fake_manifest(): diff --git a/projects/policyengine-simulation-executor/tests/test_modal_scripts.py b/projects/policyengine-simulation-executor/tests/test_modal_scripts.py index 7a2dbffca..31cc46e68 100644 --- a/projects/policyengine-simulation-executor/tests/test_modal_scripts.py +++ b/projects/policyengine-simulation-executor/tests/test_modal_scripts.py @@ -686,6 +686,19 @@ def test_deploy_workflow_runs_precompute_before_deploy(self): in executor_job ) + def test_stage12_deploy_does_not_add_another_precompute(self): + reusable_workflow = ( + REPO_ROOT / ".github" / "workflows" / "simulation-deploy.reusable.yml" + ).read_text(encoding="utf-8") + stage12_job = reusable_workflow[ + reusable_workflow.index( + "\n deploy_stage12_v2:\n" + ) : reusable_workflow.index("\n publish_stage12_v2_manifest:\n") + ] + + assert "modal-precompute.sh" not in stage12_job + assert "POLICYENGINE_MANIFEST_DIGEST" not in stage12_job + def test_deploy_workflow_threads_force_recompute_to_script(self): """The manual recompute flag should reach the precompute script.""" deploy_workflow = ( diff --git a/projects/policyengine-simulation-executor/tests/test_precompute.py b/projects/policyengine-simulation-executor/tests/test_precompute.py index 07d5e4d40..b6c65b219 100644 --- a/projects/policyengine-simulation-executor/tests/test_precompute.py +++ b/projects/policyengine-simulation-executor/tests/test_precompute.py @@ -490,6 +490,11 @@ def fake_cohort_identity(year, group): data_version="1.2.3", data_artifact_revision="rev-abc", default_dataset="populace_cps", + region_dataset_identities={ + "national": "populace_cps", + "state": "populace_us_2024_acs_local", + "congressional_district": "populace_us_2024_acs_local", + }, ), ) return SimpleNamespace(existing=existing) @@ -537,6 +542,20 @@ def test_plan_enumerates_years_and_cohorts_with_presence(self, planning_stubs): assert [entry.year for entry in work.datasets] == [2027, 2025] assert len(work.baselines) == 5 + @pytest.mark.parametrize("force", [False, True]) + def test_regional_datasets_are_not_prebuilt(self, planning_stubs, force): + plan = precompute.plan_artifacts_impl("bucket-x") + work = precompute.select_work(plan, force=force) + + assert len(plan.datasets) == 3 + assert [entry.year for entry in work.datasets] == ( + [2026, 2027, 2025] if force else [2027, 2025] + ) + assert all( + "acs_local" not in artifact.path + for artifact in precompute.build_manifest(plan).artifacts + ) + def test_empty_partition_fails_loudly(self, planning_stubs, monkeypatch): from policyengine_simulation_executor import national_partition diff --git a/projects/policyengine-simulation-executor/tests/test_release_bundle.py b/projects/policyengine-simulation-executor/tests/test_release_bundle.py index f03c53c9d..ff6b0985a 100644 --- a/projects/policyengine-simulation-executor/tests/test_release_bundle.py +++ b/projects/policyengine-simulation-executor/tests/test_release_bundle.py @@ -108,6 +108,43 @@ def test_country_release_bundle_exposes_model_and_data_versions(): assert bundle.default_dataset_uri == release["default_dataset_uri"] +def test_country_release_bundle_exposes_regional_dataset_provenance(monkeypatch): + current = deepcopy(get_current_bundle()) + local_dataset = "populace_us_2024_acs_local" + current["data_releases"]["us"]["datasets"][local_dataset] = { + "path": f"{local_dataset}.h5", + "repo_id": "policyengine/populace-us", + "repo_type": "dataset", + "revision": "acs-local-release", + "sha256": "b" * 64, + } + current["data_releases"]["us"]["region_datasets"].update( + { + "state": {"path_template": f"{local_dataset}.h5"}, + "congressional_district": {"path_template": f"{local_dataset}.h5"}, + } + ) + local_reference = current["data_releases"]["us"]["datasets"][local_dataset] + monkeypatch.setattr( + release_bundle_module, + "_current_policyengine_bundle", + lambda: current, + ) + get_country_release_bundle.cache_clear() + + bundle = get_country_release_bundle("us") + + assert bundle.region_dataset_identities == { + "national": "populace_us_2024", + "state": local_dataset, + "congressional_district": local_dataset, + } + assert bundle.dataset_revisions[local_dataset] == local_reference["revision"] + assert bundle.dataset_data_versions[local_dataset] == local_reference["revision"] + assert bundle.dataset_sha256s[local_dataset] == local_reference["sha256"] + get_country_release_bundle.cache_clear() + + def test_policyengine_data_release_keeps_build_and_package_versions_distinct( policyengine_uk_data_release, ): diff --git a/projects/policyengine-simulation-executor/tests/test_simulation_output_builder.py b/projects/policyengine-simulation-executor/tests/test_simulation_output_builder.py index 4f0f9e51d..1e442839f 100644 --- a/projects/policyengine-simulation-executor/tests/test_simulation_output_builder.py +++ b/projects/policyengine-simulation-executor/tests/test_simulation_output_builder.py @@ -1127,7 +1127,7 @@ def serialize(self): "dataset": dataset, "baseline": baseline_simulation, "reform": reform_simulation, - "resolved_data_version": None, + "resolved_data_version": get_country_release_bundle("us").data_version, "resolved_region_code": "us", "runtime": runtime, "uk_local_authority_metadata": None, @@ -1287,8 +1287,11 @@ def test_dataset_selection_identifies_the_actual_bundle_default(): assert isinstance(omitted, DatasetSelection) assert omitted.name == bundle.default_dataset assert omitted.is_default + assert omitted.data_version == bundle.data_version + assert omitted.artifact_revision == bundle.dataset_revisions[bundle.default_dataset] assert alternate.name == nondefault assert not alternate.is_default + assert alternate.data_version == bundle.dataset_data_versions.get(nondefault) def test_dataset_selection_rejects_a_bundled_uri_as_public_data(): diff --git a/projects/policyengine-simulation-executor/tests/test_smoke_app.py b/projects/policyengine-simulation-executor/tests/test_smoke_app.py index 1e37adc7a..952a3abf7 100644 --- a/projects/policyengine-simulation-executor/tests/test_smoke_app.py +++ b/projects/policyengine-simulation-executor/tests/test_smoke_app.py @@ -2,6 +2,8 @@ import importlib import sys +from collections.abc import Callable +from dataclasses import dataclass import pytest from policyengine_simulation_contract.uk_geography import ( @@ -106,11 +108,68 @@ def test_uk_dataset_smoke_validates_installed_codes_against_packaged_resources( monkeypatch.setattr( smoke_module, "detect_uk_local_authority_metadata_from_hdf", - lambda path: observed.append(path) - or UKLocalAuthorityMetadata( - boundary_version=UKLocalAuthorityBoundaryVersion.LAD22 + lambda path: ( + observed.append(path) + or UKLocalAuthorityMetadata( + boundary_version=UKLocalAuthorityBoundaryVersion.LAD22 + ) ), ) assert smoke_module._validate_installed_uk_local_authority_dataset() == "lad22" assert observed == ["/installed/enhanced_frs_2024_25.h5"] + + +@dataclass +class StubRemote: + callback: Callable[[], dict[str, str | int]] + + def remote(self) -> dict[str, str | int]: + return self.callback() + + +def test_image_smoke_requires_real_regional_check_before_success( + monkeypatch, smoke_module, capsys +) -> None: + calls: list[str] = [] + + def imported() -> dict[str, str | int]: + calls.append("imports") + return {} + + def calculated() -> dict[str, str | int]: + calls.append("regional_calculation") + return {} + + monkeypatch.setattr(smoke_module, "smoke_import_executor", StubRemote(imported)) + monkeypatch.setattr( + smoke_module, "smoke_us_regional_dataset", StubRemote(calculated) + ) + smoke_module.main() + + assert calls == ["imports", "regional_calculation"] + assert ( + "US regional dataset preparation and calculation OK" in capsys.readouterr().out + ) + + +def test_real_regional_failure_cannot_be_reported_as_image_smoke_success( + monkeypatch, smoke_module, capsys +) -> None: + def failed() -> dict[str, str | int]: + raise ValueError("stored WIC participation has missing values") + + monkeypatch.setattr(smoke_module, "smoke_import_executor", StubRemote(dict)) + monkeypatch.setattr(smoke_module, "smoke_us_regional_dataset", StubRemote(failed)) + + with pytest.raises(ValueError, match="missing values"): + smoke_module.main() + assert "executor image smoke OK" not in capsys.readouterr().out + + +def test_regional_check_has_bounded_production_sized_resources(smoke_module) -> None: + options = dict(smoke_module.app.function_calls)["smoke_us_regional_dataset"] + assert options["memory"] == 65536 + assert options["cpu"] == 8.0 + assert options["timeout"] == 1800 + assert options["max_containers"] == 1 diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py b/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py index f92d3c3c4..beb6f0f44 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py @@ -16,6 +16,7 @@ normalize_stage12_bundle, resolve_stage12_bundle, ) +from policyengine_simulation_contract.stage12_bundle import select_stage12_dataset def _bundle() -> dict: @@ -48,6 +49,36 @@ def test_complete_packaged_bundle_selects_every_worker_dependency_and_dataset() assert len(json.dumps(bundle.model_dump(mode="json"), sort_keys=True)) > 100 +def test_normalized_bundle_selects_regional_dataset_from_packaged_metadata() -> None: + raw = _bundle() + local_identity = "populace_us_2024_acs_local" + raw["data_releases"]["us"]["datasets"][local_identity] = { + "path": f"{local_identity}.h5", + "repo_id": "policyengine/populace-us", + "repo_type": "dataset", + "revision": "acs-local-release", + "sha256": "b" * 64, + } + raw["data_releases"]["us"]["region_datasets"].update( + { + "state": {"path_template": f"{local_identity}.h5"}, + "congressional_district": {"path_template": f"{local_identity}.h5"}, + } + ) + + bundle = normalize_stage12_bundle(raw) + us = next(country for country in bundle.countries if country.country == "us") + + assert bundle.schema_version == 2 + assert select_stage12_dataset(us, "DC").identity == local_identity + assert select_stage12_dataset(us, "state/DC").identity == local_identity + assert ( + select_stage12_dataset(us, "congressional_district/DC-01").identity + == local_identity + ) + assert select_stage12_dataset(us, "us").identity == us.default_dataset + + @pytest.mark.parametrize( "mutate", [ @@ -70,6 +101,9 @@ def test_complete_packaged_bundle_selects_every_worker_dependency_and_dataset() lambda bundle: bundle["data_releases"]["us"]["certified_data_artifact"].update( {"sha256": "0" * 64} ), + lambda bundle: bundle["data_releases"]["us"]["region_datasets"].update( + {"state": {"path_template": "absent.h5"}} + ), ], ids=[ "missing-data-releases", @@ -79,6 +113,7 @@ def test_complete_packaged_bundle_selects_every_worker_dependency_and_dataset() "dataset-uri-revision-mismatch", "malformed-dataset-digest", "certified-artifact-mismatch", + "absent-regional-dataset", ], ) def test_missing_or_internally_inconsistent_bundle_values_fail_closed(mutate) -> None: diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_modal_app.py b/projects/policyengine-simulation-executor/tests/test_stage12_modal_app.py index ade681c12..7234aec18 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_modal_app.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_modal_app.py @@ -144,6 +144,10 @@ def test_v2_app_name_and_images_are_separate_and_bundle_derived(monkeypatch) -> "/root/policyengine_simulation_executor/static_runtime_files" ) assert static_runtime_files[3] == {"copy": True} + assert not any( + call[0] == "run_function" and call[1] == "fetch_artifacts" + for call in image.calls + ) def test_v2_image_retains_required_runtime_dependencies() -> None: @@ -199,6 +203,8 @@ def test_v2_app_declares_validation_workers_and_non_http_coordinator( assert functions["run_single_simulation_uk"]["image"] is module.uk_worker_image assert functions["run_single_simulation_us"]["timeout"] == 3000 assert functions["run_single_simulation_uk"]["timeout"] == 3000 + assert functions["run_single_simulation_us"]["memory"] == 65536 + assert functions["run_single_simulation_uk"]["memory"] == 32768 assert functions["run_single_simulation_us"]["max_containers"] == 10 assert functions["run_single_simulation_uk"]["max_containers"] == 10 assert functions["coordinate_report"]["image"] is module.coordinator_image diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py b/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py index 909d91a9b..3f4f58d90 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py @@ -23,6 +23,7 @@ StdoutLogDestination, configure, ) +from policyengine_simulation_contract.stage12_bundle import Stage12Dataset from policyengine_simulation_contract.stage12_execution import ( ArtifactMediaType, ArtifactReference, @@ -472,7 +473,21 @@ def test_calculator_uses_the_current_dataset_selection_contract(monkeypatch) -> country_module = object() region = type("Region", (), {"code": "us", "scoping_strategy": None})() - selection = object() + selection = simulation_runtime.DatasetSelection( + name="populace_us_2024", + uri="hf://policyengine/data/populace_us_2024.h5@revision", + is_default=True, + data_version="release", + artifact_revision="revision", + sha256="a" * 64, + ) + declared_dataset = Stage12Dataset( + identity=selection.name, + uri=selection.uri, + artifact_revision="revision", + sha256="a" * 64, + repo_type="dataset", + ) dataset = object() received: dict[str, object] = {} @@ -494,7 +509,11 @@ def ensure(self) -> None: received["ensured"] = True received["extras_at_ensure"] = self.extra_variables.copy() - monkeypatch.setattr(worker, "_require_installed_bundle", lambda _: None) + monkeypatch.setattr( + worker, + "_require_installed_bundle", + lambda _: declared_dataset, + ) monkeypatch.setattr( simulation_runtime, "setup_gcp_credentials", diff --git a/projects/policyengine-simulation-executor/tests/test_update_version_registry.py b/projects/policyengine-simulation-executor/tests/test_update_version_registry.py index 9bd1544c0..80f9670ac 100644 --- a/projects/policyengine-simulation-executor/tests/test_update_version_registry.py +++ b/projects/policyengine-simulation-executor/tests/test_update_version_registry.py @@ -86,6 +86,10 @@ def fake_country_bundle_metadata( "default_dataset_uri": f"hf://datasets/policyengine/{country}/default", "dataset_uris": {"default": f"hf://datasets/policyengine/{country}"}, "dataset_repo_types": {"default": "dataset"}, + "dataset_data_versions": {"default": "test-data-version"}, + "dataset_revisions": {"default": "test-revision"}, + "dataset_sha256s": {"default": "a" * 64}, + "region_dataset_identities": {"national": "default"}, } monkeypatch.setattr( diff --git a/projects/policyengine-simulation-executor/tests/test_us_regional_dataset_check.py b/projects/policyengine-simulation-executor/tests/test_us_regional_dataset_check.py new file mode 100644 index 000000000..52e1e5a6b --- /dev/null +++ b/projects/policyengine-simulation-executor/tests/test_us_regional_dataset_check.py @@ -0,0 +1,111 @@ +"""Unit coverage for reporting real ACS preparation and calculation results. + +The real data-loading check runs separately in the PR Modal image check. +These small frames test its acceptance conditions, not dataset loading. +""" + +from dataclasses import replace + +import numpy as np +import pandas as pd +import pytest + +from policyengine_simulation_executor.simulation_runtime import DatasetSelection +from src.modal.us_regional_dataset_check import summarize_us_regional_calculation + + +@pytest.fixture +def selection() -> DatasetSelection: + return DatasetSelection( + name="populace_us_2024_acs_local", + uri="hf://policyengine/populace-us/acs.h5@fixture-revision", + is_default=False, + artifact_revision="fixture-revision", + sha256="a" * 64, + ) + + +@pytest.fixture +def frames() -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]: + return ( + pd.DataFrame({"takes_up_wic_if_eligible": [True, False, True]}), + pd.DataFrame({"state_fips": [49], "household_net_income": [100.0]}), + pd.DataFrame({"takes_up_wic_if_eligible": [True, False], "wic": [10.0, 0.0]}), + ) + + +def test_reports_selected_source_and_nonempty_utah_calculation(selection, frames): + report = summarize_us_regional_calculation( + selection, *frames, year=2026, policyengine_version="fixture-version" + ) + + assert report.dataset == selection.name + assert report.source_sha256 == selection.sha256 + assert report.artifact_revision == "fixture-revision" + assert report.prepared_person_count == 3 + assert report.household_count == 1 + assert report.person_count == 2 + + +@pytest.mark.parametrize( + "change", [{"is_default": True}, {"sha256": None}, {"artifact_revision": None}] +) +def test_refuses_national_or_unpinned_source(selection, frames, change): + with pytest.raises(ValueError, match="non-default certified"): + summarize_us_regional_calculation( + replace(selection, **change), + *frames, + year=2026, + policyengine_version="fixture-version", + ) + + +@pytest.mark.parametrize("scope", [0, 2]) +def test_rejects_missing_prepared_or_calculated_wic_decisions(selection, frames, scope): + frames[scope]["takes_up_wic_if_eligible"] = frames[scope][ + "takes_up_wic_if_eligible" + ].astype("boolean") + frames[scope].loc[0, "takes_up_wic_if_eligible"] = None + with pytest.raises(ValueError, match="WIC participation"): + summarize_us_regional_calculation( + selection, *frames, year=2026, policyengine_version="fixture-version" + ) + + +@pytest.mark.parametrize( + "frame_index,column", [(1, "household_net_income"), (2, "wic")] +) +@pytest.mark.parametrize("invalid", [np.nan, np.inf, -np.inf]) +def test_rejects_nonfinite_calculated_values( + selection, frames, frame_index, column, invalid +): + frames[frame_index].loc[0, column] = invalid + with pytest.raises(ValueError, match="finite"): + summarize_us_regional_calculation( + selection, *frames, year=2026, policyengine_version="fixture-version" + ) + + +def test_refuses_a_calculation_that_was_not_scoped_to_utah(selection, frames): + frames[1].loc[0, "state_fips"] = 6 + with pytest.raises(ValueError, match="Utah"): + summarize_us_regional_calculation( + selection, *frames, year=2026, policyengine_version="fixture-version" + ) + + +@pytest.mark.parametrize("scope", [0, 2]) +def test_rejects_a_missing_participation_column(selection, frames, scope): + frames[scope].drop(columns="takes_up_wic_if_eligible", inplace=True) + with pytest.raises(ValueError, match="WIC participation"): + summarize_us_regional_calculation( + selection, *frames, year=2026, policyengine_version="fixture-version" + ) + + +def test_rejects_missing_household_geography(selection, frames): + frames[1]["state_fips"] = pd.Series([pd.NA], dtype="Int64") + with pytest.raises(ValueError, match="Utah"): + summarize_us_regional_calculation( + selection, *frames, year=2026, policyengine_version="fixture-version" + ) diff --git a/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py b/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py index 7caf5e078..12d93222d 100644 --- a/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py +++ b/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py @@ -32,6 +32,26 @@ "populace_us_2024_acs_local": "hf://policyengine/populace-us/populace_us_2024_acs_local.h5@us-local-revision", "calibration_diagnostics": "hf://policyengine/populace-us/calibration_diagnostics.json@us-artifact-revision", }, + "dataset_data_versions": { + "populace_us_2024": "populace-us-2024-test", + "populace_us_2024_acs_local": "us-local-revision", + "calibration_diagnostics": "us-artifact-revision", + }, + "dataset_revisions": { + "populace_us_2024": "us-artifact-revision", + "populace_us_2024_acs_local": "us-local-revision", + "calibration_diagnostics": "us-artifact-revision", + }, + "dataset_sha256s": { + "populace_us_2024": "a" * 64, + "populace_us_2024_acs_local": "b" * 64, + "calibration_diagnostics": "c" * 64, + }, + "region_dataset_identities": { + "national": "populace_us_2024", + "state": "populace_us_2024_acs_local", + "congressional_district": "populace_us_2024_acs_local", + }, }, "uk": { "model_version": "2.66.0", @@ -44,6 +64,19 @@ "populace_uk_2023": "hf://policyengine/populace-uk-private/populace_uk_2023.h5@uk-artifact-revision", "local_authority_weights": "hf://policyengine/policyengine-uk-data-private/local_authority_weights.h5@uk-artifact-revision", }, + "dataset_data_versions": { + "populace_uk_2023": "populace-uk-2023-test", + "local_authority_weights": "uk-artifact-revision", + }, + "dataset_revisions": { + "populace_uk_2023": "uk-artifact-revision", + "local_authority_weights": "uk-artifact-revision", + }, + "dataset_sha256s": { + "populace_uk_2023": "d" * 64, + "local_authority_weights": "e" * 64, + }, + "region_dataset_identities": {"national": "populace_uk_2023"}, }, } diff --git a/projects/policyengine-simulation-gateway/src/policyengine_simulation_gateway/endpoints.py b/projects/policyengine-simulation-gateway/src/policyengine_simulation_gateway/endpoints.py index c8315c881..f59b68ee2 100644 --- a/projects/policyengine-simulation-gateway/src/policyengine_simulation_gateway/endpoints.py +++ b/projects/policyengine-simulation-gateway/src/policyengine_simulation_gateway/endpoints.py @@ -131,29 +131,115 @@ def _revision_from_dataset_uri(dataset_uri: str | None) -> str | None: def _bundle_response_data_version( *, country_bundle: dict, + dataset_name: str | None, resolved_dataset: str | None, ) -> str | None: + dataset_versions = country_bundle.get("dataset_data_versions") + if isinstance(dataset_versions, dict) and isinstance(dataset_name, str): + data_version = dataset_versions.get(dataset_name) + if isinstance(data_version, str): + return data_version return _country_bundle_data_version(country_bundle) or _revision_from_dataset_uri( resolved_dataset ) +def _regional_dataset_type( + *, + country: str, + region: str | None, + region_group: list[str] | None, +) -> str | None: + if region_group is not None or region is None: + return None + normalized = region.strip().lower() + if normalized == country.lower(): + return None + if country.lower() == "us" and len(normalized) == 2 and normalized.isalpha(): + return "state" + region_type, separator, _ = normalized.partition("/") + return region_type if separator else None + + +def _resolve_dataset_name_from_app_bundle( + *, + country_bundle: dict, + country: str, + region: str | None, + region_group: list[str] | None, +) -> str | None: + default_dataset = country_bundle.get("default_dataset") + if not isinstance(default_dataset, str): + return None + region_type = _regional_dataset_type( + country=country, + region=region, + region_group=region_group, + ) + region_dataset_identities = country_bundle.get("region_dataset_identities") + if region_type is not None and isinstance(region_dataset_identities, dict): + regional_dataset = region_dataset_identities.get(region_type) + if isinstance(regional_dataset, str): + return regional_dataset + return default_dataset + + def _resolve_dataset_uri_from_app_bundle( *, app_bundle: dict, country: str, -) -> str | None: + region: str | None, + region_group: list[str] | None, +) -> tuple[str | None, str | None]: country_bundle = app_bundle.get(country.lower()) if not isinstance(country_bundle, dict): - return None - dataset_uri = country_bundle.get("default_dataset_uri") + return None, None + dataset_name = _resolve_dataset_name_from_app_bundle( + country_bundle=country_bundle, + country=country, + region=region, + region_group=region_group, + ) + dataset_uris = country_bundle.get("dataset_uris") + dataset_uri = ( + dataset_uris.get(dataset_name) + if isinstance(dataset_uris, dict) and isinstance(dataset_name, str) + else None + ) + if not isinstance(dataset_uri, str) and dataset_name == country_bundle.get( + "default_dataset" + ): + dataset_uri = country_bundle.get("default_dataset_uri") if not isinstance(dataset_uri, str): - return None - return runtime_dataset_uri( - dataset_uri, - default_revision=_country_bundle_data_package_version(country_bundle), - artifact_revision=_country_bundle_data_artifact_revision(country_bundle), - validate_hf=False, + return dataset_name, None + dataset_versions = country_bundle.get("dataset_data_versions") + dataset_revisions = country_bundle.get("dataset_revisions") + data_version = ( + dataset_versions.get(dataset_name) + if isinstance(dataset_versions, dict) + else None + ) + artifact_revision = ( + dataset_revisions.get(dataset_name) + if isinstance(dataset_revisions, dict) + else None + ) + return ( + dataset_name, + runtime_dataset_uri( + dataset_uri, + default_revision=( + data_version + if isinstance(data_version, str) + else _country_bundle_data_package_version(country_bundle) + ), + artifact_revision=( + artifact_revision + if isinstance(artifact_revision, str) + else _country_bundle_data_artifact_revision(country_bundle) + ), + validate_hf=False, + ), ) @@ -498,17 +584,23 @@ def _resolve_from_legacy_dicts( def _build_policyengine_bundle( country: str, resolution: RouteResolution, + *, + region: str | None = None, + region_group: list[str] | None = None, ) -> PolicyEngineBundle: app_bundle = resolution.bundle_manifest country_bundle = app_bundle.get(country.lower()) if not isinstance(country_bundle, dict): country_bundle = {} - resolved_dataset = _resolve_dataset_uri_from_app_bundle( + dataset_name, resolved_dataset = _resolve_dataset_uri_from_app_bundle( app_bundle=app_bundle, country=country, + region=region, + region_group=region_group, ) data_version = _bundle_response_data_version( country_bundle=country_bundle, + dataset_name=dataset_name, resolved_dataset=resolved_dataset, ) model_version = country_bundle.get("model_version") or resolution.response_version @@ -748,6 +840,8 @@ async def submit_simulation( bundle = _build_policyengine_bundle( request.country, route, + region=request.region, + region_group=request.region_group, ) _resolve_request_spm(request, bundle, route) except (ValueError, HuggingFaceDatasetReferenceError) as exc: @@ -844,6 +938,7 @@ async def submit_budget_window_batch( bundle = _build_policyengine_bundle( request.country, route, + region=request.region, ) _resolve_request_spm(request, bundle, route) except (ValueError, HuggingFaceDatasetReferenceError) as exc: diff --git a/projects/policyengine-simulation-gateway/tests/test_endpoints.py b/projects/policyengine-simulation-gateway/tests/test_endpoints.py index 098174be5..27cf4d772 100644 --- a/projects/policyengine-simulation-gateway/tests/test_endpoints.py +++ b/projects/policyengine-simulation-gateway/tests/test_endpoints.py @@ -472,7 +472,7 @@ def test__given_submission_without_dataset_name__then_bundle_uses_manifest_uri( ) assert "data" not in mock_modal["func"].last_payload - def test__given_bundled_us_overlay__then_rejects_data_key( + def test__given_certified_us_regional_dataset__then_rejects_data_key( self, mock_modal, client: TestClient ): response = client.post( @@ -574,8 +574,9 @@ def test__given_legacy_alias_in_bundle_snapshot__then_gateway_rejects_it( assert response.json()["detail"][0]["type"] == "extra_forbidden" assert mock_modal["func"].last_payload is None - def test__given_us_state_region_without_data__then_keeps_contract_and_uses_default_dataset( - self, mock_modal, client: TestClient + @pytest.mark.parametrize("region", ["UT", "state/UT"]) + def test__given_us_state_region__then_uses_regional_dataset_without_request_override( + self, mock_modal, client: TestClient, region: str ): mock_modal["dicts"]["simulation-api-us-versions"] = { "latest": "1.500.0", @@ -587,7 +588,53 @@ def test__given_us_state_region_without_data__then_keeps_contract_and_uses_defau json={ "country": "us", "scope": "macro", - "region": "state/UT", + "region": region, + "reform": {}, + }, + ) + + assert response.status_code == 200 + assert response.json()["policyengine_bundle"] == expected_bundle( + "us", + "1.500.0", + dataset="populace_us_2024_acs_local", + data_version="us-local-revision", + ) + assert "data" not in mock_modal["func"].last_payload + assert "data_version" not in mock_modal["func"].last_payload + + def test__given_us_congressional_district__then_uses_regional_dataset( + self, mock_modal, client: TestClient + ): + response = client.post( + "/simulate/economy/comparison", + json={ + "country": "us", + "scope": "macro", + "region": "congressional_district/DC-01", + "reform": {}, + }, + ) + + assert response.status_code == 200 + assert response.json()["policyengine_bundle"] == expected_bundle( + "us", + "1.500.0", + dataset="populace_us_2024_acs_local", + data_version="us-local-revision", + ) + assert "data" not in mock_modal["func"].last_payload + assert "data_version" not in mock_modal["func"].last_payload + + def test__given_us_region_group__then_keeps_national_dataset( + self, mock_modal, client: TestClient + ): + response = client.post( + "/simulate/economy/comparison", + json={ + "country": "us", + "scope": "macro", + "region_group": ["state/CA", "state/NY"], "reform": {}, }, ) @@ -1271,6 +1318,31 @@ def test__given_no_dataset_override__then_parent_batch_omits_dataset( assert response.status_code == 200 assert "data" not in mock_modal["func"].last_payload + def test__given_state_budget_window__then_reports_regional_dataset( + self, mock_modal, client: TestClient + ): + response = client.post( + "/simulate/economy/budget-window", + json={ + "country": "us", + "region": "state/DC", + "scope": "macro", + "reform": {}, + "start_year": "2026", + "window_size": 2, + }, + ) + + assert response.status_code == 200 + assert response.json()["policyengine_bundle"] == expected_bundle( + "us", + "1.500.0", + dataset="populace_us_2024_acs_local", + data_version="us-local-revision", + ) + assert "data" not in mock_modal["func"].last_payload + assert "data_version" not in mock_modal["func"].last_payload + def test__given_budget_window_submission__then_returns_parent_batch_job_id( self, mock_modal, client: TestClient ):