From 3eda91eb3aed82aa2d9f1db9a615c8e31aa1e978 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Mon, 5 Oct 2026 23:15:44 +0400 Subject: [PATCH 01/10] Route regional simulations through bundle defaults --- .../modal/utils/update_version_registry.py | 8 ++ .../release_bundle.py | 70 +++++++++++ .../simulation_runtime.py | 45 +++++-- .../tests/test_canonical_spm.py | 4 +- .../tests/test_release_bundle.py | 36 ++++++ .../tests/test_simulation_output_builder.py | 5 +- .../tests/test_update_version_registry.py | 4 + .../fixtures/gateway_endpoints.py | 32 +++++ .../endpoints.py | 115 ++++++++++++++++-- .../tests/test_endpoints.py | 78 +++++++++++- 10 files changed, 374 insertions(+), 23 deletions(-) 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..97c823f7a 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] + regional_dataset_defaults: 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), + "regional_dataset_defaults": dict(bundle.regional_dataset_defaults), } 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..45adfd9f4 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] + regional_dataset_defaults: 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,33 @@ 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 _regional_dataset_defaults(bundle: Mapping, country: str) -> dict[str, str]: + all_defaults = _mapping(bundle.get("regional_dataset_defaults")) + country_defaults = _mapping(all_defaults.get(country)) + return { + str(region_type): dataset + for region_type, dataset in country_defaults.items() + if isinstance(dataset, str) and dataset + } + + def _current_policyengine_bundle() -> Mapping | None: try: from policyengine.bundle import get_current_bundle @@ -250,6 +296,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) + regional_dataset_defaults: dict[str, str] = {} data_package: Mapping = {} data_release: Mapping = {} if bundle_metadata is not None: @@ -287,6 +335,12 @@ 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) + regional_dataset_defaults = _regional_dataset_defaults(bundle_manifest, country) data_package_version = _derive_data_package_version( bundled_data_package=data_package, manifest_data_package_version=manifest.data_package.version, @@ -303,6 +357,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(regional_dataset_defaults.values()).difference( + dataset_uris + ) + if unknown_regional_datasets: + raise ValueError( + "PolicyEngine.py bundle regional defaults reference unknown datasets: " + + ", ".join(sorted(unknown_regional_datasets)) + ) return CountryReleaseBundle( country=country, @@ -317,6 +383,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, + regional_dataset_defaults=regional_dataset_defaults, ) 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/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_release_bundle.py b/projects/policyengine-simulation-executor/tests/test_release_bundle.py index f03c53c9d..e92db0a52 100644 --- a/projects/policyengine-simulation-executor/tests/test_release_bundle.py +++ b/projects/policyengine-simulation-executor/tests/test_release_bundle.py @@ -108,6 +108,42 @@ 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["regional_dataset_defaults"] = { + "us": { + "state": local_dataset, + "congressional_district": local_dataset, + } + } + 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, + } + 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.regional_dataset_defaults == { + "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_update_version_registry.py b/projects/policyengine-simulation-executor/tests/test_update_version_registry.py index 9bd1544c0..9fdf59a99 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}, + "regional_dataset_defaults": {}, } monkeypatch.setattr( diff --git a/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py b/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py index 7caf5e078..6369254a6 100644 --- a/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py +++ b/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py @@ -32,6 +32,25 @@ "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, + }, + "regional_dataset_defaults": { + "state": "populace_us_2024_acs_local", + "congressional_district": "populace_us_2024_acs_local", + }, }, "uk": { "model_version": "2.66.0", @@ -44,6 +63,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, + }, + "regional_dataset_defaults": {}, }, } 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..56d6aa138 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, + ) + regional_defaults = country_bundle.get("regional_dataset_defaults") + if region_type is not None and isinstance(regional_defaults, dict): + regional_dataset = regional_defaults.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..ca61c53df 100644 --- a/projects/policyengine-simulation-gateway/tests/test_endpoints.py +++ b/projects/policyengine-simulation-gateway/tests/test_endpoints.py @@ -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 ): From 569f4113fb4c24c362650238407b552453dab776 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Mon, 5 Oct 2026 23:15:51 +0400 Subject: [PATCH 02/10] Select regional datasets in Stage 12 --- .../stage12_bundle.py | 33 +++++++++++- .../stage12_adapter.py | 11 ++-- .../tests/stage12_fixtures.py | 37 +++++++++---- .../tests/test_stage12_adapter.py | 20 +++++++ .../stage12_bundle/manifest.py | 5 ++ .../stage12_bundle/parsing.py | 25 +++++++++ .../stage12_runtime/simulation.py | 54 ++++++++++++++++--- .../tests/test_stage12_bundle.py | 35 ++++++++++++ .../tests/test_stage12_runtime.py | 23 +++++++- 9 files changed, 219 insertions(+), 24 deletions(-) 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..984284cb1 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,15 @@ class Stage12CountryBundle(StrictBundleModel): default_dataset: NonEmptyText default_dataset_uri: NonEmptyText datasets: tuple[Stage12Dataset, ...] + regional_dataset_defaults: dict[NonEmptyText, NonEmptyText] = Field( + default_factory=dict + ) 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 +57,31 @@ 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() + dataset_identity = country.default_dataset + if normalized_region != country.country: + if ( + country.country == "us" + and len(normalized_region) == 2 + and normalized_region.isalpha() + ): + dataset_identity = country.regional_dataset_defaults.get( + "state", country.default_dataset + ) + else: + region_type, separator, _ = normalized_region.partition("/") + if separator: + dataset_identity = country.regional_dataset_defaults.get( + region_type, country.default_dataset + ) + return next( + dataset for dataset in country.datasets if dataset.identity == dataset_identity + ) 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..bc7035feb 100644 --- a/projects/policyengine-simulation-entry/tests/stage12_fixtures.py +++ b/projects/policyengine-simulation-entry/tests/stage12_fixtures.py @@ -36,6 +36,32 @@ 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", + ) + ] + regional_dataset_defaults = {} + 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", + ) + ) + regional_dataset_defaults = { + "state": local_dataset, + "congressional_district": local_dataset, + } countries.append( Stage12CountryBundle( country=country, @@ -48,15 +74,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), + regional_dataset_defaults=regional_dataset_defaults, ) ) 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/src/policyengine_simulation_executor/stage12_bundle/manifest.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_bundle/manifest.py index 93c3bd737..8e18f2b33 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,6 +58,10 @@ def _normalize_stage12_bundle( ) countries = require_mapping(bundle.get("countries"), "countries") data_releases = require_mapping(bundle.get("data_releases"), "data_releases") + regional_dataset_defaults = require_mapping( + bundle.get("regional_dataset_defaults", {}), + "regional_dataset_defaults", + ) normalized_countries = tuple( normalize_country_bundle( @@ -66,6 +70,7 @@ def _normalize_stage12_bundle( countries=countries, packages=packages, data_releases=data_releases, + regional_dataset_defaults=regional_dataset_defaults, ) for country in supported_countries ) 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..336fd6cf4 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 @@ -68,6 +68,7 @@ def normalize_country_bundle( countries: Mapping[str, object], packages: Mapping[str, object], data_releases: Mapping[str, object], + regional_dataset_defaults: Mapping[str, object], ) -> Stage12CountryBundle: """Validate and normalize one country's model and dataset release.""" @@ -204,6 +205,29 @@ def normalize_country_bundle( f"PolicyEngine.py bundle default dataset URI is inconsistent for {country!r}" ) + country_regional_defaults = require_mapping( + regional_dataset_defaults.get(country, {}), + f"regional_dataset_defaults.{country}", + ) + normalized_regional_defaults: dict[str, str] = {} + for region_type, dataset_identity in country_regional_defaults.items(): + normalized_region_type = require_text( + region_type, + f"regional_dataset_defaults.{country}.", + ) + normalized_dataset_identity = require_text( + dataset_identity, + f"regional_dataset_defaults.{country}.{normalized_region_type}", + ) + if normalized_dataset_identity not in datasets_by_identity: + raise Stage12BundleError( + "PolicyEngine.py bundle regional dataset default " + f"{country}.{normalized_region_type} names an absent dataset" + ) + normalized_regional_defaults[normalized_region_type] = ( + normalized_dataset_identity + ) + certified = require_mapping( release.get("certified_data_artifact"), f"data_releases.{country}.certified_data_artifact", @@ -243,4 +267,5 @@ def normalize_country_bundle( default_dataset=default_dataset, default_dataset_uri=default_dataset_uri, datasets=tuple(datasets), + regional_dataset_defaults=normalized_regional_defaults, ) 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/test_stage12_bundle.py b/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py index f92d3c3c4..9ab65ee31 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["regional_dataset_defaults"] = { + "us": { + "state": local_identity, + "congressional_district": local_identity, + } + } + + 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.update( + {"regional_dataset_defaults": {"us": {"state": "absent"}}} + ), ], 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_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", From eaf402b2b65423a30c965405e1c19128c7a2899e Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Mon, 5 Oct 2026 23:16:01 +0400 Subject: [PATCH 03/10] Bake regional datasets into simulation images --- .../workflows/simulation-deploy.reusable.yml | 11 ++ .../fixtures/identity_stubs.py | 4 + .../src/modal/_image_setup.py | 9 +- .../src/modal/app.py | 24 +--- .../src/modal/artifact_manifest.py | 31 +++++ .../src/modal/v2_app.py | 41 +++++-- .../artifact_keys.py | 47 +++++--- .../precompute.py | 74 +++++++++--- .../precompute_models.py | 18 ++- .../tests/test_image_setup_fetch.py | 22 ++++ .../tests/test_modal_bundle_image.py | 3 + .../tests/test_modal_scripts.py | 6 +- .../tests/test_precompute.py | 107 +++++++++++++++--- .../tests/test_precompute_app.py | 5 + .../tests/test_stage12_modal_app.py | 16 +++ 15 files changed, 332 insertions(+), 86 deletions(-) create mode 100644 projects/policyengine-simulation-executor/src/modal/artifact_manifest.py diff --git a/.github/workflows/simulation-deploy.reusable.yml b/.github/workflows/simulation-deploy.reusable.yml index ac6d6e40a..46e5b2827 100644 --- a/.github/workflows/simulation-deploy.reusable.yml +++ b/.github/workflows/simulation-deploy.reusable.yml @@ -561,6 +561,15 @@ jobs: name="$(uv run python -c 'from policyengine_simulation_contract.stage12_manifest import v2_application_name; from policyengine_simulation_executor.stage12_bundle import load_stage12_bundle; print(v2_application_name(load_stage12_bundle().bundle.policyengine_version))')" echo "name=${name}" >> "$GITHUB_OUTPUT" + - name: Run Stage 12 artifact precompute + id: precompute + working-directory: projects/policyengine-simulation-executor + env: + MODAL_TOKEN_ID: ${{ secrets.MODAL_TOKEN_ID }} + MODAL_TOKEN_SECRET: ${{ secrets.MODAL_TOKEN_SECRET }} + POLICYENGINE_ARTIFACT_BUCKET: ${{ vars.POLICYENGINE_ARTIFACT_BUCKET }} + run: ../../.github/scripts/modal-precompute.sh "${{ inputs.modal_environment }}" "${{ inputs.force_recompute }}" + - name: Synchronize the separate Stage 12 runtime secret working-directory: projects/policyengine-simulation-executor env: @@ -582,6 +591,8 @@ jobs: STAGE12_EXPECT_US_VERSION: ${{ needs.prepare.outputs.us_version }} STAGE12_EXPECT_UK_VERSION: ${{ needs.prepare.outputs.uk_version }} STAGE12_ARTIFACT_BUCKET: ${{ vars.STAGE12_ARTIFACT_BUCKET }} + POLICYENGINE_MANIFEST_DIGEST: ${{ steps.precompute.outputs.manifest_digest }} + POLICYENGINE_ARTIFACT_BUCKET: ${{ vars.POLICYENGINE_ARTIFACT_BUCKET }} OBSERVABILITY_SERVICE_NAMESPACE: ${{ vars.OBSERVABILITY_SERVICE_NAMESPACE }} OBSERVABILITY_TRACE_PROJECT_ID: ${{ vars.OBSERVABILITY_TRACE_PROJECT_ID }} OBSERVABILITY_LOGGING_PROJECT_ID: ${{ vars.OBSERVABILITY_LOGGING_PROJECT_ID }} diff --git a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py index d260dd5ad..ba00f3acb 100644 --- a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py +++ b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py @@ -44,6 +44,10 @@ 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}, + regional_dataset_defaults={}, ) receipt_entry = { "country": "us", diff --git a/projects/policyengine-simulation-executor/src/modal/_image_setup.py b/projects/policyengine-simulation-executor/src/modal/_image_setup.py index f8f7183c3..d80e412b1 100644 --- a/projects/policyengine-simulation-executor/src/modal/_image_setup.py +++ b/projects/policyengine-simulation-executor/src/modal/_image_setup.py @@ -142,14 +142,17 @@ def fetch_artifacts(bucket: str, manifest, *, client=None): # survive into the image, and the layer starts empty of these # files anyway. for artifact in artifacts: - target = data_folder / artifact["filename"] + target = Path(artifact.get("destination") or data_folder / artifact["filename"]) + target.parent.mkdir(parents=True, exist_ok=True) bucket_handle.blob(artifact["path"]).download_to_filename(str(target)) logger.info("Fetched %s (%.1f MB)", target, target.stat().st_size / 1e6) absent = [ - str(data_folder / artifact["filename"]) + str(Path(artifact.get("destination") or data_folder / artifact["filename"])) for artifact in artifacts - if not (data_folder / artifact["filename"]).exists() + if not Path( + artifact.get("destination") or data_folder / artifact["filename"] + ).exists() ] if absent: raise RuntimeError(f"Artifact fetch did not produce expected files: {absent}") diff --git a/projects/policyengine-simulation-executor/src/modal/app.py b/projects/policyengine-simulation-executor/src/modal/app.py index 4e96b2276..2dc9817c0 100644 --- a/projects/policyengine-simulation-executor/src/modal/app.py +++ b/projects/policyengine-simulation-executor/src/modal/app.py @@ -34,6 +34,7 @@ get_bundled_package_version, ) from src.modal._image_setup import fetch_artifacts, snapshot_models +from src.modal.artifact_manifest import deploy_time_artifact_inputs from src.modal.bundle_data import bundle_data_install_command from src.modal.logging_redaction import redact_params_for_logging from src.modal.static_runtime_files import add_static_runtime_files @@ -141,26 +142,7 @@ def _deploy_time_artifact_inputs() -> tuple[str, dict | None]: model by design (the layer is import-restricted), so the dict exists only at that serialization boundary. """ - if not modal.is_local(): - return "", None - digest = os.environ.get("POLICYENGINE_MANIFEST_DIGEST") - bucket = os.environ.get("POLICYENGINE_ARTIFACT_BUCKET", "") - if not digest: - return bucket, None - from policyengine_simulation_executor.artifact_store import ArtifactStore - from policyengine_simulation_executor.precompute_models import ArtifactManifest - - store = ArtifactStore(bucket or None) - payload = store.read_manifest(digest) - if payload is None: - raise RuntimeError( - f"Artifact manifest {digest} is not in the store: the deploy " - "must consume a digest published by the precompute run." - ) - # Validate on the runner so a corrupt or foreign manifest fails the - # deploy here, not mid-image-build. - manifest = ArtifactManifest.model_validate(payload) - return store.bucket_name, manifest.canonical_payload() + return deploy_time_artifact_inputs() _ARTIFACT_BUCKET, _DEPLOY_MANIFEST = _deploy_time_artifact_inputs() @@ -261,7 +243,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/artifact_manifest.py b/projects/policyengine-simulation-executor/src/modal/artifact_manifest.py new file mode 100644 index 000000000..7c0119789 --- /dev/null +++ b/projects/policyengine-simulation-executor/src/modal/artifact_manifest.py @@ -0,0 +1,31 @@ +"""Resolve a content-addressed precompute manifest during Modal deployment.""" + +from __future__ import annotations + +import os + + +def deploy_time_artifact_inputs() -> tuple[str, dict | None]: + """Return the artifact bucket and validated manifest on the deploy runner.""" + + import modal + + if not modal.is_local(): + return "", None + digest = os.environ.get("POLICYENGINE_MANIFEST_DIGEST") + bucket = os.environ.get("POLICYENGINE_ARTIFACT_BUCKET", "") + if not digest: + return bucket, None + + from policyengine_simulation_executor.artifact_store import ArtifactStore + from policyengine_simulation_executor.precompute_models import ArtifactManifest + + store = ArtifactStore(bucket or None) + payload = store.read_manifest(digest) + if payload is None: + raise RuntimeError( + f"Artifact manifest {digest} is not in the store: the deploy " + "must consume a digest published by the precompute run." + ) + manifest = ArtifactManifest.model_validate(payload) + return store.bucket_name, manifest.canonical_payload() diff --git a/projects/policyengine-simulation-executor/src/modal/v2_app.py b/projects/policyengine-simulation-executor/src/modal/v2_app.py index 357670a2a..2fce9b61c 100644 --- a/projects/policyengine-simulation-executor/src/modal/v2_app.py +++ b/projects/policyengine-simulation-executor/src/modal/v2_app.py @@ -37,10 +37,12 @@ assertion_values, load_stage12_bundle, ) +from src.modal._image_setup import fetch_artifacts +from src.modal.artifact_manifest import deploy_time_artifact_inputs from src.modal.bundle_data import bundle_data_install_command from src.modal.static_runtime_files import add_static_runtime_files -STAGE12_DATA_DIR = "/opt/policyengine/stage12-data" +STAGE12_DATA_DIR = "/opt/policyengine/data" _UV_PROJECT_DIR = str(Path(__file__).resolve().parents[2]) if modal.is_local() else "." RESOLVED_BUNDLE = load_stage12_bundle() BUNDLE_VALUES = assertion_values(RESOLVED_BUNDLE.bundle) @@ -82,6 +84,7 @@ def _external_assertions(environment: dict[str, str]) -> dict[str, str]: hf_secret, comparison_runtime_secret, ] +_ARTIFACT_BUCKET, _DEPLOY_MANIFEST = deploy_time_artifact_inputs() def _country_bundle(country: CountryId): @@ -90,7 +93,11 @@ def _country_bundle(country: CountryId): ) -def build_v2_image(countries: tuple[CountryId, ...]) -> modal.Image: +def build_v2_image( + countries: tuple[CountryId, ...], + *, + include_us_artifacts: bool = False, +) -> modal.Image: country_values = { country: _country_bundle(country).model_dump(mode="json") for country in countries @@ -114,6 +121,9 @@ def build_v2_image(countries: tuple[CountryId, ...]) -> modal.Image: { **modal_image_environment(), "POLICYENGINE_DATA_FOLDER": STAGE12_DATA_DIR, + "POLICYENGINE_BUNDLE_RECEIPT": ( + f"{STAGE12_DATA_DIR}/.policyengine-bundle-receipt.json" + ), "STAGE12_BUNDLE_MANIFEST_SHA256": ( RESOLVED_BUNDLE.bundle_manifest_sha256 ), @@ -124,19 +134,28 @@ def build_v2_image(countries: tuple[CountryId, ...]) -> modal.Image: ), } ) - .add_local_python_source( - "src.modal", - "policyengine_simulation_executor", - "policyengine_simulation_observability", - "policyengine_simulation_contract", - "policyengine_stage12_persistence", - copy=True, + ) + if include_us_artifacts: + image = image.run_function( + fetch_artifacts, + args=(_ARTIFACT_BUCKET, _DEPLOY_MANIFEST), + secrets=[gcp_secret], + cpu=2.0, + memory=4096, + timeout=900, ) + image = image.add_local_python_source( + "src.modal", + "policyengine_simulation_executor", + "policyengine_simulation_observability", + "policyengine_simulation_contract", + "policyengine_stage12_persistence", + copy=True, ) return add_static_runtime_files(image, uv_project_dir=_UV_PROJECT_DIR) -us_worker_image = build_v2_image(("us",)) +us_worker_image = build_v2_image(("us",), include_us_artifacts=True) uk_worker_image = build_v2_image(("uk",)) coordinator_image = build_v2_image(("us", "uk")) @@ -179,7 +198,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/artifact_keys.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py index 5cdcffba9..077255930 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py @@ -302,8 +302,12 @@ def _certification_fingerprint(country: str) -> Optional[str]: return fingerprint if isinstance(fingerprint, str) and fingerprint else None -def collect_dataset_identity(country: str, year: int) -> DatasetArtifactIdentity: - """Identity of the default single-year dataset for this installed bundle. +def collect_dataset_identity( + country: str, + year: int, + dataset: str | None = None, +) -> DatasetArtifactIdentity: + """Identity of one certified single-year dataset for this installed bundle. Reads the same sources the runtime trusts: the release bundle (versions), the bundle receipt (content sha), and the manifest certification @@ -311,32 +315,47 @@ def collect_dataset_identity(country: str, year: int) -> DatasetArtifactIdentity runtime lookup uses, so filename coupling is inherited, not re-derived. """ - from policyengine.provenance.manifest import ( - dataset_logical_name, - resolve_dataset_reference, - ) + from policyengine.provenance.manifest import dataset_logical_name from policyengine_simulation_executor.release_bundle import ( get_country_release_bundle, ) bundle = get_country_release_bundle(country) - stem = dataset_logical_name( - resolve_dataset_reference(bundle.country, bundle.default_dataset) - ) + selected_dataset = dataset or bundle.default_dataset + try: + selected_uri = bundle.dataset_uris[selected_dataset] + except KeyError as exc: + raise ValueError( + f"Unknown certified dataset {selected_dataset!r} for {bundle.country!r}" + ) from exc + stem = dataset_logical_name(selected_uri) + data_version = bundle.dataset_data_versions.get(selected_dataset) + artifact_revision = bundle.dataset_revisions.get(selected_dataset) + if not data_version or not artifact_revision: + raise ValueError( + f"Certified dataset {selected_dataset!r} has incomplete provenance" + ) + is_default = selected_dataset == bundle.default_dataset from policyengine_simulation_executor.spm import normalize_runtime_spm selection = normalize_runtime_spm({"country": country, "time_period": year}) return DatasetArtifactIdentity( spm=selection, country=bundle.country, - dataset=bundle.default_dataset, + dataset=selected_dataset, stem=stem, year=int(year), - data_version=bundle.data_version, - data_artifact_revision=bundle.data_artifact_revision, - source_sha256=_receipt_source_sha256(bundle.country, bundle.data_version), - data_build_fingerprint=_certification_fingerprint(bundle.country), + data_version=data_version, + data_artifact_revision=artifact_revision, + source_sha256=( + _receipt_source_sha256(bundle.country, bundle.data_version) + if is_default + else bundle.dataset_sha256s.get(selected_dataset) + ), + data_build_fingerprint=( + _certification_fingerprint(bundle.country) if is_default else None + ), model_version=bundle.model_version, policyengine_version=bundle.policyengine_version, ) 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..b34026339 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py @@ -61,7 +61,7 @@ # orders spawn submission, but keep it intentional. PRECOMPUTE_YEARS = [2026, 2027, 2025] -MANIFEST_SCHEMA = "mf1" +MANIFEST_SCHEMA = "mf2" # The stdout contract the CI deploy job parses; keep format changes in # lockstep with the workflow side (covered by a regression test). @@ -151,6 +151,7 @@ def build_manifest(plan: PrecomputePlan) -> ArtifactManifest: filename=entry.filename, year=entry.year, digest=entry.digest, + destination=entry.runtime_destination, ) for entry in plan.datasets ] + [ @@ -160,6 +161,7 @@ def build_manifest(plan: PrecomputePlan) -> ArtifactManifest: filename=entry.path.rsplit("/", maxsplit=1)[-1], year=entry.year, digest=entry.digest, + destination=entry.runtime_destination, ) for entry in plan.baselines ] @@ -173,6 +175,8 @@ def build_manifest(plan: PrecomputePlan) -> ArtifactManifest: def plan_artifacts_impl(bucket: str) -> PrecomputePlan: """Compute every expected artifact identity and its store presence.""" + from pathlib import Path + from policyengine_simulation_executor.artifact_keys import ( collect_dataset_identity, ) @@ -183,6 +187,11 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: from policyengine_simulation_executor.release_bundle import ( get_country_release_bundle, ) + from policyengine_simulation_executor.simulation_runtime import ( + bundle_dataset_selection, + dataset_data_folder, + resolve_data_folder, + ) store = ArtifactStore(bucket) groups = national_region_groups(PRECOMPUTE_COUNTRY) @@ -191,17 +200,40 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: datasets: list[DatasetPlanEntry] = [] baselines: list[BaselinePlanEntry] = [] + bundle = get_country_release_bundle(PRECOMPUTE_COUNTRY) + dataset_names = [ + bundle.default_dataset, + *sorted( + set(bundle.regional_dataset_defaults.values()).difference( + {bundle.default_dataset} + ) + ), + ] for year in PRECOMPUTE_YEARS: - dataset_identity = collect_dataset_identity(PRECOMPUTE_COUNTRY, year) - datasets.append( - DatasetPlanEntry( - year=year, - digest=dataset_identity.digest, - path=dataset_identity.store_path, - filename=dataset_identity.filename, - exists=store.exists(dataset_identity.store_path), + for dataset_name in dataset_names: + dataset_identity = collect_dataset_identity( + PRECOMPUTE_COUNTRY, + year, + dataset_name, + ) + selection = bundle_dataset_selection( + PRECOMPUTE_COUNTRY, + dataset_name, + ) + datasets.append( + DatasetPlanEntry( + year=year, + digest=dataset_identity.digest, + path=dataset_identity.store_path, + filename=dataset_identity.filename, + dataset=dataset_name, + runtime_destination=str( + Path(dataset_data_folder(PRECOMPUTE_COUNTRY, selection)) + / dataset_identity.filename + ), + exists=store.exists(dataset_identity.store_path), + ) ) - ) for group in groups: identity = cohort_identity(year, group) baselines.append( @@ -212,11 +244,14 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: digest=identity.digest, path=identity.store_path, simulation_id=identity.simulation_id, + runtime_destination=str( + Path(resolve_data_folder()) + / identity.store_path.rsplit("/", maxsplit=1)[-1] + ), exists=store.exists(identity.store_path), ) ) - bundle = get_country_release_bundle(PRECOMPUTE_COUNTRY) receipt = BundleVersionIdentity( policyengine_version=bundle.policyengine_version, model_version=bundle.model_version, @@ -235,15 +270,17 @@ def build_dataset_impl(bucket: str, expected: DatasetPlanEntry) -> DatasetBuildR collect_dataset_identity, ) from policyengine_simulation_executor.artifact_store import ArtifactStore - from policyengine_simulation_executor.release_bundle import ( - get_country_release_bundle, - ) from policyengine_simulation_executor.simulation_runtime import ( _country_module, - resolve_data_folder, + bundle_dataset_selection, + dataset_data_folder, ) - identity = collect_dataset_identity(PRECOMPUTE_COUNTRY, expected.year) + identity = collect_dataset_identity( + PRECOMPUTE_COUNTRY, + expected.year, + expected.dataset, + ) if identity.store_path != expected.path: raise RuntimeError( "Planned and in-container dataset identities disagree " @@ -251,11 +288,12 @@ def build_dataset_impl(bucket: str, expected: DatasetPlanEntry) -> DatasetBuildR "under a mismatched key." ) - data_folder = resolve_data_folder() + selection = bundle_dataset_selection(PRECOMPUTE_COUNTRY, expected.dataset) + data_folder = dataset_data_folder(PRECOMPUTE_COUNTRY, selection) country_module = _country_module(PRECOMPUTE_COUNTRY) started = time.monotonic() country_module.ensure_datasets( - datasets=[get_country_release_bundle(PRECOMPUTE_COUNTRY).default_dataset], + datasets=[expected.dataset], years=[expected.year], data_folder=data_folder, ) diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py index fd5332fd2..0cfee633b 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py @@ -22,7 +22,7 @@ from typing import Any, Literal, Optional, Protocol -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator # Mirrors baseline_artifacts.OUTCOME_*; the artifact outcome a worker # observed when ensuring a baseline. @@ -56,6 +56,8 @@ class DatasetPlanEntry(_StrictModel): digest: str path: str filename: str + dataset: str + runtime_destination: str exists: bool @@ -68,6 +70,7 @@ class BaselinePlanEntry(_StrictModel): digest: str path: str simulation_id: str + runtime_destination: str exists: bool @@ -94,6 +97,7 @@ class ManifestArtifact(_StrictModel): filename: str year: int digest: str + destination: str | None = None class ArtifactManifest(_StrictModel): @@ -103,14 +107,22 @@ class ArtifactManifest(_StrictModel): # "schema" is the wire key; the Python name avoids shadowing # BaseModel's deprecated .schema attribute. - manifest_schema: str = Field(alias="schema") + manifest_schema: Literal["mf1", "mf2"] = Field(alias="schema") country: str receipt: BundleVersionIdentity artifacts: list[ManifestArtifact] + @model_validator(mode="after") + def require_mf2_destinations(self) -> "ArtifactManifest": + if self.manifest_schema == "mf2" and any( + artifact.destination is None for artifact in self.artifacts + ): + raise ValueError("mf2 artifacts must declare a runtime destination") + return self + def canonical_payload(self) -> dict[str, Any]: """The exact dict shape that gets digested and stored.""" - return self.model_dump(by_alias=True) + return self.model_dump(by_alias=True, exclude_none=True) class DeployedMarker(_StrictModel): diff --git a/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py b/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py index 2f4a71fe5..ac0d5ad3a 100644 --- a/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py +++ b/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py @@ -107,6 +107,28 @@ def test_downloads_every_artifact_under_its_runtime_filename( target = data_folder / artifact["filename"] assert target.read_bytes() == f"bytes-{artifact['digest']}".encode() + def test_mf2_downloads_to_each_declared_runtime_destination( + self, fake_client, data_folder, tmp_path + ): + manifest = _manifest() + manifest["schema"] = "mf2" + destinations = [ + tmp_path / "national" / "populace_year_2026.h5", + tmp_path / "regional" / "bl1-aaaa.h5", + ] + for artifact, destination in zip( + manifest["artifacts"], destinations, strict=True + ): + artifact["destination"] = str(destination) + + validated = ArtifactManifest.model_validate(manifest).canonical_payload() + fetch_artifacts("test-bucket", validated, client=fake_client) + + assert [destination.read_bytes() for destination in destinations] == [ + b"bytes-d1", + b"bytes-b1", + ] + def test_missing_store_object_fails_before_any_download( self, fake_client, data_folder ): 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..74204688f 100644 --- a/projects/policyengine-simulation-executor/tests/test_modal_scripts.py +++ b/projects/policyengine-simulation-executor/tests/test_modal_scripts.py @@ -728,13 +728,13 @@ def test_deploy_workflow_threads_manifest_and_store_credentials_to_deploy(self): assert "GCP_CREDENTIALS_JSON: ${{ secrets.GCP_CREDENTIALS_JSON }}" in ( deploy_step ) - # The precompute, deploy, and record-marker steps each carry the - # bucket var. + # The direct precompute, direct deploy, Stage 12 precompute, + # Stage 12 deploy, and record-marker steps each carry the bucket var. assert ( reusable_workflow.count( "POLICYENGINE_ARTIFACT_BUCKET: ${{ vars.POLICYENGINE_ARTIFACT_BUCKET }}" ) - == 3 + == 5 ) def test_precompute_is_gated_by_the_shared_secret_sync(self): diff --git a/projects/policyengine-simulation-executor/tests/test_precompute.py b/projects/policyengine-simulation-executor/tests/test_precompute.py index 07d5e4d40..4d6ba9220 100644 --- a/projects/policyengine-simulation-executor/tests/test_precompute.py +++ b/projects/policyengine-simulation-executor/tests/test_precompute.py @@ -46,6 +46,8 @@ def _plan() -> PrecomputePlan: digest="d1", path="datasets/us/d1/populace_year_2026.h5", filename="populace_year_2026.h5", + dataset="populace_cps", + runtime_destination="/opt/policyengine/data/populace_year_2026.h5", exists=True, ), DatasetPlanEntry( @@ -53,6 +55,8 @@ def _plan() -> PrecomputePlan: digest="d2", path="datasets/us/d2/populace_year_2027.h5", filename="populace_year_2027.h5", + dataset="populace_cps", + runtime_destination="/opt/policyengine/data/populace_year_2027.h5", exists=False, ), ], @@ -64,6 +68,7 @@ def _plan() -> PrecomputePlan: digest="b1", path="baselines/us/b1/bl1-aaaa.h5", simulation_id="bl1-aaaa", + runtime_destination="/opt/policyengine/data/bl1-aaaa.h5", exists=True, ), BaselinePlanEntry( @@ -73,6 +78,7 @@ def _plan() -> PrecomputePlan: digest="b2", path="baselines/us/b2/bl1-bbbb.h5", simulation_id="bl1-bbbb", + runtime_destination="/opt/policyengine/data/bl1-bbbb.h5", exists=False, ), ], @@ -111,6 +117,13 @@ def test_artifact_outcome_mirrors_runtime_constants(self): ba.OUTCOME_MISS, } + def test_mf2_manifest_requires_runtime_destinations(self): + manifest = precompute.build_manifest(_plan()).canonical_payload() + del manifest["artifacts"][0]["destination"] + + with pytest.raises(ValidationError, match="runtime destination"): + ArtifactManifest.model_validate(manifest) + class TestPlanning: def test_select_work_computes_only_misses(self): @@ -157,7 +170,7 @@ def test_manifest_payload_is_wire_compatible(self): identically to the raw dict shape already published to the store (key names included — note "schema", not "manifest_schema").""" raw = { - "schema": "mf1", + "schema": "mf2", "country": "us", "receipt": { "policyengine_version": "4.22.0", @@ -173,6 +186,7 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "populace_year_2026.h5", "year": 2026, "digest": "d1", + "destination": "/opt/policyengine/data/populace_year_2026.h5", }, { "type": "dataset", @@ -180,6 +194,7 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "populace_year_2027.h5", "year": 2027, "digest": "d2", + "destination": "/opt/policyengine/data/populace_year_2027.h5", }, { "type": "baseline", @@ -187,6 +202,7 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "bl1-aaaa.h5", "year": 2026, "digest": "b1", + "destination": "/opt/policyengine/data/bl1-aaaa.h5", }, { "type": "baseline", @@ -194,6 +210,7 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "bl1-bbbb.h5", "year": 2027, "digest": "b2", + "destination": "/opt/policyengine/data/bl1-bbbb.h5", }, ], } @@ -441,6 +458,7 @@ def planning_stubs(self, monkeypatch): artifact_store, national_partition, release_bundle, + simulation_runtime, ) existing = { @@ -464,12 +482,36 @@ def exists(self, path): monkeypatch.setattr( artifact_keys, "collect_dataset_identity", - lambda country, year: SimpleNamespace( - digest=f"ds-{year}", - store_path=f"datasets/us/ds-{year}/populace_year_{year}.h5", - filename=f"populace_year_{year}.h5", + lambda country, year, dataset=None: SimpleNamespace( + digest=f"ds-{dataset}-{year}", + store_path=( + f"datasets/us/ds-{dataset}-{year}/{dataset}_year_{year}.h5" + ), + filename=f"{dataset}_year_{year}.h5", ), ) + monkeypatch.setattr( + simulation_runtime, + "bundle_dataset_selection", + lambda country, dataset: SimpleNamespace( + name=dataset, + is_default=dataset == "populace_cps", + ), + ) + monkeypatch.setattr( + simulation_runtime, + "dataset_data_folder", + lambda country, selection: ( + "/opt/policyengine/data" + if selection.is_default + else "/tmp/policyengine-alternate-data/local" + ), + ) + monkeypatch.setattr( + simulation_runtime, + "resolve_data_folder", + lambda: "/opt/policyengine/data", + ) def fake_cohort_identity(year, group): tag = f"{year}-{len(group)}" @@ -490,6 +532,10 @@ def fake_cohort_identity(year, group): data_version="1.2.3", data_artifact_revision="rev-abc", default_dataset="populace_cps", + regional_dataset_defaults={ + "state": "populace_us_2024_acs_local", + "congressional_district": "populace_us_2024_acs_local", + }, ), ) return SimpleNamespace(existing=existing) @@ -497,11 +543,24 @@ def fake_cohort_identity(year, group): def test_plan_enumerates_years_and_cohorts_with_presence(self, planning_stubs): plan = precompute.plan_artifacts_impl("bucket-x") - assert [entry.year for entry in plan.datasets] == [2026, 2027, 2025] - assert [entry.exists for entry in plan.datasets] == [True, False, False] - assert plan.datasets[0].digest == "ds-2026" - assert plan.datasets[0].path == "datasets/us/ds-2026/populace_year_2026.h5" - assert plan.datasets[0].filename == "populace_year_2026.h5" + assert [entry.year for entry in plan.datasets] == [ + 2026, + 2026, + 2027, + 2027, + 2025, + 2025, + ] + assert [entry.dataset for entry in plan.datasets] == [ + "populace_cps", + "populace_us_2024_acs_local", + ] * 3 + assert all(not entry.exists for entry in plan.datasets) + assert plan.datasets[0].digest == "ds-populace_cps-2026" + assert plan.datasets[0].filename == "populace_cps_year_2026.h5" + assert plan.datasets[1].runtime_destination.startswith( + "/tmp/policyengine-alternate-data/local/" + ) # Year-major, partition order inside each year. assert [entry.simulation_id for entry in plan.baselines] == [ @@ -534,7 +593,14 @@ def test_plan_enumerates_years_and_cohorts_with_presence(self, planning_stubs): # An exists-flag inversion would make select_work recompute (or, # worse, skip) the wrong entries — lock the wiring end to end. work = precompute.select_work(plan, force=False) - assert [entry.year for entry in work.datasets] == [2027, 2025] + assert [entry.year for entry in work.datasets] == [ + 2026, + 2026, + 2027, + 2027, + 2025, + 2025, + ] assert len(work.baselines) == 5 def test_empty_partition_fails_loudly(self, planning_stubs, monkeypatch): @@ -603,10 +669,21 @@ def fake_ensure_datasets(*, datasets, years, data_folder): (tmp_path / state.identity.filename).write_bytes(b"h5-bytes") monkeypatch.setattr( - artifact_keys, "collect_dataset_identity", lambda c, y: state.identity + artifact_keys, + "collect_dataset_identity", + lambda c, y, dataset=None: state.identity, ) monkeypatch.setattr(artifact_store, "ArtifactStore", FakeStore) - monkeypatch.setattr(sr, "resolve_data_folder", lambda: str(tmp_path)) + monkeypatch.setattr( + sr, + "bundle_dataset_selection", + lambda country, dataset: SimpleNamespace(name=dataset), + ) + monkeypatch.setattr( + sr, + "dataset_data_folder", + lambda country, selection: str(tmp_path), + ) monkeypatch.setattr( sr, "_country_module", @@ -625,6 +702,8 @@ def _entry(self, path="datasets/us/ds-2026/populace_year_2026.h5"): digest="ds-2026", path=path, filename="populace_year_2026.h5", + dataset="populace_cps", + runtime_destination="/tmp/populace_year_2026.h5", exists=False, ) @@ -767,6 +846,7 @@ def _entry(self, sim_id="bl1-cohort"): digest="bl-d", path=f"baselines/us/bl-d/{sim_id}.h5", simulation_id=sim_id, + runtime_destination=f"/opt/policyengine/data/{sim_id}.h5", exists=False, ) @@ -987,6 +1067,7 @@ def run(self): digest="d", path="baselines/us/d/bl1-verify.h5", simulation_id="bl1-verify", + runtime_destination="/opt/policyengine/data/bl1-verify.h5", exists=True, ) verdict = precompute.verify_determinism_impl("bucket-x", entry) diff --git a/projects/policyengine-simulation-executor/tests/test_precompute_app.py b/projects/policyengine-simulation-executor/tests/test_precompute_app.py index 1ecfdd18b..212a58ca7 100644 --- a/projects/policyengine-simulation-executor/tests/test_precompute_app.py +++ b/projects/policyengine-simulation-executor/tests/test_precompute_app.py @@ -110,6 +110,10 @@ def _sample_plan(): "digest": "d1", "path": "datasets/us/d1/populace_year_2026.h5", "filename": "populace_year_2026.h5", + "dataset": "populace_cps", + "runtime_destination": ( + "/opt/policyengine/data/populace_year_2026.h5" + ), "exists": False, } ], @@ -121,6 +125,7 @@ def _sample_plan(): "digest": "b1", "path": "baselines/us/b1/bl1-aaaa.h5", "simulation_id": "bl1-aaaa", + "runtime_destination": ("/opt/policyengine/data/bl1-aaaa.h5"), "exists": False, } ], 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..06ca14310 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_modal_app.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_modal_app.py @@ -145,6 +145,20 @@ def test_v2_app_name_and_images_are_separate_and_bundle_derived(monkeypatch) -> ) assert static_runtime_files[3] == {"copy": True} + us_fetches = [ + call + for call in module.us_worker_image.calls + if call[0] == "run_function" and call[1] == "fetch_artifacts" + ] + assert len(us_fetches) == 1 + assert us_fetches[0][2]["secrets"] == [module.gcp_secret] + for image in (module.uk_worker_image, module.coordinator_image): + assert not [ + call + for call in image.calls + if call[0] == "run_function" and call[1] == "fetch_artifacts" + ] + def test_v2_image_retains_required_runtime_dependencies() -> None: project = tomllib.loads( @@ -199,6 +213,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 From ae77ef9eb18f787219a07f35c1787d6854e2f1f9 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Tue, 6 Oct 2026 18:45:03 +0400 Subject: [PATCH 04/10] Require policyengine 6.2.2 --- projects/policyengine-simulation-executor/pyproject.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) 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", From 2d0726e585f4892afde3073d06bad4859f432eeb Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Wed, 7 Oct 2026 18:03:08 +0400 Subject: [PATCH 05/10] fix: consume certified regional dataset metadata --- .../stage12_bundle.py | 36 +++++++----- .../tests/test_stage12_manifest.py | 1 + .../tests/stage12_fixtures.py | 7 ++- .../fixtures/identity_stubs.py | 2 +- .../modal/utils/update_version_registry.py | 4 +- .../precompute.py | 2 +- .../release_bundle.py | 56 ++++++++++++++----- .../stage12_bundle/manifest.py | 6 -- .../stage12_bundle/parsing.py | 56 +++++++++++++------ .../tests/test_precompute.py | 3 +- .../tests/test_release_bundle.py | 15 ++--- .../tests/test_stage12_bundle.py | 14 ++--- .../tests/test_update_version_registry.py | 2 +- .../fixtures/gateway_endpoints.py | 5 +- .../endpoints.py | 6 +- .../tests/test_endpoints.py | 2 +- 16 files changed, 140 insertions(+), 77 deletions(-) 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 984284cb1..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,9 +36,7 @@ class Stage12CountryBundle(StrictBundleModel): default_dataset: NonEmptyText default_dataset_uri: NonEmptyText datasets: tuple[Stage12Dataset, ...] - regional_dataset_defaults: dict[NonEmptyText, NonEmptyText] = Field( - default_factory=dict - ) + region_dataset_identities: dict[NonEmptyText, NonEmptyText] class Stage12BundleManifest(StrictBundleModel): @@ -66,22 +64,34 @@ def select_stage12_dataset( """Select the certified dataset declared for a concrete region type.""" normalized_region = region.strip().lower() - dataset_identity = country.default_dataset + region_type = "national" if normalized_region != country.country: if ( country.country == "us" and len(normalized_region) == 2 and normalized_region.isalpha() ): - dataset_identity = country.regional_dataset_defaults.get( - "state", country.default_dataset - ) + region_type = "state" else: - region_type, separator, _ = normalized_region.partition("/") + requested_region_type, separator, _ = normalized_region.partition("/") if separator: - dataset_identity = country.regional_dataset_defaults.get( - region_type, country.default_dataset - ) - return next( - dataset for dataset in country.datasets if dataset.identity == dataset_identity + 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..05a7b300e 100644 --- a/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py +++ b/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py @@ -69,6 +69,7 @@ def _resolved_bundle(version: str = "5.2.0") -> ResolvedStage12Bundle: repo_type="dataset", ), ), + region_dataset_identities={"national": dataset}, ) ) bundle = Stage12BundleManifest( diff --git a/projects/policyengine-simulation-entry/tests/stage12_fixtures.py b/projects/policyengine-simulation-entry/tests/stage12_fixtures.py index bc7035feb..2e43715c1 100644 --- a/projects/policyengine-simulation-entry/tests/stage12_fixtures.py +++ b/projects/policyengine-simulation-entry/tests/stage12_fixtures.py @@ -45,7 +45,7 @@ def worker() -> V2WorkerVersion: repo_type="dataset", ) ] - regional_dataset_defaults = {} + region_dataset_identities = {"national": dataset} if country == "us": local_dataset = "populace_us_2024_acs_local" local_revision = f"{local_dataset}-revision" @@ -58,7 +58,8 @@ def worker() -> V2WorkerVersion: repo_type="dataset", ) ) - regional_dataset_defaults = { + region_dataset_identities = { + "national": dataset, "state": local_dataset, "congressional_district": local_dataset, } @@ -75,7 +76,7 @@ def worker() -> V2WorkerVersion: default_dataset=dataset, default_dataset_uri=f"hf://policyengine/data/{dataset}.h5@{revision}", datasets=tuple(datasets), - regional_dataset_defaults=regional_dataset_defaults, + region_dataset_identities=region_dataset_identities, ) ) country_workers.append( diff --git a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py index ba00f3acb..b10620edb 100644 --- a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py +++ b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py @@ -47,7 +47,7 @@ def install_identity_stubs(monkeypatch): dataset_data_versions={"populace_cps": "1.2.3"}, dataset_revisions={"populace_cps": "rev-abc"}, dataset_sha256s={"populace_cps": "feedbead" * 8}, - regional_dataset_defaults={}, + region_dataset_identities={"national": "populace_cps"}, ) receipt_entry = { "country": "us", 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 97c823f7a..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 @@ -34,7 +34,7 @@ class CountryBundleMetadata(TypedDict): dataset_data_versions: dict[str, str] dataset_revisions: dict[str, str] dataset_sha256s: dict[str, str] - regional_dataset_defaults: dict[str, str] + region_dataset_identities: dict[str, str] class BundleManifestMetadata(TypedDict): @@ -104,7 +104,7 @@ def _country_bundle_metadata(country: str) -> CountryBundleMetadata: "dataset_data_versions": dict(bundle.dataset_data_versions), "dataset_revisions": dict(bundle.dataset_revisions), "dataset_sha256s": dict(bundle.dataset_sha256s), - "regional_dataset_defaults": dict(bundle.regional_dataset_defaults), + "region_dataset_identities": dict(bundle.region_dataset_identities), } 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 b34026339..faa446c69 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute.py @@ -204,7 +204,7 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: dataset_names = [ bundle.default_dataset, *sorted( - set(bundle.regional_dataset_defaults.values()).difference( + set(bundle.region_dataset_identities.values()).difference( {bundle.default_dataset} ) ), 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 45adfd9f4..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 @@ -43,7 +43,7 @@ class CountryReleaseBundle: dataset_data_versions: Mapping[str, str] dataset_revisions: Mapping[str, str] dataset_sha256s: Mapping[str, str] - regional_dataset_defaults: Mapping[str, str] + region_dataset_identities: Mapping[str, str] def _normalise_country(country: str) -> str: @@ -196,14 +196,42 @@ def _dataset_provenance_from_release( return revisions, sha256s -def _regional_dataset_defaults(bundle: Mapping, country: str) -> dict[str, str]: - all_defaults = _mapping(bundle.get("regional_dataset_defaults")) - country_defaults = _mapping(all_defaults.get(country)) - return { - str(region_type): dataset - for region_type, dataset in country_defaults.items() - if isinstance(dataset, str) and dataset +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: @@ -297,7 +325,7 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: 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) - regional_dataset_defaults: dict[str, str] = {} + region_dataset_identities = _region_dataset_identities_from_manifest(manifest) data_package: Mapping = {} data_release: Mapping = {} if bundle_metadata is not None: @@ -340,7 +368,9 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: ) dataset_revisions.update(release_dataset_revisions) dataset_sha256s.update(release_dataset_sha256s) - regional_dataset_defaults = _regional_dataset_defaults(bundle_manifest, country) + 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, @@ -361,12 +391,12 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: dataset_data_versions = dict(dataset_revisions) dataset_data_versions[default_dataset] = str(data_version) - unknown_regional_datasets = set(regional_dataset_defaults.values()).difference( + unknown_regional_datasets = set(region_dataset_identities.values()).difference( dataset_uris ) if unknown_regional_datasets: raise ValueError( - "PolicyEngine.py bundle regional defaults reference unknown datasets: " + "PolicyEngine.py bundle region mappings reference unknown datasets: " + ", ".join(sorted(unknown_regional_datasets)) ) @@ -386,7 +416,7 @@ def get_country_release_bundle(country: str) -> CountryReleaseBundle: dataset_data_versions=dataset_data_versions, dataset_revisions=dataset_revisions, dataset_sha256s=dataset_sha256s, - regional_dataset_defaults=regional_dataset_defaults, + region_dataset_identities=region_dataset_identities, ) 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 8e18f2b33..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,11 +58,6 @@ def _normalize_stage12_bundle( ) countries = require_mapping(bundle.get("countries"), "countries") data_releases = require_mapping(bundle.get("data_releases"), "data_releases") - regional_dataset_defaults = require_mapping( - bundle.get("regional_dataset_defaults", {}), - "regional_dataset_defaults", - ) - normalized_countries = tuple( normalize_country_bundle( country=country, @@ -70,7 +65,6 @@ def _normalize_stage12_bundle( countries=countries, packages=packages, data_releases=data_releases, - regional_dataset_defaults=regional_dataset_defaults, ) for country in supported_countries ) 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 336fd6cf4..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 @@ -68,7 +68,6 @@ def normalize_country_bundle( countries: Mapping[str, object], packages: Mapping[str, object], data_releases: Mapping[str, object], - regional_dataset_defaults: Mapping[str, object], ) -> Stage12CountryBundle: """Validate and normalize one country's model and dataset release.""" @@ -205,27 +204,52 @@ def normalize_country_bundle( f"PolicyEngine.py bundle default dataset URI is inconsistent for {country!r}" ) - country_regional_defaults = require_mapping( - regional_dataset_defaults.get(country, {}), - f"regional_dataset_defaults.{country}", + region_dataset_templates = require_mapping( + release.get("region_datasets"), + f"data_releases.{country}.region_datasets", ) - normalized_regional_defaults: dict[str, str] = {} - for region_type, dataset_identity in country_regional_defaults.items(): + 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"regional_dataset_defaults.{country}.", + f"data_releases.{country}.region_datasets.", + ) + template = require_mapping( + raw_template, + f"data_releases.{country}.region_datasets.{normalized_region_type}", ) - normalized_dataset_identity = require_text( - dataset_identity, - f"regional_dataset_defaults.{country}.{normalized_region_type}", + path_template = require_text( + template.get("path_template"), + "data_releases." + f"{country}.region_datasets.{normalized_region_type}.path_template", ) - if normalized_dataset_identity not in datasets_by_identity: + matching_datasets = [ + identity + for identity, path in dataset_paths.items() + if path == path_template + ] + if len(matching_datasets) != 1: raise Stage12BundleError( - "PolicyEngine.py bundle regional dataset default " - f"{country}.{normalized_region_type} names an absent dataset" + "PolicyEngine.py bundle region dataset " + f"{country}.{normalized_region_type} must identify exactly one dataset" ) - normalized_regional_defaults[normalized_region_type] = ( - normalized_dataset_identity + 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( @@ -267,5 +291,5 @@ def normalize_country_bundle( default_dataset=default_dataset, default_dataset_uri=default_dataset_uri, datasets=tuple(datasets), - regional_dataset_defaults=normalized_regional_defaults, + region_dataset_identities=normalized_region_dataset_identities, ) diff --git a/projects/policyengine-simulation-executor/tests/test_precompute.py b/projects/policyengine-simulation-executor/tests/test_precompute.py index 4d6ba9220..92dead0af 100644 --- a/projects/policyengine-simulation-executor/tests/test_precompute.py +++ b/projects/policyengine-simulation-executor/tests/test_precompute.py @@ -532,7 +532,8 @@ def fake_cohort_identity(year, group): data_version="1.2.3", data_artifact_revision="rev-abc", default_dataset="populace_cps", - regional_dataset_defaults={ + region_dataset_identities={ + "national": "populace_cps", "state": "populace_us_2024_acs_local", "congressional_district": "populace_us_2024_acs_local", }, diff --git a/projects/policyengine-simulation-executor/tests/test_release_bundle.py b/projects/policyengine-simulation-executor/tests/test_release_bundle.py index e92db0a52..ff6b0985a 100644 --- a/projects/policyengine-simulation-executor/tests/test_release_bundle.py +++ b/projects/policyengine-simulation-executor/tests/test_release_bundle.py @@ -111,12 +111,6 @@ def test_country_release_bundle_exposes_model_and_data_versions(): def test_country_release_bundle_exposes_regional_dataset_provenance(monkeypatch): current = deepcopy(get_current_bundle()) local_dataset = "populace_us_2024_acs_local" - current["regional_dataset_defaults"] = { - "us": { - "state": local_dataset, - "congressional_district": local_dataset, - } - } current["data_releases"]["us"]["datasets"][local_dataset] = { "path": f"{local_dataset}.h5", "repo_id": "policyengine/populace-us", @@ -124,6 +118,12 @@ def test_country_release_bundle_exposes_regional_dataset_provenance(monkeypatch) "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, @@ -134,7 +134,8 @@ def test_country_release_bundle_exposes_regional_dataset_provenance(monkeypatch) bundle = get_country_release_bundle("us") - assert bundle.regional_dataset_defaults == { + assert bundle.region_dataset_identities == { + "national": "populace_us_2024", "state": local_dataset, "congressional_district": local_dataset, } diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py b/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py index 9ab65ee31..beb6f0f44 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_bundle.py @@ -59,12 +59,12 @@ def test_normalized_bundle_selects_regional_dataset_from_packaged_metadata() -> "revision": "acs-local-release", "sha256": "b" * 64, } - raw["regional_dataset_defaults"] = { - "us": { - "state": local_identity, - "congressional_district": local_identity, + 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") @@ -101,8 +101,8 @@ def test_normalized_bundle_selects_regional_dataset_from_packaged_metadata() -> lambda bundle: bundle["data_releases"]["us"]["certified_data_artifact"].update( {"sha256": "0" * 64} ), - lambda bundle: bundle.update( - {"regional_dataset_defaults": {"us": {"state": "absent"}}} + lambda bundle: bundle["data_releases"]["us"]["region_datasets"].update( + {"state": {"path_template": "absent.h5"}} ), ], ids=[ 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 9fdf59a99..80f9670ac 100644 --- a/projects/policyengine-simulation-executor/tests/test_update_version_registry.py +++ b/projects/policyengine-simulation-executor/tests/test_update_version_registry.py @@ -89,7 +89,7 @@ def fake_country_bundle_metadata( "dataset_data_versions": {"default": "test-data-version"}, "dataset_revisions": {"default": "test-revision"}, "dataset_sha256s": {"default": "a" * 64}, - "regional_dataset_defaults": {}, + "region_dataset_identities": {"national": "default"}, } monkeypatch.setattr( diff --git a/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py b/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py index 6369254a6..12d93222d 100644 --- a/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py +++ b/projects/policyengine-simulation-gateway/fixtures/gateway_endpoints.py @@ -47,7 +47,8 @@ "populace_us_2024_acs_local": "b" * 64, "calibration_diagnostics": "c" * 64, }, - "regional_dataset_defaults": { + "region_dataset_identities": { + "national": "populace_us_2024", "state": "populace_us_2024_acs_local", "congressional_district": "populace_us_2024_acs_local", }, @@ -75,7 +76,7 @@ "populace_uk_2023": "d" * 64, "local_authority_weights": "e" * 64, }, - "regional_dataset_defaults": {}, + "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 56d6aa138..f59b68ee2 100644 --- a/projects/policyengine-simulation-gateway/src/policyengine_simulation_gateway/endpoints.py +++ b/projects/policyengine-simulation-gateway/src/policyengine_simulation_gateway/endpoints.py @@ -176,9 +176,9 @@ def _resolve_dataset_name_from_app_bundle( region=region, region_group=region_group, ) - regional_defaults = country_bundle.get("regional_dataset_defaults") - if region_type is not None and isinstance(regional_defaults, dict): - regional_dataset = regional_defaults.get(region_type) + 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 diff --git a/projects/policyengine-simulation-gateway/tests/test_endpoints.py b/projects/policyengine-simulation-gateway/tests/test_endpoints.py index ca61c53df..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( From 59d405ed2a07232753e37fed35dafdca3576cdac Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Thu, 8 Oct 2026 18:24:05 +0400 Subject: [PATCH 06/10] Keep artifact precompute limited to national datasets --- .../workflows/simulation-deploy.reusable.yml | 11 -- .../fixtures/identity_stubs.py | 1 - .../src/modal/_image_setup.py | 9 +- .../src/modal/app.py | 22 +++- .../src/modal/artifact_manifest.py | 31 ----- .../src/modal/v2_app.py | 39 ++---- .../artifact_keys.py | 47 +++---- .../precompute.py | 76 +++--------- .../precompute_models.py | 18 +-- .../tests/test_image_setup_fetch.py | 22 ---- .../tests/test_modal_scripts.py | 19 ++- .../tests/test_precompute.py | 117 ++++-------------- .../tests/test_precompute_app.py | 5 - .../tests/test_stage12_modal_app.py | 16 +-- 14 files changed, 116 insertions(+), 317 deletions(-) delete mode 100644 projects/policyengine-simulation-executor/src/modal/artifact_manifest.py diff --git a/.github/workflows/simulation-deploy.reusable.yml b/.github/workflows/simulation-deploy.reusable.yml index 46e5b2827..ac6d6e40a 100644 --- a/.github/workflows/simulation-deploy.reusable.yml +++ b/.github/workflows/simulation-deploy.reusable.yml @@ -561,15 +561,6 @@ jobs: name="$(uv run python -c 'from policyengine_simulation_contract.stage12_manifest import v2_application_name; from policyengine_simulation_executor.stage12_bundle import load_stage12_bundle; print(v2_application_name(load_stage12_bundle().bundle.policyengine_version))')" echo "name=${name}" >> "$GITHUB_OUTPUT" - - name: Run Stage 12 artifact precompute - id: precompute - working-directory: projects/policyengine-simulation-executor - env: - MODAL_TOKEN_ID: ${{ secrets.MODAL_TOKEN_ID }} - MODAL_TOKEN_SECRET: ${{ secrets.MODAL_TOKEN_SECRET }} - POLICYENGINE_ARTIFACT_BUCKET: ${{ vars.POLICYENGINE_ARTIFACT_BUCKET }} - run: ../../.github/scripts/modal-precompute.sh "${{ inputs.modal_environment }}" "${{ inputs.force_recompute }}" - - name: Synchronize the separate Stage 12 runtime secret working-directory: projects/policyengine-simulation-executor env: @@ -591,8 +582,6 @@ jobs: STAGE12_EXPECT_US_VERSION: ${{ needs.prepare.outputs.us_version }} STAGE12_EXPECT_UK_VERSION: ${{ needs.prepare.outputs.uk_version }} STAGE12_ARTIFACT_BUCKET: ${{ vars.STAGE12_ARTIFACT_BUCKET }} - POLICYENGINE_MANIFEST_DIGEST: ${{ steps.precompute.outputs.manifest_digest }} - POLICYENGINE_ARTIFACT_BUCKET: ${{ vars.POLICYENGINE_ARTIFACT_BUCKET }} OBSERVABILITY_SERVICE_NAMESPACE: ${{ vars.OBSERVABILITY_SERVICE_NAMESPACE }} OBSERVABILITY_TRACE_PROJECT_ID: ${{ vars.OBSERVABILITY_TRACE_PROJECT_ID }} OBSERVABILITY_LOGGING_PROJECT_ID: ${{ vars.OBSERVABILITY_LOGGING_PROJECT_ID }} diff --git a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py index b10620edb..a028b1a93 100644 --- a/projects/policyengine-simulation-executor/fixtures/identity_stubs.py +++ b/projects/policyengine-simulation-executor/fixtures/identity_stubs.py @@ -47,7 +47,6 @@ def install_identity_stubs(monkeypatch): dataset_data_versions={"populace_cps": "1.2.3"}, dataset_revisions={"populace_cps": "rev-abc"}, dataset_sha256s={"populace_cps": "feedbead" * 8}, - region_dataset_identities={"national": "populace_cps"}, ) receipt_entry = { "country": "us", diff --git a/projects/policyengine-simulation-executor/src/modal/_image_setup.py b/projects/policyengine-simulation-executor/src/modal/_image_setup.py index d80e412b1..f8f7183c3 100644 --- a/projects/policyengine-simulation-executor/src/modal/_image_setup.py +++ b/projects/policyengine-simulation-executor/src/modal/_image_setup.py @@ -142,17 +142,14 @@ def fetch_artifacts(bucket: str, manifest, *, client=None): # survive into the image, and the layer starts empty of these # files anyway. for artifact in artifacts: - target = Path(artifact.get("destination") or data_folder / artifact["filename"]) - target.parent.mkdir(parents=True, exist_ok=True) + target = data_folder / artifact["filename"] bucket_handle.blob(artifact["path"]).download_to_filename(str(target)) logger.info("Fetched %s (%.1f MB)", target, target.stat().st_size / 1e6) absent = [ - str(Path(artifact.get("destination") or data_folder / artifact["filename"])) + str(data_folder / artifact["filename"]) for artifact in artifacts - if not Path( - artifact.get("destination") or data_folder / artifact["filename"] - ).exists() + if not (data_folder / artifact["filename"]).exists() ] if absent: raise RuntimeError(f"Artifact fetch did not produce expected files: {absent}") diff --git a/projects/policyengine-simulation-executor/src/modal/app.py b/projects/policyengine-simulation-executor/src/modal/app.py index 2dc9817c0..7e72dbb7a 100644 --- a/projects/policyengine-simulation-executor/src/modal/app.py +++ b/projects/policyengine-simulation-executor/src/modal/app.py @@ -34,7 +34,6 @@ get_bundled_package_version, ) from src.modal._image_setup import fetch_artifacts, snapshot_models -from src.modal.artifact_manifest import deploy_time_artifact_inputs from src.modal.bundle_data import bundle_data_install_command from src.modal.logging_redaction import redact_params_for_logging from src.modal.static_runtime_files import add_static_runtime_files @@ -142,7 +141,26 @@ def _deploy_time_artifact_inputs() -> tuple[str, dict | None]: model by design (the layer is import-restricted), so the dict exists only at that serialization boundary. """ - return deploy_time_artifact_inputs() + if not modal.is_local(): + return "", None + digest = os.environ.get("POLICYENGINE_MANIFEST_DIGEST") + bucket = os.environ.get("POLICYENGINE_ARTIFACT_BUCKET", "") + if not digest: + return bucket, None + from policyengine_simulation_executor.artifact_store import ArtifactStore + from policyengine_simulation_executor.precompute_models import ArtifactManifest + + store = ArtifactStore(bucket or None) + payload = store.read_manifest(digest) + if payload is None: + raise RuntimeError( + f"Artifact manifest {digest} is not in the store: the deploy " + "must consume a digest published by the precompute run." + ) + # Validate on the runner so a corrupt or foreign manifest fails the + # deploy here, not mid-image-build. + manifest = ArtifactManifest.model_validate(payload) + return store.bucket_name, manifest.canonical_payload() _ARTIFACT_BUCKET, _DEPLOY_MANIFEST = _deploy_time_artifact_inputs() diff --git a/projects/policyengine-simulation-executor/src/modal/artifact_manifest.py b/projects/policyengine-simulation-executor/src/modal/artifact_manifest.py deleted file mode 100644 index 7c0119789..000000000 --- a/projects/policyengine-simulation-executor/src/modal/artifact_manifest.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Resolve a content-addressed precompute manifest during Modal deployment.""" - -from __future__ import annotations - -import os - - -def deploy_time_artifact_inputs() -> tuple[str, dict | None]: - """Return the artifact bucket and validated manifest on the deploy runner.""" - - import modal - - if not modal.is_local(): - return "", None - digest = os.environ.get("POLICYENGINE_MANIFEST_DIGEST") - bucket = os.environ.get("POLICYENGINE_ARTIFACT_BUCKET", "") - if not digest: - return bucket, None - - from policyengine_simulation_executor.artifact_store import ArtifactStore - from policyengine_simulation_executor.precompute_models import ArtifactManifest - - store = ArtifactStore(bucket or None) - payload = store.read_manifest(digest) - if payload is None: - raise RuntimeError( - f"Artifact manifest {digest} is not in the store: the deploy " - "must consume a digest published by the precompute run." - ) - manifest = ArtifactManifest.model_validate(payload) - return store.bucket_name, manifest.canonical_payload() diff --git a/projects/policyengine-simulation-executor/src/modal/v2_app.py b/projects/policyengine-simulation-executor/src/modal/v2_app.py index 2fce9b61c..60cbee22f 100644 --- a/projects/policyengine-simulation-executor/src/modal/v2_app.py +++ b/projects/policyengine-simulation-executor/src/modal/v2_app.py @@ -37,12 +37,10 @@ assertion_values, load_stage12_bundle, ) -from src.modal._image_setup import fetch_artifacts -from src.modal.artifact_manifest import deploy_time_artifact_inputs from src.modal.bundle_data import bundle_data_install_command from src.modal.static_runtime_files import add_static_runtime_files -STAGE12_DATA_DIR = "/opt/policyengine/data" +STAGE12_DATA_DIR = "/opt/policyengine/stage12-data" _UV_PROJECT_DIR = str(Path(__file__).resolve().parents[2]) if modal.is_local() else "." RESOLVED_BUNDLE = load_stage12_bundle() BUNDLE_VALUES = assertion_values(RESOLVED_BUNDLE.bundle) @@ -84,7 +82,6 @@ def _external_assertions(environment: dict[str, str]) -> dict[str, str]: hf_secret, comparison_runtime_secret, ] -_ARTIFACT_BUCKET, _DEPLOY_MANIFEST = deploy_time_artifact_inputs() def _country_bundle(country: CountryId): @@ -93,11 +90,7 @@ def _country_bundle(country: CountryId): ) -def build_v2_image( - countries: tuple[CountryId, ...], - *, - include_us_artifacts: bool = False, -) -> modal.Image: +def build_v2_image(countries: tuple[CountryId, ...]) -> modal.Image: country_values = { country: _country_bundle(country).model_dump(mode="json") for country in countries @@ -121,9 +114,6 @@ def build_v2_image( { **modal_image_environment(), "POLICYENGINE_DATA_FOLDER": STAGE12_DATA_DIR, - "POLICYENGINE_BUNDLE_RECEIPT": ( - f"{STAGE12_DATA_DIR}/.policyengine-bundle-receipt.json" - ), "STAGE12_BUNDLE_MANIFEST_SHA256": ( RESOLVED_BUNDLE.bundle_manifest_sha256 ), @@ -134,28 +124,19 @@ def build_v2_image( ), } ) - ) - if include_us_artifacts: - image = image.run_function( - fetch_artifacts, - args=(_ARTIFACT_BUCKET, _DEPLOY_MANIFEST), - secrets=[gcp_secret], - cpu=2.0, - memory=4096, - timeout=900, + .add_local_python_source( + "src.modal", + "policyengine_simulation_executor", + "policyengine_simulation_observability", + "policyengine_simulation_contract", + "policyengine_stage12_persistence", + copy=True, ) - image = image.add_local_python_source( - "src.modal", - "policyengine_simulation_executor", - "policyengine_simulation_observability", - "policyengine_simulation_contract", - "policyengine_stage12_persistence", - copy=True, ) return add_static_runtime_files(image, uv_project_dir=_UV_PROJECT_DIR) -us_worker_image = build_v2_image(("us",), include_us_artifacts=True) +us_worker_image = build_v2_image(("us",)) uk_worker_image = build_v2_image(("uk",)) coordinator_image = build_v2_image(("us", "uk")) diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py index 077255930..5cdcffba9 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/artifact_keys.py @@ -302,12 +302,8 @@ def _certification_fingerprint(country: str) -> Optional[str]: return fingerprint if isinstance(fingerprint, str) and fingerprint else None -def collect_dataset_identity( - country: str, - year: int, - dataset: str | None = None, -) -> DatasetArtifactIdentity: - """Identity of one certified single-year dataset for this installed bundle. +def collect_dataset_identity(country: str, year: int) -> DatasetArtifactIdentity: + """Identity of the default single-year dataset for this installed bundle. Reads the same sources the runtime trusts: the release bundle (versions), the bundle receipt (content sha), and the manifest certification @@ -315,47 +311,32 @@ def collect_dataset_identity( runtime lookup uses, so filename coupling is inherited, not re-derived. """ - from policyengine.provenance.manifest import dataset_logical_name + from policyengine.provenance.manifest import ( + dataset_logical_name, + resolve_dataset_reference, + ) from policyengine_simulation_executor.release_bundle import ( get_country_release_bundle, ) bundle = get_country_release_bundle(country) - selected_dataset = dataset or bundle.default_dataset - try: - selected_uri = bundle.dataset_uris[selected_dataset] - except KeyError as exc: - raise ValueError( - f"Unknown certified dataset {selected_dataset!r} for {bundle.country!r}" - ) from exc - stem = dataset_logical_name(selected_uri) - data_version = bundle.dataset_data_versions.get(selected_dataset) - artifact_revision = bundle.dataset_revisions.get(selected_dataset) - if not data_version or not artifact_revision: - raise ValueError( - f"Certified dataset {selected_dataset!r} has incomplete provenance" - ) - is_default = selected_dataset == bundle.default_dataset + stem = dataset_logical_name( + resolve_dataset_reference(bundle.country, bundle.default_dataset) + ) from policyengine_simulation_executor.spm import normalize_runtime_spm selection = normalize_runtime_spm({"country": country, "time_period": year}) return DatasetArtifactIdentity( spm=selection, country=bundle.country, - dataset=selected_dataset, + dataset=bundle.default_dataset, stem=stem, year=int(year), - data_version=data_version, - data_artifact_revision=artifact_revision, - source_sha256=( - _receipt_source_sha256(bundle.country, bundle.data_version) - if is_default - else bundle.dataset_sha256s.get(selected_dataset) - ), - data_build_fingerprint=( - _certification_fingerprint(bundle.country) if is_default else None - ), + data_version=bundle.data_version, + data_artifact_revision=bundle.data_artifact_revision, + source_sha256=_receipt_source_sha256(bundle.country, bundle.data_version), + data_build_fingerprint=_certification_fingerprint(bundle.country), model_version=bundle.model_version, policyengine_version=bundle.policyengine_version, ) 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 faa446c69..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,12 +56,14 @@ 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. PRECOMPUTE_YEARS = [2026, 2027, 2025] -MANIFEST_SCHEMA = "mf2" +MANIFEST_SCHEMA = "mf1" # The stdout contract the CI deploy job parses; keep format changes in # lockstep with the workflow side (covered by a regression test). @@ -151,7 +153,6 @@ def build_manifest(plan: PrecomputePlan) -> ArtifactManifest: filename=entry.filename, year=entry.year, digest=entry.digest, - destination=entry.runtime_destination, ) for entry in plan.datasets ] + [ @@ -161,7 +162,6 @@ def build_manifest(plan: PrecomputePlan) -> ArtifactManifest: filename=entry.path.rsplit("/", maxsplit=1)[-1], year=entry.year, digest=entry.digest, - destination=entry.runtime_destination, ) for entry in plan.baselines ] @@ -175,8 +175,6 @@ def build_manifest(plan: PrecomputePlan) -> ArtifactManifest: def plan_artifacts_impl(bucket: str) -> PrecomputePlan: """Compute every expected artifact identity and its store presence.""" - from pathlib import Path - from policyengine_simulation_executor.artifact_keys import ( collect_dataset_identity, ) @@ -187,11 +185,6 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: from policyengine_simulation_executor.release_bundle import ( get_country_release_bundle, ) - from policyengine_simulation_executor.simulation_runtime import ( - bundle_dataset_selection, - dataset_data_folder, - resolve_data_folder, - ) store = ArtifactStore(bucket) groups = national_region_groups(PRECOMPUTE_COUNTRY) @@ -200,40 +193,17 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: datasets: list[DatasetPlanEntry] = [] baselines: list[BaselinePlanEntry] = [] - bundle = get_country_release_bundle(PRECOMPUTE_COUNTRY) - dataset_names = [ - bundle.default_dataset, - *sorted( - set(bundle.region_dataset_identities.values()).difference( - {bundle.default_dataset} - ) - ), - ] for year in PRECOMPUTE_YEARS: - for dataset_name in dataset_names: - dataset_identity = collect_dataset_identity( - PRECOMPUTE_COUNTRY, - year, - dataset_name, - ) - selection = bundle_dataset_selection( - PRECOMPUTE_COUNTRY, - dataset_name, - ) - datasets.append( - DatasetPlanEntry( - year=year, - digest=dataset_identity.digest, - path=dataset_identity.store_path, - filename=dataset_identity.filename, - dataset=dataset_name, - runtime_destination=str( - Path(dataset_data_folder(PRECOMPUTE_COUNTRY, selection)) - / dataset_identity.filename - ), - exists=store.exists(dataset_identity.store_path), - ) + dataset_identity = collect_dataset_identity(PRECOMPUTE_COUNTRY, year) + datasets.append( + DatasetPlanEntry( + year=year, + digest=dataset_identity.digest, + path=dataset_identity.store_path, + filename=dataset_identity.filename, + exists=store.exists(dataset_identity.store_path), ) + ) for group in groups: identity = cohort_identity(year, group) baselines.append( @@ -244,14 +214,11 @@ def plan_artifacts_impl(bucket: str) -> PrecomputePlan: digest=identity.digest, path=identity.store_path, simulation_id=identity.simulation_id, - runtime_destination=str( - Path(resolve_data_folder()) - / identity.store_path.rsplit("/", maxsplit=1)[-1] - ), exists=store.exists(identity.store_path), ) ) + bundle = get_country_release_bundle(PRECOMPUTE_COUNTRY) receipt = BundleVersionIdentity( policyengine_version=bundle.policyengine_version, model_version=bundle.model_version, @@ -270,17 +237,15 @@ def build_dataset_impl(bucket: str, expected: DatasetPlanEntry) -> DatasetBuildR collect_dataset_identity, ) from policyengine_simulation_executor.artifact_store import ArtifactStore + from policyengine_simulation_executor.release_bundle import ( + get_country_release_bundle, + ) from policyengine_simulation_executor.simulation_runtime import ( _country_module, - bundle_dataset_selection, - dataset_data_folder, + resolve_data_folder, ) - identity = collect_dataset_identity( - PRECOMPUTE_COUNTRY, - expected.year, - expected.dataset, - ) + identity = collect_dataset_identity(PRECOMPUTE_COUNTRY, expected.year) if identity.store_path != expected.path: raise RuntimeError( "Planned and in-container dataset identities disagree " @@ -288,12 +253,11 @@ def build_dataset_impl(bucket: str, expected: DatasetPlanEntry) -> DatasetBuildR "under a mismatched key." ) - selection = bundle_dataset_selection(PRECOMPUTE_COUNTRY, expected.dataset) - data_folder = dataset_data_folder(PRECOMPUTE_COUNTRY, selection) + data_folder = resolve_data_folder() country_module = _country_module(PRECOMPUTE_COUNTRY) started = time.monotonic() country_module.ensure_datasets( - datasets=[expected.dataset], + datasets=[get_country_release_bundle(PRECOMPUTE_COUNTRY).default_dataset], years=[expected.year], data_folder=data_folder, ) diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py index 0cfee633b..fd5332fd2 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/precompute_models.py @@ -22,7 +22,7 @@ from typing import Any, Literal, Optional, Protocol -from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic import BaseModel, ConfigDict, Field # Mirrors baseline_artifacts.OUTCOME_*; the artifact outcome a worker # observed when ensuring a baseline. @@ -56,8 +56,6 @@ class DatasetPlanEntry(_StrictModel): digest: str path: str filename: str - dataset: str - runtime_destination: str exists: bool @@ -70,7 +68,6 @@ class BaselinePlanEntry(_StrictModel): digest: str path: str simulation_id: str - runtime_destination: str exists: bool @@ -97,7 +94,6 @@ class ManifestArtifact(_StrictModel): filename: str year: int digest: str - destination: str | None = None class ArtifactManifest(_StrictModel): @@ -107,22 +103,14 @@ class ArtifactManifest(_StrictModel): # "schema" is the wire key; the Python name avoids shadowing # BaseModel's deprecated .schema attribute. - manifest_schema: Literal["mf1", "mf2"] = Field(alias="schema") + manifest_schema: str = Field(alias="schema") country: str receipt: BundleVersionIdentity artifacts: list[ManifestArtifact] - @model_validator(mode="after") - def require_mf2_destinations(self) -> "ArtifactManifest": - if self.manifest_schema == "mf2" and any( - artifact.destination is None for artifact in self.artifacts - ): - raise ValueError("mf2 artifacts must declare a runtime destination") - return self - def canonical_payload(self) -> dict[str, Any]: """The exact dict shape that gets digested and stored.""" - return self.model_dump(by_alias=True, exclude_none=True) + return self.model_dump(by_alias=True) class DeployedMarker(_StrictModel): diff --git a/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py b/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py index ac0d5ad3a..2f4a71fe5 100644 --- a/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py +++ b/projects/policyengine-simulation-executor/tests/test_image_setup_fetch.py @@ -107,28 +107,6 @@ def test_downloads_every_artifact_under_its_runtime_filename( target = data_folder / artifact["filename"] assert target.read_bytes() == f"bytes-{artifact['digest']}".encode() - def test_mf2_downloads_to_each_declared_runtime_destination( - self, fake_client, data_folder, tmp_path - ): - manifest = _manifest() - manifest["schema"] = "mf2" - destinations = [ - tmp_path / "national" / "populace_year_2026.h5", - tmp_path / "regional" / "bl1-aaaa.h5", - ] - for artifact, destination in zip( - manifest["artifacts"], destinations, strict=True - ): - artifact["destination"] = str(destination) - - validated = ArtifactManifest.model_validate(manifest).canonical_payload() - fetch_artifacts("test-bucket", validated, client=fake_client) - - assert [destination.read_bytes() for destination in destinations] == [ - b"bytes-d1", - b"bytes-b1", - ] - def test_missing_store_object_fails_before_any_download( self, fake_client, data_folder ): diff --git a/projects/policyengine-simulation-executor/tests/test_modal_scripts.py b/projects/policyengine-simulation-executor/tests/test_modal_scripts.py index 74204688f..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 = ( @@ -728,13 +741,13 @@ def test_deploy_workflow_threads_manifest_and_store_credentials_to_deploy(self): assert "GCP_CREDENTIALS_JSON: ${{ secrets.GCP_CREDENTIALS_JSON }}" in ( deploy_step ) - # The direct precompute, direct deploy, Stage 12 precompute, - # Stage 12 deploy, and record-marker steps each carry the bucket var. + # The precompute, deploy, and record-marker steps each carry the + # bucket var. assert ( reusable_workflow.count( "POLICYENGINE_ARTIFACT_BUCKET: ${{ vars.POLICYENGINE_ARTIFACT_BUCKET }}" ) - == 5 + == 3 ) def test_precompute_is_gated_by_the_shared_secret_sync(self): diff --git a/projects/policyengine-simulation-executor/tests/test_precompute.py b/projects/policyengine-simulation-executor/tests/test_precompute.py index 92dead0af..b6c65b219 100644 --- a/projects/policyengine-simulation-executor/tests/test_precompute.py +++ b/projects/policyengine-simulation-executor/tests/test_precompute.py @@ -46,8 +46,6 @@ def _plan() -> PrecomputePlan: digest="d1", path="datasets/us/d1/populace_year_2026.h5", filename="populace_year_2026.h5", - dataset="populace_cps", - runtime_destination="/opt/policyengine/data/populace_year_2026.h5", exists=True, ), DatasetPlanEntry( @@ -55,8 +53,6 @@ def _plan() -> PrecomputePlan: digest="d2", path="datasets/us/d2/populace_year_2027.h5", filename="populace_year_2027.h5", - dataset="populace_cps", - runtime_destination="/opt/policyengine/data/populace_year_2027.h5", exists=False, ), ], @@ -68,7 +64,6 @@ def _plan() -> PrecomputePlan: digest="b1", path="baselines/us/b1/bl1-aaaa.h5", simulation_id="bl1-aaaa", - runtime_destination="/opt/policyengine/data/bl1-aaaa.h5", exists=True, ), BaselinePlanEntry( @@ -78,7 +73,6 @@ def _plan() -> PrecomputePlan: digest="b2", path="baselines/us/b2/bl1-bbbb.h5", simulation_id="bl1-bbbb", - runtime_destination="/opt/policyengine/data/bl1-bbbb.h5", exists=False, ), ], @@ -117,13 +111,6 @@ def test_artifact_outcome_mirrors_runtime_constants(self): ba.OUTCOME_MISS, } - def test_mf2_manifest_requires_runtime_destinations(self): - manifest = precompute.build_manifest(_plan()).canonical_payload() - del manifest["artifacts"][0]["destination"] - - with pytest.raises(ValidationError, match="runtime destination"): - ArtifactManifest.model_validate(manifest) - class TestPlanning: def test_select_work_computes_only_misses(self): @@ -170,7 +157,7 @@ def test_manifest_payload_is_wire_compatible(self): identically to the raw dict shape already published to the store (key names included — note "schema", not "manifest_schema").""" raw = { - "schema": "mf2", + "schema": "mf1", "country": "us", "receipt": { "policyengine_version": "4.22.0", @@ -186,7 +173,6 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "populace_year_2026.h5", "year": 2026, "digest": "d1", - "destination": "/opt/policyengine/data/populace_year_2026.h5", }, { "type": "dataset", @@ -194,7 +180,6 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "populace_year_2027.h5", "year": 2027, "digest": "d2", - "destination": "/opt/policyengine/data/populace_year_2027.h5", }, { "type": "baseline", @@ -202,7 +187,6 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "bl1-aaaa.h5", "year": 2026, "digest": "b1", - "destination": "/opt/policyengine/data/bl1-aaaa.h5", }, { "type": "baseline", @@ -210,7 +194,6 @@ def test_manifest_payload_is_wire_compatible(self): "filename": "bl1-bbbb.h5", "year": 2027, "digest": "b2", - "destination": "/opt/policyengine/data/bl1-bbbb.h5", }, ], } @@ -458,7 +441,6 @@ def planning_stubs(self, monkeypatch): artifact_store, national_partition, release_bundle, - simulation_runtime, ) existing = { @@ -482,36 +464,12 @@ def exists(self, path): monkeypatch.setattr( artifact_keys, "collect_dataset_identity", - lambda country, year, dataset=None: SimpleNamespace( - digest=f"ds-{dataset}-{year}", - store_path=( - f"datasets/us/ds-{dataset}-{year}/{dataset}_year_{year}.h5" - ), - filename=f"{dataset}_year_{year}.h5", - ), - ) - monkeypatch.setattr( - simulation_runtime, - "bundle_dataset_selection", - lambda country, dataset: SimpleNamespace( - name=dataset, - is_default=dataset == "populace_cps", - ), - ) - monkeypatch.setattr( - simulation_runtime, - "dataset_data_folder", - lambda country, selection: ( - "/opt/policyengine/data" - if selection.is_default - else "/tmp/policyengine-alternate-data/local" + lambda country, year: SimpleNamespace( + digest=f"ds-{year}", + store_path=f"datasets/us/ds-{year}/populace_year_{year}.h5", + filename=f"populace_year_{year}.h5", ), ) - monkeypatch.setattr( - simulation_runtime, - "resolve_data_folder", - lambda: "/opt/policyengine/data", - ) def fake_cohort_identity(year, group): tag = f"{year}-{len(group)}" @@ -544,24 +502,11 @@ def fake_cohort_identity(year, group): def test_plan_enumerates_years_and_cohorts_with_presence(self, planning_stubs): plan = precompute.plan_artifacts_impl("bucket-x") - assert [entry.year for entry in plan.datasets] == [ - 2026, - 2026, - 2027, - 2027, - 2025, - 2025, - ] - assert [entry.dataset for entry in plan.datasets] == [ - "populace_cps", - "populace_us_2024_acs_local", - ] * 3 - assert all(not entry.exists for entry in plan.datasets) - assert plan.datasets[0].digest == "ds-populace_cps-2026" - assert plan.datasets[0].filename == "populace_cps_year_2026.h5" - assert plan.datasets[1].runtime_destination.startswith( - "/tmp/policyengine-alternate-data/local/" - ) + assert [entry.year for entry in plan.datasets] == [2026, 2027, 2025] + assert [entry.exists for entry in plan.datasets] == [True, False, False] + assert plan.datasets[0].digest == "ds-2026" + assert plan.datasets[0].path == "datasets/us/ds-2026/populace_year_2026.h5" + assert plan.datasets[0].filename == "populace_year_2026.h5" # Year-major, partition order inside each year. assert [entry.simulation_id for entry in plan.baselines] == [ @@ -594,16 +539,23 @@ def test_plan_enumerates_years_and_cohorts_with_presence(self, planning_stubs): # An exists-flag inversion would make select_work recompute (or, # worse, skip) the wrong entries — lock the wiring end to end. work = precompute.select_work(plan, force=False) - assert [entry.year for entry in work.datasets] == [ - 2026, - 2026, - 2027, - 2027, - 2025, - 2025, - ] + 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 @@ -670,21 +622,10 @@ def fake_ensure_datasets(*, datasets, years, data_folder): (tmp_path / state.identity.filename).write_bytes(b"h5-bytes") monkeypatch.setattr( - artifact_keys, - "collect_dataset_identity", - lambda c, y, dataset=None: state.identity, + artifact_keys, "collect_dataset_identity", lambda c, y: state.identity ) monkeypatch.setattr(artifact_store, "ArtifactStore", FakeStore) - monkeypatch.setattr( - sr, - "bundle_dataset_selection", - lambda country, dataset: SimpleNamespace(name=dataset), - ) - monkeypatch.setattr( - sr, - "dataset_data_folder", - lambda country, selection: str(tmp_path), - ) + monkeypatch.setattr(sr, "resolve_data_folder", lambda: str(tmp_path)) monkeypatch.setattr( sr, "_country_module", @@ -703,8 +644,6 @@ def _entry(self, path="datasets/us/ds-2026/populace_year_2026.h5"): digest="ds-2026", path=path, filename="populace_year_2026.h5", - dataset="populace_cps", - runtime_destination="/tmp/populace_year_2026.h5", exists=False, ) @@ -847,7 +786,6 @@ def _entry(self, sim_id="bl1-cohort"): digest="bl-d", path=f"baselines/us/bl-d/{sim_id}.h5", simulation_id=sim_id, - runtime_destination=f"/opt/policyengine/data/{sim_id}.h5", exists=False, ) @@ -1068,7 +1006,6 @@ def run(self): digest="d", path="baselines/us/d/bl1-verify.h5", simulation_id="bl1-verify", - runtime_destination="/opt/policyengine/data/bl1-verify.h5", exists=True, ) verdict = precompute.verify_determinism_impl("bucket-x", entry) diff --git a/projects/policyengine-simulation-executor/tests/test_precompute_app.py b/projects/policyengine-simulation-executor/tests/test_precompute_app.py index 212a58ca7..1ecfdd18b 100644 --- a/projects/policyengine-simulation-executor/tests/test_precompute_app.py +++ b/projects/policyengine-simulation-executor/tests/test_precompute_app.py @@ -110,10 +110,6 @@ def _sample_plan(): "digest": "d1", "path": "datasets/us/d1/populace_year_2026.h5", "filename": "populace_year_2026.h5", - "dataset": "populace_cps", - "runtime_destination": ( - "/opt/policyengine/data/populace_year_2026.h5" - ), "exists": False, } ], @@ -125,7 +121,6 @@ def _sample_plan(): "digest": "b1", "path": "baselines/us/b1/bl1-aaaa.h5", "simulation_id": "bl1-aaaa", - "runtime_destination": ("/opt/policyengine/data/bl1-aaaa.h5"), "exists": False, } ], 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 06ca14310..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,20 +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} - - us_fetches = [ - call - for call in module.us_worker_image.calls - if call[0] == "run_function" and call[1] == "fetch_artifacts" - ] - assert len(us_fetches) == 1 - assert us_fetches[0][2]["secrets"] == [module.gcp_secret] - for image in (module.uk_worker_image, module.coordinator_image): - assert not [ - call + assert not any( + call[0] == "run_function" and call[1] == "fetch_artifacts" for call in image.calls - if call[0] == "run_function" and call[1] == "fetch_artifacts" - ] + ) def test_v2_image_retains_required_runtime_dependencies() -> None: From a955846bbc9426bd2898c38f4f759e20b79ad04a Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Thu, 8 Oct 2026 18:37:20 +0400 Subject: [PATCH 07/10] Cover subnational datasets in Stage 12 manifest fixtures --- .../tests/test_stage12_manifest.py | 80 ++++++++++++++++--- 1 file changed, 67 insertions(+), 13 deletions(-) diff --git a/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py b/libs/policyengine-simulation-contract/tests/test_stage12_manifest.py index 05a7b300e..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,17 +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", - ), - ), - region_dataset_identities={"national": dataset}, + default_dataset_uri=national_dataset.uri, + datasets=tuple(datasets), + region_dataset_identities=region_dataset_identities, ) ) bundle = Stage12BundleManifest( @@ -132,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()) From 96d21293ff9540bca599a002952969256cce8529 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Fri, 9 Oct 2026 02:05:46 +0400 Subject: [PATCH 08/10] Validate real ACS preparation and Utah calculation in PR CI --- .github/workflows/pr-image-smoke.yml | 7 +- docs/us-regional-dataset-pr-validation.md | 46 ++++++ .../src/modal/smoke_app.py | 21 ++- .../src/modal/us_regional_dataset_check.py | 150 ++++++++++++++++++ .../integration/test_image_smoke_modal.py | 5 +- .../tests/test_smoke_app.py | 65 +++++++- .../tests/test_us_regional_dataset_check.py | 111 +++++++++++++ 7 files changed, 395 insertions(+), 10 deletions(-) create mode 100644 docs/us-regional-dataset-pr-validation.md create mode 100644 projects/policyengine-simulation-executor/src/modal/us_regional_dataset_check.py create mode 100644 projects/policyengine-simulation-executor/tests/test_us_regional_dataset_check.py 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..25c1303d3 --- /dev/null +++ b/docs/us-regional-dataset-pr-validation.md @@ -0,0 +1,46 @@ +# 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. First publish the narrowly scoped temporary compatibility fix +in PolicyEngine/policyengine.py#566, then update the executor's two `.py` pins +and frozen lockfile to that actual published release and rerun the check. 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/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_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_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" + ) From 3a89f4178a445395eb3ffa162510a89e5ce43d55 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Fri, 9 Oct 2026 18:40:06 +0400 Subject: [PATCH 09/10] Target PolicyEngine 6.2.5 pending WIC fix publication --- docs/us-regional-dataset-pr-validation.md | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/docs/us-regional-dataset-pr-validation.md b/docs/us-regional-dataset-pr-validation.md index 25c1303d3..16d11ba15 100644 --- a/docs/us-regional-dataset-pr-validation.md +++ b/docs/us-regional-dataset-pr-validation.md @@ -41,6 +41,16 @@ 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. First publish the narrowly scoped temporary compatibility fix -in PolicyEngine/policyengine.py#566, then update the executor's two `.py` pins -and frozen lockfile to that actual published release and rerun the check. +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 still resolves 6.2.2. Frozen +installs therefore still use 6.2.2, not the new target; this intermediate state +is intentionally incomplete and must not be merged or deployed. From bc470c274387d3283ef636b36a94c29e78c8cad1 Mon Sep 17 00:00:00 2001 From: Anthony Volk <14987227+anth-volk@users.noreply.github.com> Date: Fri, 9 Oct 2026 18:59:26 +0400 Subject: [PATCH 10/10] Document current main lockfile after PR rebase --- docs/us-regional-dataset-pr-validation.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/us-regional-dataset-pr-validation.md b/docs/us-regional-dataset-pr-validation.md index 16d11ba15..08859a8c0 100644 --- a/docs/us-regional-dataset-pr-validation.md +++ b/docs/us-regional-dataset-pr-validation.md @@ -51,6 +51,6 @@ 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 still resolves 6.2.2. Frozen -installs therefore still use 6.2.2, not the new target; this intermediate state +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.