diff --git a/docs/migration/stage-12-modal-worker-foundation.md b/docs/migration/stage-12-modal-worker-foundation.md index ef15f3b94..4e3a47231 100644 --- a/docs/migration/stage-12-modal-worker-foundation.md +++ b/docs/migration/stage-12-modal-worker-foundation.md @@ -165,6 +165,43 @@ identifies an artifact declared by that same bundle. The request model converts the Public API's `_telemetry` field to `telemetry`; the adapter accepts that correlation metadata but does not include it in either simulation input. +## Report output planning + +Before starting either child simulation, the Modal report coordinator resolves +one immutable, strongly typed output plan for the complete Stage 12 report. The +plan records the requested aggregate profile, whether cliff analysis was +requested, whether either policy activates labor-supply responses, and the +required columns for every country-model entity. + +The resolver constructs data-free planning `Simulation` objects and invokes +the output-configuration functions supplied by the pinned PolicyEngine 5.2.0 +bundle. For US reports this includes the budgetary-impact variables; both +countries use PolicyEngine's conditional cliff and labor-supply configuration. +Planning never loads a dataset or executes a simulation. The former +`variables=("*",)` internal marker is not an output expansion mechanism and is +no longer used. + +The coordinator includes the exact same output plan in the baseline and reform +child inputs. Each child adds the plan's additional variables before calling +`Simulation.ensure()`, then verifies that every planned entity and column is +present before writing its Parquet artifact. The artifact descriptor records a +digest of the plan. After both calls complete, the coordinator checks both plan +digests and independently validates both loaded Parquet schemas before it runs +aggregate calculations. Extra columns are permitted; missing planned columns +fail the temporary Stage 12 report without affecting the production result. + +The plan distinguishes calculated variables from columns copied directly from +the source dataset. UK geographic reports require the household columns +`constituency_code_oa` and `la_code_oa`. Workers do not request those columns as +calculated PolicyEngine variables, but they must be present in each materialized +artifact; otherwise validation fails before geographic aggregation begins. + +Aggregation uses the retained output tables without rerunning either model. +The precomputed baseline and reform simulations retain their original policy +metadata so PolicyEngine can correctly recognize conditional labor-supply +analysis. `include_cliffs=true` is supported by this path and is carried through +both output planning and aggregation. + ## Temporary direct runner endpoint > **Temporary Stage 12 interface:** The authenticated routes in this section diff --git a/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_execution.py b/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_execution.py index 5b8b44c89..6a6f3ad15 100644 --- a/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_execution.py +++ b/libs/policyengine-simulation-contract/src/policyengine_simulation_contract/stage12_execution.py @@ -2,9 +2,11 @@ from __future__ import annotations +import json from dataclasses import dataclass from datetime import datetime, timedelta from enum import StrEnum +from hashlib import sha256 from typing import Annotated, Literal from uuid import UUID @@ -155,18 +157,6 @@ def require_complete_filter(self) -> GeographySelection: return self -class RequestedSimulationOutput(StrictContractModel): - schema_version: Literal[1] = 1 - variables: Annotated[tuple[ContractText, ...], Field(min_length=1)] - - @field_validator("variables") - @classmethod - def require_unique_variables(cls, value: tuple[str, ...]) -> tuple[str, ...]: - if len(value) != len(set(value)): - raise ValueError("requested variables must be unique") - return value - - class SimulationExecutionInput(StrictContractModel): contract_version: Literal[1] = 1 evaluation_id: UUID @@ -177,7 +167,6 @@ class SimulationExecutionInput(StrictContractModel): year: Annotated[int, Field(ge=1900, le=2200)] geography: GeographySelection options: dict[str, JsonValue] = Field(default_factory=dict) - requested_output: RequestedSimulationOutput bundle: BundleProvenance @model_validator(mode="before") @@ -190,6 +179,127 @@ def reject_combined_policy_input(cls, value: object) -> object: return value +class ReportAggregate(StrEnum): + BUDGET = "budget" + POVERTY = "poverty" + INEQUALITY = "inequality" + DISTRIBUTIONAL = "distributional" + WINNERS_AND_LOSERS = "winners_and_losers" + GEOGRAPHIC = "geographic" + PROGRAM_STATISTICS = "program_statistics" + + +class ReportOutputRequirements(StrictContractModel): + """Report features that determine which simulation columns must exist.""" + + schema_version: Literal[1] = 1 + aggregates: Annotated[tuple[ReportAggregate, ...], Field(min_length=1)] + include_cliff_impacts: bool + labor_supply_response_active: bool + + @field_validator("aggregates") + @classmethod + def require_complete_aggregate_profile( + cls, + value: tuple[ReportAggregate, ...], + ) -> tuple[ReportAggregate, ...]: + if len(value) != len(set(value)): + raise ValueError("report output aggregates must be unique") + if value != tuple(ReportAggregate): + raise ValueError( + "Stage 12 currently requires the complete aggregate profile" + ) + return value + + +class EntityOutputPlan(StrictContractModel): + """Required materialized columns for one country-model entity.""" + + entity: ContractText + materialized_variables: Annotated[ + tuple[ContractText, ...], + Field(min_length=1), + ] + additional_variables: tuple[ContractText, ...] = () + dataset_variables: tuple[ContractText, ...] = () + + @field_validator( + "materialized_variables", + "additional_variables", + "dataset_variables", + ) + @classmethod + def require_canonical_variables(cls, value: tuple[str, ...]) -> tuple[str, ...]: + if len(value) != len(set(value)): + raise ValueError("output-plan variables must be unique") + if value != tuple(sorted(value)): + raise ValueError("output-plan variables must use canonical order") + return value + + @model_validator(mode="after") + def require_variable_subsets(self) -> EntityOutputPlan: + if not set(self.additional_variables).issubset(self.materialized_variables): + raise ValueError( + "additional output-plan variables must be a materialized subset" + ) + if not set(self.dataset_variables).issubset(self.materialized_variables): + raise ValueError( + "dataset output-plan variables must be a materialized subset" + ) + if set(self.additional_variables).intersection(self.dataset_variables): + raise ValueError( + "calculated additional variables and dataset variables must be disjoint" + ) + return self + + +class Stage12OutputPlan(StrictContractModel): + """One immutable output schema shared by both report simulations.""" + + schema_version: Literal[1] = 1 + country: CountryId + requirements: ReportOutputRequirements + entities: Annotated[tuple[EntityOutputPlan, ...], Field(min_length=1)] + + @field_validator("entities") + @classmethod + def require_canonical_entities( + cls, + value: tuple[EntityOutputPlan, ...], + ) -> tuple[EntityOutputPlan, ...]: + names = tuple(entity.entity for entity in value) + if len(names) != len(set(names)): + raise ValueError("output-plan entities must be unique") + if names != tuple(sorted(names)): + raise ValueError("output-plan entities must use canonical order") + return value + + +def stage12_output_plan_sha256(plan: Stage12OutputPlan) -> str: + """Return the stable digest recorded by each Stage 12 child artifact.""" + + payload = json.dumps( + plan.model_dump(mode="json"), + sort_keys=True, + separators=(",", ":"), + ensure_ascii=False, + allow_nan=False, + ).encode("utf-8") + return sha256(payload).hexdigest() + + +class PlannedSimulationExecutionInput(SimulationExecutionInput): + """Internal coordinator-to-worker input with a resolved output schema.""" + + output_plan: Stage12OutputPlan + + @model_validator(mode="after") + def require_output_plan_country(self) -> PlannedSimulationExecutionInput: + if self.output_plan.country != self.geography.country: + raise ValueError("output plan country must match simulation geography") + return self + + class RowIdentity(StrictContractModel): schema_version: Literal[1] = 1 identifier_columns: Annotated[tuple[ContractText, ...], Field(min_length=1)] @@ -211,21 +321,12 @@ class SimulationArtifactDescriptor(StrictContractModel): role: SimulationRole artifact: ArtifactReference output_schema_version: Literal[1] = 1 + output_plan_sha256: Sha256Digest row_identity: RowIdentity bundle: BundleProvenance calculation_provenance: dict[str, JsonValue] | None = None -class ReportAggregate(StrEnum): - BUDGET = "budget" - POVERTY = "poverty" - INEQUALITY = "inequality" - DISTRIBUTIONAL = "distributional" - WINNERS_AND_LOSERS = "winners_and_losers" - GEOGRAPHIC = "geographic" - PROGRAM_STATISTICS = "program_statistics" - - class ReportExecutionInput(StrictContractModel): contract_version: Literal[1] = 1 evaluation_id: UUID @@ -250,7 +351,6 @@ def require_aligned_simulations(self) -> ReportExecutionInput: "year", "geography", "options", - "requested_output", "bundle", ): if getattr(self.baseline, field_name) != getattr(self.reform, field_name): diff --git a/libs/policyengine-simulation-contract/tests/test_stage12_execution.py b/libs/policyengine-simulation-contract/tests/test_stage12_execution.py new file mode 100644 index 000000000..500bf35e1 --- /dev/null +++ b/libs/policyengine-simulation-contract/tests/test_stage12_execution.py @@ -0,0 +1,117 @@ +"""Tests for typed Stage 12 execution contracts.""" + +from __future__ import annotations + +import pytest + +from policyengine_simulation_contract.stage12_execution import ( + EntityOutputPlan, + ReportAggregate, + ReportOutputRequirements, + Stage12OutputPlan, + stage12_output_plan_sha256, +) + + +def _plan() -> Stage12OutputPlan: + return Stage12OutputPlan( + country="us", + requirements=ReportOutputRequirements( + aggregates=tuple(ReportAggregate), + include_cliff_impacts=False, + labor_supply_response_active=False, + ), + entities=( + EntityOutputPlan( + entity="household", + materialized_variables=("household_id", "household_net_income"), + additional_variables=(), + ), + EntityOutputPlan( + entity="person", + materialized_variables=("federal_benefit_cost", "person_id"), + additional_variables=("federal_benefit_cost",), + ), + ), + ) + + +def test_output_plan_has_a_deterministic_digest() -> None: + plan = _plan() + + assert stage12_output_plan_sha256(plan) == stage12_output_plan_sha256( + Stage12OutputPlan.model_validate(plan.model_dump(mode="json")) + ) + + +@pytest.mark.parametrize( + ("field", "value", "message"), + [ + ( + "materialized_variables", + ("person_id", "federal_benefit_cost"), + "canonical order", + ), + ( + "materialized_variables", + ("person_id", "person_id"), + "unique", + ), + ( + "additional_variables", + ("missing",), + "subset", + ), + ( + "dataset_variables", + ("missing",), + "subset", + ), + ], +) +def test_entity_output_plan_rejects_noncanonical_variables( + field: str, + value: tuple[str, ...], + message: str, +) -> None: + values = { + "entity": "person", + "materialized_variables": ("federal_benefit_cost", "person_id"), + "additional_variables": ("federal_benefit_cost",), + "dataset_variables": (), + } + values[field] = value + + with pytest.raises(ValueError, match=message): + EntityOutputPlan.model_validate(values) + + +def test_output_plan_rejects_duplicate_or_unsorted_entities() -> None: + entity = _plan().entities[0] + + with pytest.raises(ValueError, match="canonical order"): + Stage12OutputPlan( + country="us", + requirements=_plan().requirements, + entities=( + _plan().entities[1], + entity, + ), + ) + + with pytest.raises(ValueError, match="unique"): + Stage12OutputPlan( + country="us", + requirements=_plan().requirements, + entities=(entity, entity), + ) + + +def test_entity_output_plan_separates_calculated_and_dataset_variables() -> None: + with pytest.raises(ValueError, match="must be disjoint"): + EntityOutputPlan( + entity="household", + materialized_variables=("constituency_code_oa", "household_id"), + additional_variables=("constituency_code_oa",), + dataset_variables=("constituency_code_oa",), + ) 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 140215c03..c51e7aa2e 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 @@ -17,7 +17,6 @@ GeographySelection, ReportAggregate, ReportExecutionInput, - RequestedSimulationOutput, SimulationExecutionInput, SimulationRole, ) @@ -28,7 +27,6 @@ class ComparisonSkipReason(StrEnum): UNSUPPORTED_FLOW = "unsupported_flow" UNSUPPORTED_SCOPE = "unsupported_scope" UNSUPPORTED_BUDGET_WINDOW = "unsupported_budget_window" - UNSUPPORTED_CLIFF_CALCULATION = "unsupported_cliff_calculation" UNSUPPORTED_COUNTRY = "unsupported_country" UNSUPPORTED_DATASET = "unsupported_dataset" MISSING_BUNDLE_PROVENANCE = "missing_bundle_provenance" @@ -98,8 +96,9 @@ def adapt_annual_comparison( return _skip(ComparisonSkipReason.UNSUPPORTED_OPTIONS) if payload.get("scope") != "macro": return _skip(ComparisonSkipReason.UNSUPPORTED_SCOPE) - if payload.get("include_cliffs") is True: - return _skip(ComparisonSkipReason.UNSUPPORTED_CLIFF_CALCULATION) + include_cliffs = payload.get("include_cliffs") + if include_cliffs is not None and not isinstance(include_cliffs, bool): + return _skip(ComparisonSkipReason.UNSUPPORTED_REQUEST_SHAPE) if payload.get("segmented") not in {None, False}: return _skip(ComparisonSkipReason.UNSUPPORTED_OPTIONS) country = payload.get("country") @@ -177,7 +176,8 @@ def adapt_annual_comparison( region=region, ) options = {key: payload[key] for key in ("spm",) if payload.get(key) is not None} - output = RequestedSimulationOutput(variables=("*",)) + if include_cliffs is True: + options["include_cliffs"] = True def simulation(role: SimulationRole, policy: dict) -> SimulationExecutionInput: return SimulationExecutionInput( @@ -189,7 +189,6 @@ def simulation(role: SimulationRole, policy: dict) -> SimulationExecutionInput: year=int(year_text), geography=geography, options=options, - requested_output=output, bundle=provenance, ) diff --git a/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py b/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py index 891f8bb80..16c435d9d 100644 --- a/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py +++ b/projects/policyengine-simulation-entry/tests/test_stage12_adapter.py @@ -93,13 +93,28 @@ def test_normalized_production_telemetry_does_not_change_eligibility() -> None: assert result.report.reform.options == {} +def test_cliff_analysis_is_forwarded_to_both_simulations() -> None: + payload = {**eligible_payload(), "include_cliffs": True} + + result = adapt_annual_comparison( + payload, + evaluation_id=EVALUATION_ID, + worker=worker(), + ) + + assert result.report is not None + assert result.skip_reason is None + assert result.report.baseline.options == {"include_cliffs": True} + assert result.report.reform.options == {"include_cliffs": True} + + @pytest.mark.parametrize( ("update", "reason"), [ ({"scope": "household"}, ComparisonSkipReason.UNSUPPORTED_SCOPE), ( - {"include_cliffs": True}, - ComparisonSkipReason.UNSUPPORTED_CLIFF_CALCULATION, + {"include_cliffs": "true"}, + ComparisonSkipReason.UNSUPPORTED_REQUEST_SHAPE, ), ({"country": "ca"}, ComparisonSkipReason.UNSUPPORTED_COUNTRY), ({"data": "another"}, ComparisonSkipReason.UNSUPPORTED_DATASET), diff --git a/projects/policyengine-simulation-entry/tests/test_stage12_backend.py b/projects/policyengine-simulation-entry/tests/test_stage12_backend.py index c465bb47d..b71fda113 100644 --- a/projects/policyengine-simulation-entry/tests/test_stage12_backend.py +++ b/projects/policyengine-simulation-entry/tests/test_stage12_backend.py @@ -167,7 +167,7 @@ def test_repeated_automatic_submission_uses_one_deterministic_report_identity() ) -def test_unsupported_automatic_input_does_not_dispatch_or_write() -> None: +def test_cliff_automatic_input_dispatches_without_entrypoint_writes() -> None: store = FakeStore() invoker = FakeInvoker() @@ -180,7 +180,10 @@ def test_unsupported_automatic_input_does_not_dispatch_or_write() -> None: ) assert store.events == [] - assert invoker.calls == [] + assert len(invoker.calls) == 1 + report = invoker.calls[0]["report_payload"] + assert report["baseline"]["options"] == {"include_cliffs": True} + assert report["reform"]["options"] == {"include_cliffs": True} def test_dispatch_failure_is_sanitized_without_database_cleanup() -> None: @@ -305,12 +308,12 @@ def test_temporary_submission_rejects_unsupported_input_without_side_effects() - with pytest.raises(TemporaryStage12UnsupportedRequest) as error: asyncio.run( _backend(store, invoker).submit_temporary_report( - request_payload={**eligible_payload(), "include_cliffs": True}, + request_payload={**eligible_payload(), "include_cliffs": "true"}, request_id="request-1", ) ) - assert error.value.reason == "unsupported_cliff_calculation" + assert error.value.reason == "unsupported_request_shape" assert store.events == [] assert invoker.calls == [] diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_artifacts.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_artifacts.py index 490318650..7427b208d 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_artifacts.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_artifacts.py @@ -18,10 +18,11 @@ AggregateReportArtifactPayload, ArtifactMediaType, ArtifactReference, + PlannedSimulationExecutionInput, ResultComparisonArtifactPayload, RowIdentity, SimulationArtifactDescriptor, - SimulationExecutionInput, + stage12_output_plan_sha256, ) from pydantic import JsonValue @@ -166,9 +167,17 @@ def deserialize_simulation_frames(payload: bytes) -> dict[str, pd.DataFrame]: if raw_dtypes is None: raise ValueError("Stage 12 simulation artifact has no dtype metadata") dtypes = json.loads(raw_dtypes) + if not isinstance(dtypes, dict): + raise TypeError("Stage 12 simulation artifact dtype metadata is invalid") combined = table.to_pandas() frames: dict[str, pd.DataFrame] = {} for entity in sorted(combined[PARQUET_CONTRACT.entity_column].unique()): + entity_name = str(entity) + entity_dtypes = dtypes.get(entity_name) + if not isinstance(entity_dtypes, dict): + raise ValueError( + f"Stage 12 simulation artifact has no dtype schema for {entity_name}" + ) frame = combined.loc[combined[PARQUET_CONTRACT.entity_column] == entity].copy() frame = frame.sort_values(PARQUET_CONTRACT.row_order_column, kind="stable") frame = frame.drop( @@ -177,11 +186,14 @@ def deserialize_simulation_frames(payload: bytes) -> dict[str, pd.DataFrame]: PARQUET_CONTRACT.row_order_column, ] ) - frame = frame.dropna(axis=1, how="all").reset_index(drop=True) - for column, dtype in dtypes[entity].items(): + # The dtype metadata is the authoritative per-entity schema. Inferring + # ownership by dropping all-null columns loses legitimate planned + # outputs whose values happen to be null for every row. + frame = frame.reindex(columns=sorted(entity_dtypes)).reset_index(drop=True) + for column, dtype in entity_dtypes.items(): if column in frame.columns and str(frame[column].dtype) != dtype: frame[column] = frame[column].astype(dtype) - frames[str(entity)] = frame.reindex(columns=sorted(frame.columns)) + frames[entity_name] = frame return frames @@ -237,7 +249,7 @@ def write_input( self, *, prefix: str, - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, ) -> ArtifactReference: return self._write_immutable( input_path(prefix=prefix, role=simulation.role.value), @@ -249,7 +261,7 @@ def write_simulation( self, *, prefix: str, - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, frames: Mapping[str, pd.DataFrame], calculation_provenance: Mapping[str, Any] | None = None, ) -> SimulationArtifactDescriptor: @@ -275,6 +287,7 @@ def write_simulation( simulation_execution_id=simulation.simulation_execution_id, role=simulation.role, artifact=artifact, + output_plan_sha256=stage12_output_plan_sha256(simulation.output_plan), row_identity=row_identity, bundle=simulation.bundle, calculation_provenance=normalized_provenance, diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_qualification.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_qualification.py index c124fc3fe..e6999ea57 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_qualification.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_qualification.py @@ -13,9 +13,10 @@ from policyengine_simulation_contract.stage12_execution import ( ArtifactMediaType, ArtifactReference, + PlannedSimulationExecutionInput, ReportExecutionInput, SimulationArtifactDescriptor, - SimulationExecutionInput, + stage12_output_plan_sha256, ) from policyengine_simulation_executor.simulation_microdata import ( rebuild_entity_frame, @@ -34,6 +35,10 @@ SimulationCalculation, simulation_input_sha256, ) +from policyengine_simulation_executor.stage12_runtime.output_planning import ( + plan_simulation_input, + resolve_report_output_plan, +) class Stage12ParityReceipt(BaseModel): @@ -58,7 +63,7 @@ class Stage12ParityReceipt(BaseModel): ExistingRunner = Callable[[dict[str, Any]], Mapping[str, Any]] SingleSimulationRunner = Callable[ - [SimulationExecutionInput], + [PlannedSimulationExecutionInput], Mapping[str, pd.DataFrame] | SimulationCalculation, ] AggregateBuilder = Callable[..., dict[str, Any]] @@ -147,7 +152,7 @@ def _calculation( def _descriptor( - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, calculation: SimulationCalculation, ) -> SimulationArtifactDescriptor: payload, row_identity = serialize_simulation_frames( @@ -167,6 +172,7 @@ def _descriptor( content_sha256=sha256(payload).hexdigest(), size_bytes=len(payload), ), + output_plan_sha256=stage12_output_plan_sha256(simulation.output_plan), row_identity=row_identity, bundle=simulation.bundle, calculation_provenance=cast( @@ -190,6 +196,9 @@ def qualify_report_parity( report = ReportExecutionInput.model_validate(report_payload) _validate_matching_input(existing_request, report) + output_plan = resolve_report_output_plan(report) + baseline_input = plan_simulation_input(report.baseline, output_plan) + reform_input = plan_simulation_input(report.reform, output_plan) if existing_runner is None: from policyengine_simulation_executor.simulation_runtime import ( @@ -230,8 +239,8 @@ def qualify_report_parity( ) with ThreadPoolExecutor(max_workers=2) as executor: - baseline_future = executor.submit(single_simulation_runner, report.baseline) - reform_future = executor.submit(single_simulation_runner, report.reform) + baseline_future = executor.submit(single_simulation_runner, baseline_input) + reform_future = executor.submit(single_simulation_runner, reform_input) baseline_calculation = _calculation(baseline_future.result()) reform_calculation = _calculation(reform_future.result()) @@ -246,8 +255,8 @@ def qualify_report_parity( tolerances=simulation_tolerances, ) - baseline_descriptor = _descriptor(report.baseline, baseline_calculation) - reform_descriptor = _descriptor(report.reform, reform_calculation) + baseline_descriptor = _descriptor(baseline_input, baseline_calculation) + reform_descriptor = _descriptor(reform_input, reform_calculation) v2_report = aggregate_builder( report=report, baseline_frames=baseline_calculation.frames, @@ -277,8 +286,8 @@ def qualify_report_parity( return Stage12ParityReceipt( evaluation_id=str(report.evaluation_id), bundle_manifest_sha256=bundle.bundle_manifest_sha256, - baseline_input_sha256=simulation_input_sha256(report.baseline), - reform_input_sha256=simulation_input_sha256(report.reform), + baseline_input_sha256=simulation_input_sha256(baseline_input), + reform_input_sha256=simulation_input_sha256(reform_input), existing_baseline_output_sha256=sha256(existing_baseline_payload).hexdigest(), v2_baseline_output_sha256=sha256(v2_baseline_payload).hexdigest(), existing_reform_output_sha256=sha256(existing_reform_payload).hexdigest(), diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/aggregation.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/aggregation.py index f2d6a5bae..18e9b98ac 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/aggregation.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/aggregation.py @@ -10,11 +10,15 @@ from policyengine_simulation_contract.stage12_execution import ( ReportExecutionInput, SimulationArtifactDescriptor, + Stage12OutputPlan, + stage12_output_plan_sha256, ) +from pydantic import JsonValue def validate_aligned_outputs( report: ReportExecutionInput, + output_plan: Stage12OutputPlan, baseline: SimulationArtifactDescriptor, reform: SimulationArtifactDescriptor, ) -> None: @@ -29,6 +33,12 @@ def validate_aligned_outputs( raise ValueError("simulation artifacts have incompatible bundle provenance") if baseline.output_schema_version != reform.output_schema_version: raise ValueError("simulation artifacts use incompatible output schemas") + expected_plan_sha256 = stage12_output_plan_sha256(output_plan) + if ( + baseline.output_plan_sha256 != expected_plan_sha256 + or reform.output_plan_sha256 != expected_plan_sha256 + ): + raise ValueError("simulation artifacts do not satisfy the report output plan") if baseline.row_identity != reform.row_identity: raise ValueError("simulation artifacts have incompatible stable row identities") @@ -90,7 +100,10 @@ def build_aggregate_report( from policyengine_simulation_executor.simulation_output_builder import ( SimulationOutputBuilder, ) - from policyengine_simulation_executor.simulation_runtime import _country_module + from policyengine_simulation_executor.simulation_runtime import ( + _country_module, + _normalise_policy, + ) country = report.baseline.geography.country country_module = _country_module(country) @@ -107,11 +120,11 @@ def build_aggregate_report( ), } - def stand_in(dataset): + def stand_in(dataset, policy: dict[str, JsonValue]): simulation = PrecomputedSimulation( dataset=dataset, tax_benefit_model_version=country_module.model, - policy=None, + policy=_normalise_policy(policy), ) simulation.output_dataset = dataset return simulation @@ -130,8 +143,8 @@ def stand_in(dataset): simulation_params=params, country_module=country_module, dataset=datasets["baseline"], - baseline=stand_in(datasets["baseline"]), - reform=stand_in(datasets["reform"]), + baseline=stand_in(datasets["baseline"], report.baseline.policy), + reform=stand_in(datasets["reform"], report.reform.policy), resolved_data_version=report.baseline.bundle.dataset.artifact_revision, resolved_region_code=report.baseline.geography.region, ).serialize() diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/coordination.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/coordination.py index 67b4d9f90..de94e4f97 100644 --- a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/coordination.py +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/coordination.py @@ -14,10 +14,11 @@ ComparisonRunAggregationStatus, ComparisonRunLifecycleStatus, ComparisonSimulationRecord, + PlannedSimulationExecutionInput, ReportExecutionInput, ResultComparisonStatus, SimulationArtifactDescriptor, - SimulationExecutionInput, + Stage12OutputPlan, Stage12InvocationContext, ) @@ -38,13 +39,18 @@ runtime_store, ) from .simulation import descriptor_from_record, simulation_input_sha256 +from .output_planning import ( + plan_simulation_input, + resolve_report_output_plan, + validate_output_frames, +) SIMULATION_WAIT_TIMEOUT_SECONDS = 3_000 def _child_record( *, - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, context: Stage12InvocationContext, function_name: str, ) -> ComparisonSimulationRecord: @@ -179,6 +185,9 @@ def coordinate_report( artifacts: Stage12ArtifactStore | None = None, invoker: ChildInvoker | None = None, aggregator: Callable[..., dict[str, Any]] = build_aggregate_report, + output_plan_resolver: Callable[ + [ReportExecutionInput], Stage12OutputPlan + ] = resolve_report_output_plan, ) -> dict[str, Any]: report = ReportExecutionInput.model_validate(payload) context = Stage12InvocationContext.model_validate(context_payload) @@ -220,10 +229,14 @@ def coordinate_report( } ) function_name = context.simulation_callable - simulations = (report.baseline, report.reform) descriptors: dict[str, SimulationArtifactDescriptor] = {} calls: dict[str, ChildCall] = {} try: + output_plan = output_plan_resolver(report) + simulations = ( + plan_simulation_input(report.baseline, output_plan), + plan_simulation_input(report.reform, output_plan), + ) children = { simulation.role: persistence.create_or_resolve_simulation( _child_record( @@ -335,17 +348,21 @@ def coordinate_report( raise baseline = descriptors["baseline"] reform = descriptors["reform"] - validate_aligned_outputs(report, baseline, reform) + validate_aligned_outputs(report, output_plan, baseline, reform) baseline_payload = artifact_storage.read(baseline.artifact.uri) reform_payload = artifact_storage.read(reform.artifact.uri) if sha256(baseline_payload).hexdigest() != baseline.artifact.content_sha256: raise ValueError("baseline artifact digest mismatch") if sha256(reform_payload).hexdigest() != reform.artifact.content_sha256: raise ValueError("reform artifact digest mismatch") + baseline_frames = deserialize_simulation_frames(baseline_payload) + reform_frames = deserialize_simulation_frames(reform_payload) + validate_output_frames(baseline_frames, output_plan) + validate_output_frames(reform_frames, output_plan) aggregate = aggregator( report=report, - baseline_frames=deserialize_simulation_frames(baseline_payload), - reform_frames=deserialize_simulation_frames(reform_payload), + baseline_frames=baseline_frames, + reform_frames=reform_frames, baseline_descriptor=baseline, reform_descriptor=reform, ) diff --git a/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/output_planning.py b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/output_planning.py new file mode 100644 index 000000000..c6dea0d2b --- /dev/null +++ b/projects/policyengine-simulation-executor/src/policyengine_simulation_executor/stage12_runtime/output_planning.py @@ -0,0 +1,254 @@ +"""Resolve, apply, and verify Stage 12 simulation output schemas.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Protocol, cast + +import pandas as pd +from policyengine.core import Simulation +from policyengine.outputs import ( + configure_cliff_impact_variables, + configure_labor_supply_response_variables, + labor_supply_response_is_active, +) +from policyengine_simulation_contract.stage12_bundle import CountryId +from policyengine_simulation_contract.stage12_execution import ( + EntityOutputPlan, + PlannedSimulationExecutionInput, + ReportAggregate, + ReportExecutionInput, + ReportOutputRequirements, + SimulationExecutionInput, + Stage12OutputPlan, +) + +UK_GEOGRAPHIC_DATASET_VARIABLES: dict[str, tuple[str, ...]] = { + "household": ("constituency_code_oa", "la_code_oa"), +} + + +class OutputVariableModel(Protocol): + """Country-model behavior required by the Stage 12 planner.""" + + def resolve_entity_variables( + self, + simulation: Simulation, + ) -> dict[str, list[str]]: ... + + +def _country_model(country: CountryId) -> OutputVariableModel: + if country == "us": + from policyengine.tax_benefit_models import us + + return us.model + from policyengine.tax_benefit_models import uk + + return uk.model + + +def _planning_simulation( + simulation: SimulationExecutionInput, + *, + model: OutputVariableModel, +) -> Simulation: + """Build a data-free object used only by PolicyEngine output configurators. + + ``Simulation`` normally requires a dataset. Output planning only needs its + policy, model version, and ``extra_variables`` fields, so validation is + intentionally bypassed. This object must never call ``ensure`` or access a + dataset. + """ + + return Simulation.model_construct( + id=f"stage12-output-plan-{simulation.role.value}", + policy=dict(simulation.policy), + dynamic=None, + dataset=None, + scoping_strategy=None, + extra_variables={}, + tax_benefit_model_version=model, + output_dataset=None, + ) + + +def _include_cliff_impacts(report: ReportExecutionInput) -> bool: + value = report.baseline.options.get("include_cliffs", False) + if not isinstance(value, bool): + raise TypeError("include_cliffs must be a boolean") + return value + + +def _required_dataset_variables( + *, + country: CountryId, + aggregates: tuple[ReportAggregate, ...], +) -> dict[str, tuple[str, ...]]: + """Return input-dataset columns copied into the materialized output.""" + + if country == "uk" and ReportAggregate.GEOGRAPHIC in aggregates: + return UK_GEOGRAPHIC_DATASET_VARIABLES + return {} + + +def _configure_country_outputs( + *, + country: CountryId, + baseline: Simulation, + reform: Simulation, + include_cliff_impacts: bool, +) -> bool: + labor_supply_active = labor_supply_response_is_active( + baseline, + reform, + country_code=country, + ) + configure_labor_supply_response_variables( + baseline, + reform, + country_code=country, + ) + if include_cliff_impacts: + configure_cliff_impact_variables(baseline, reform) + if country == "us": + from policyengine.tax_benefit_models.us.analysis import ( + configure_budgetary_impact_variables, + ) + + configure_budgetary_impact_variables(baseline, reform) + return labor_supply_active + + +def resolve_report_output_plan(report: ReportExecutionInput) -> Stage12OutputPlan: + """Resolve one country-owned output schema for both report simulations.""" + + country = report.baseline.geography.country + model = _country_model(country) + baseline = _planning_simulation(report.baseline, model=model) + reform = _planning_simulation(report.reform, model=model) + include_cliff_impacts = _include_cliff_impacts(report) + labor_supply_active = _configure_country_outputs( + country=country, + baseline=baseline, + reform=reform, + include_cliff_impacts=include_cliff_impacts, + ) + requirements = ReportOutputRequirements( + aggregates=report.requested_aggregates, + include_cliff_impacts=include_cliff_impacts, + labor_supply_response_active=labor_supply_active, + ) + baseline_variables = model.resolve_entity_variables(baseline) + reform_variables = model.resolve_entity_variables(reform) + dataset_variables = _required_dataset_variables( + country=country, + aggregates=requirements.aggregates, + ) + entities = [] + for entity in sorted( + set(baseline_variables) | set(reform_variables) | set(dataset_variables) + ): + entity_dataset_variables = dataset_variables.get(entity, ()) + materialized = tuple( + sorted( + set(baseline_variables.get(entity, ())) + | set(reform_variables.get(entity, ())) + | set(entity_dataset_variables) + ) + ) + additional = tuple( + sorted( + set((baseline.extra_variables or {}).get(entity, ())) + | set((reform.extra_variables or {}).get(entity, ())) + ) + ) + entities.append( + EntityOutputPlan( + entity=entity, + materialized_variables=materialized, + additional_variables=additional, + dataset_variables=entity_dataset_variables, + ) + ) + return Stage12OutputPlan( + country=country, + requirements=requirements, + entities=tuple(entities), + ) + + +def plan_simulation_input( + simulation: SimulationExecutionInput, + output_plan: Stage12OutputPlan, +) -> PlannedSimulationExecutionInput: + """Attach a coordinator-resolved schema to one child input.""" + + return PlannedSimulationExecutionInput.model_validate( + { + **simulation.model_dump(mode="json"), + "output_plan": output_plan.model_dump(mode="json"), + } + ) + + +def apply_output_plan( + simulation: Simulation, + output_plan: Stage12OutputPlan, +) -> None: + """Install required extra variables before a live simulation is ensured.""" + + extras = { + entity: list(variables) + for entity, variables in (simulation.extra_variables or {}).items() + } + for entity_plan in output_plan.entities: + entity_extras = extras.setdefault(entity_plan.entity, []) + for variable in entity_plan.additional_variables: + if variable not in entity_extras: + entity_extras.append(variable) + simulation.extra_variables = extras + model = cast(OutputVariableModel, simulation.tax_benefit_model_version) + resolved = model.resolve_entity_variables(simulation) + missing = { + entity_plan.entity: sorted( + ( + set(entity_plan.materialized_variables) + - set(entity_plan.dataset_variables) + ) + - set(resolved.get(entity_plan.entity, ())) + ) + for entity_plan in output_plan.entities + } + missing = {entity: variables for entity, variables in missing.items() if variables} + if missing: + raise ValueError( + f"configured simulation does not satisfy output plan: {missing}" + ) + + +def validate_output_frames( + frames: Mapping[str, pd.DataFrame], + output_plan: Stage12OutputPlan, +) -> None: + """Require every planned entity and column in materialized output frames.""" + + missing_entities = [ + entity.entity for entity in output_plan.entities if entity.entity not in frames + ] + if missing_entities: + raise ValueError( + f"simulation output is missing planned entities: {missing_entities}" + ) + missing_columns = { + entity.entity: sorted( + set(entity.materialized_variables) - set(frames[entity.entity].columns) + ) + for entity in output_plan.entities + } + missing_columns = { + entity: columns for entity, columns in missing_columns.items() if columns + } + if missing_columns: + raise ValueError( + f"simulation output is missing planned columns: {missing_columns}" + ) 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 d064cf7e5..6c5230baa 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 @@ -13,9 +13,11 @@ from policyengine_simulation_contract.stage12_execution import ( ComparisonRunLifecycleStatus, ComparisonSimulationRecord, + PlannedSimulationExecutionInput, SimulationArtifactDescriptor, SimulationExecutionInput, Stage12InvocationContext, + stage12_output_plan_sha256, ) from policyengine_simulation_executor.stage12_artifacts import ( @@ -25,6 +27,7 @@ from policyengine_simulation_executor.stage12_bundle import load_stage12_bundle from .dependencies import ComparisonStore, artifact_store, runtime_store +from .output_planning import apply_output_plan, validate_output_frames @dataclass(frozen=True) @@ -42,7 +45,7 @@ def simulation_input_sha256(simulation: SimulationExecutionInput) -> str: def _require_context( - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, context: Stage12InvocationContext, *, required_country: CountryId, @@ -57,10 +60,8 @@ def _require_context( raise ValueError( "single-simulation artifact prefix names another comparison run" ) - if simulation.requested_output.variables != ("*",): - raise ValueError( - "Stage 12 report simulations require the complete output table" - ) + if simulation.output_plan.country != required_country: + raise ValueError("single-simulation output plan names another country") def _require_installed_bundle(simulation: SimulationExecutionInput) -> None: @@ -93,9 +94,9 @@ def _require_installed_bundle(simulation: SimulationExecutionInput) -> None: def calculate_simulation_frames( - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, ) -> SimulationCalculation: - """Run exactly one policy and return its complete entity output tables.""" + """Run one policy and return the coordinator-planned entity output tables.""" if simulation.population.kind != "dataset": raise ValueError("Stage 12 society-wide worker requires a dataset population") @@ -144,12 +145,14 @@ def calculate_simulation_frames( scoping_strategy=region.scoping_strategy, region_code=region.code, ) + apply_output_plan(model, simulation.output_plan) model.ensure() output_data = getattr(getattr(model, "output_dataset", None), "data", None) entity_data = getattr(output_data, "entity_data", None) if not isinstance(entity_data, Mapping): raise TypeError("simulation produced no entity output tables") frames = {entity: pd.DataFrame(frame) for entity, frame in entity_data.items()} + validate_output_frames(frames, simulation.output_plan) selection = getattr(model, "spm_config", None) calculation_provenance = None if selection is not None: @@ -177,7 +180,7 @@ def calculate_simulation_frames( def descriptor_from_record( child: ComparisonSimulationRecord, - simulation: SimulationExecutionInput, + simulation: PlannedSimulationExecutionInput, ) -> SimulationArtifactDescriptor: output_uri = child.output_uri output_sha256 = child.output_sha256 @@ -211,6 +214,7 @@ def descriptor_from_record( size_bytes=0, ), output_schema_version=1, + output_plan_sha256=stage12_output_plan_sha256(simulation.output_plan), row_identity=RowIdentity( identifier_columns=row_identity_columns, row_count=row_count, @@ -228,11 +232,11 @@ def run_single_simulation( store: ComparisonStore | None = None, artifacts: Stage12ArtifactStore | None = None, calculator: Callable[ - [SimulationExecutionInput], + [PlannedSimulationExecutionInput], Mapping[str, pd.DataFrame] | SimulationCalculation, ] = calculate_simulation_frames, ) -> dict[str, Any]: - simulation = SimulationExecutionInput.model_validate(payload) + simulation = PlannedSimulationExecutionInput.model_validate(payload) context = Stage12InvocationContext.model_validate(context_payload) _require_context(simulation, context, required_country=required_country) persistence = store or runtime_store() @@ -263,6 +267,7 @@ def run_single_simulation( else: frames = calculated calculation_provenance = None + validate_output_frames(frames, simulation.output_plan) descriptor = artifact_storage.write_simulation( prefix=context.artifact_prefix, simulation=simulation, diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_artifacts.py b/projects/policyengine-simulation-executor/tests/test_stage12_artifacts.py index fdaa9ca85..019035377 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_artifacts.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_artifacts.py @@ -49,6 +49,21 @@ def test_parquet_encoding_is_deterministic_and_preserves_rows_and_dtypes() -> No assert restored["person"]["person_id"].tolist() == [1, 2] +def test_parquet_round_trip_preserves_an_all_null_entity_column() -> None: + frames = _frames() + frames["person"]["planned_value"] = pd.Series( + [pd.NA, pd.NA], + dtype="Float64", + ) + + payload, _ = serialize_simulation_frames(frames) + restored = deserialize_simulation_frames(payload) + + assert "planned_value" in restored["person"].columns + assert str(restored["person"]["planned_value"].dtype) == "Float64" + assert restored["person"]["planned_value"].isna().all() + + def test_parquet_retains_detached_calculation_provenance() -> None: provenance = { "spm_config": {"scenario": "official"}, diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_output_planning.py b/projects/policyengine-simulation-executor/tests/test_stage12_output_planning.py new file mode 100644 index 000000000..02a857ea9 --- /dev/null +++ b/projects/policyengine-simulation-executor/tests/test_stage12_output_planning.py @@ -0,0 +1,271 @@ +"""Tests for country-aware Stage 12 output planning.""" + +from __future__ import annotations + +from uuid import UUID + +import pandas as pd +import pytest +from policyengine.core import Simulation +from policyengine_simulation_contract.stage12_bundle import CountryId +from policyengine_simulation_contract.stage12_execution import ( + BundleProvenance, + DatasetArtifactMediaType, + DatasetArtifactReference, + DatasetPopulationInput, + DatasetProvenance, + GeographySelection, + ReportAggregate, + ReportExecutionInput, + SimulationExecutionInput, + SimulationRole, + Stage12OutputPlan, +) +from pydantic import JsonValue + +from policyengine_simulation_executor.stage12_runtime.output_planning import ( + apply_output_plan, + resolve_report_output_plan, + validate_output_frames, +) + +EVALUATION_ID = UUID("00000000-0000-0000-0000-000000000001") + + +def _bundle(country: CountryId) -> BundleProvenance: + package_name, package_version, dataset = ( + ("policyengine-us", "1.764.6", "populace_us_2024") + if country == "us" + else ("policyengine-uk", "2.90.2", "populace_uk_2023") + ) + return BundleProvenance( + policyengine_version="5.2.0", + country_package_name=package_name, + country_package_version=package_version, + dataset=DatasetProvenance( + identity=dataset, + uri=f"hf://policyengine/data/{dataset}.h5@revision", + artifact_revision="revision", + data_package_name="populace-data", + data_package_version="0.1.0", + ), + bundle_manifest_sha256="b" * 64, + ) + + +def _simulation( + country: CountryId, + role: SimulationRole, + *, + policy: dict[str, JsonValue] | None = None, + options: dict[str, JsonValue] | None = None, +) -> SimulationExecutionInput: + bundle = _bundle(country) + return SimulationExecutionInput( + evaluation_id=EVALUATION_ID, + simulation_execution_id=UUID( + "00000000-0000-0000-0000-000000000002" + if role is SimulationRole.BASELINE + else "00000000-0000-0000-0000-000000000003" + ), + role=role, + policy=policy or {}, + population=DatasetPopulationInput( + artifact=DatasetArtifactReference( + uri=bundle.dataset.uri, + media_type=DatasetArtifactMediaType.HDF5, + content_sha256="a" * 64, + ) + ), + year=2026, + geography=GeographySelection(country=country, region=country), + options=options or {}, + bundle=bundle, + ) + + +def _report( + country: CountryId = "us", + *, + reform: dict[str, JsonValue] | None = None, + include_cliffs: bool = False, + aggregates: tuple[ReportAggregate, ...] = tuple(ReportAggregate), +) -> ReportExecutionInput: + options: dict[str, JsonValue] = {"include_cliffs": True} if include_cliffs else {} + return ReportExecutionInput( + evaluation_id=EVALUATION_ID, + baseline=_simulation( + country, + SimulationRole.BASELINE, + options=options, + ), + reform=_simulation( + country, + SimulationRole.REFORM, + policy=reform, + options=options, + ), + requested_aggregates=aggregates, + ) + + +def _variables(plan: Stage12OutputPlan, entity: str) -> set[str]: + entity_plan = next(item for item in plan.entities if item.entity == entity) + return set(entity_plan.materialized_variables) + + +def _additional_variables(plan: Stage12OutputPlan, entity: str) -> set[str]: + entity_plan = next(item for item in plan.entities if item.entity == entity) + return set(entity_plan.additional_variables) + + +def _dataset_variables(plan: Stage12OutputPlan, entity: str) -> set[str]: + entity_plan = next(item for item in plan.entities if item.entity == entity) + return set(entity_plan.dataset_variables) + + +def test_us_plan_adds_budget_variables_to_the_country_defaults() -> None: + plan = resolve_report_output_plan(_report()) + + person_variables = _variables(plan, "person") + assert {"person_id", "age"}.issubset(person_variables) + assert {"federal_benefit_cost", "state_benefit_cost"}.issubset(person_variables) + assert {"federal_benefit_cost", "state_benefit_cost"}.issubset( + _additional_variables(plan, "person") + ) + + +def test_uk_plan_uses_uk_defaults_without_us_budget_variables() -> None: + plan = resolve_report_output_plan(_report("uk")) + + assert "benunit" in {entity.entity for entity in plan.entities} + assert "tax_unit" not in {entity.entity for entity in plan.entities} + assert "federal_benefit_cost" not in _variables(plan, "person") + assert _dataset_variables(plan, "household") == { + "constituency_code_oa", + "la_code_oa", + } + assert _dataset_variables(plan, "household").issubset(_variables(plan, "household")) + + +def test_cliff_variables_are_conditional() -> None: + without_cliffs = resolve_report_output_plan(_report()) + with_cliffs = resolve_report_output_plan(_report(include_cliffs=True)) + + assert with_cliffs.requirements.include_cliff_impacts is True + assert {"cliff_gap", "is_on_cliff", "is_adult"}.issubset( + _additional_variables(with_cliffs, "person") + ) + assert "cliff_gap" not in _additional_variables(without_cliffs, "person") + + +def test_labor_supply_variables_are_shared_when_either_policy_activates_them() -> None: + plan = resolve_report_output_plan( + _report( + reform={"gov.simulation.labor_supply_responses.elasticities.income": 0.1} + ) + ) + + assert plan.requirements.labor_supply_response_active is True + assert { + "income_elasticity_lsr", + "substitution_elasticity_lsr", + "weekly_hours_worked_behavioural_response_income_elasticity", + "weekly_hours_worked_behavioural_response_substitution_elasticity", + }.issubset(_additional_variables(plan, "person")) + + +def test_unrelated_policy_does_not_activate_labor_supply_outputs() -> None: + plan = resolve_report_output_plan( + _report(reform={"gov.irs.credits.ctc.amount.base[0].amount": 3_000}) + ) + + assert plan.requirements.labor_supply_response_active is False + assert "income_elasticity_lsr" not in _additional_variables(plan, "person") + + +def test_planning_is_deterministic_and_never_ensures_a_simulation(monkeypatch) -> None: + def reject_ensure(_simulation): + raise AssertionError("output planning must not execute a simulation") + + monkeypatch.setattr(Simulation, "ensure", reject_ensure) + + first = resolve_report_output_plan(_report(include_cliffs=True)) + second = resolve_report_output_plan(_report(include_cliffs=True)) + + assert first == second + + +def test_incomplete_aggregate_profile_is_rejected_before_dispatch() -> None: + with pytest.raises(ValueError, match="complete aggregate profile"): + resolve_report_output_plan(_report(aggregates=(ReportAggregate.BUDGET,))) + + +def test_apply_and_validate_output_plan_reject_missing_materialized_columns() -> None: + from policyengine.tax_benefit_models import us + + plan = resolve_report_output_plan(_report()) + simulation = Simulation.model_construct( + policy={}, + dynamic=None, + dataset=None, + scoping_strategy=None, + extra_variables={}, + tax_benefit_model_version=us.model, + output_dataset=None, + ) + + apply_output_plan(simulation, plan) + resolved = us.model.resolve_entity_variables(simulation) + person_plan = next(entity for entity in plan.entities if entity.entity == "person") + assert set(person_plan.additional_variables).issubset( + set(resolved[person_plan.entity]) + ) + + frames = { + entity.entity: pd.DataFrame( + {variable: [0] for variable in entity.materialized_variables} + ) + for entity in plan.entities + } + validate_output_frames(frames, plan) + frames["person"] = frames["person"].drop(columns=["federal_benefit_cost"]) + with pytest.raises(ValueError, match="federal_benefit_cost"): + validate_output_frames(frames, plan) + + +@pytest.mark.parametrize("missing", ["constituency_code_oa", "la_code_oa"]) +def test_uk_output_validation_requires_geographic_dataset_columns( + missing: str, +) -> None: + from policyengine.tax_benefit_models import uk + + plan = resolve_report_output_plan(_report("uk")) + simulation = Simulation.model_construct( + policy={}, + dynamic=None, + dataset=None, + scoping_strategy=None, + extra_variables={}, + tax_benefit_model_version=uk.model, + output_dataset=None, + ) + + # Dataset variables are copied by the country model rather than calculated, + # so applying the plan must not try to add them to ``extra_variables``. + apply_output_plan(simulation, plan) + assert not _dataset_variables(plan, "household").intersection( + simulation.extra_variables.get("household", ()) + ) + + frames = { + entity.entity: pd.DataFrame( + {variable: [0] for variable in entity.materialized_variables} + ) + for entity in plan.entities + } + validate_output_frames(frames, plan) + frames["household"] = frames["household"].drop(columns=[missing]) + + with pytest.raises(ValueError, match=missing): + validate_output_frames(frames, plan) diff --git a/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py b/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py index 840fdf0d8..d5fd4166a 100644 --- a/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py +++ b/projects/policyengine-simulation-executor/tests/test_stage12_runtime.py @@ -25,16 +25,20 @@ DatasetArtifactReference, DatasetPopulationInput, DatasetProvenance, + EntityOutputPlan, GeographySelection, + PlannedSimulationExecutionInput, ReportAggregate, ReportExecutionInput, - RequestedSimulationOutput, + ReportOutputRequirements, ResultComparisonStatus, RowIdentity, SimulationArtifactDescriptor, SimulationExecutionInput, SimulationRole, Stage12InvocationContext, + Stage12OutputPlan, + stage12_output_plan_sha256, ) from policyengine_simulation_executor.stage12_artifacts import ( @@ -44,7 +48,8 @@ from policyengine_simulation_executor.stage12_runtime import ( SimulationCalculation, _build_spm_result, - coordinate_report, + build_aggregate_report, + coordinate_report as coordinate_report_impl, run_single_simulation, simulation_input_sha256, ) @@ -88,20 +93,54 @@ def _simulation(role: SimulationRole) -> SimulationExecutionInput: ), year=2026, geography=GeographySelection(country="us", region="us"), - requested_output=RequestedSimulationOutput(variables=("*",)), bundle=_bundle(), ) +def _output_plan() -> Stage12OutputPlan: + return Stage12OutputPlan( + country="us", + requirements=ReportOutputRequirements( + aggregates=tuple(ReportAggregate), + include_cliff_impacts=False, + labor_supply_response_active=False, + ), + entities=( + EntityOutputPlan( + entity="household", + materialized_variables=("household_id", "household_net_income"), + ), + EntityOutputPlan( + entity="person", + materialized_variables=("age", "household_id", "person_id"), + ), + ), + ) + + +def _planned_simulation(role: SimulationRole) -> PlannedSimulationExecutionInput: + return PlannedSimulationExecutionInput.model_validate( + { + **_simulation(role).model_dump(mode="json"), + "output_plan": _output_plan().model_dump(mode="json"), + } + ) + + def _report() -> ReportExecutionInput: return ReportExecutionInput( evaluation_id=EVALUATION_ID, baseline=_simulation(SimulationRole.BASELINE), reform=_simulation(SimulationRole.REFORM), - requested_aggregates=(ReportAggregate.BUDGET,), + requested_aggregates=tuple(ReportAggregate), ) +def coordinate_report(*args, **kwargs): + kwargs.setdefault("output_plan_resolver", lambda _: _output_plan()) + return coordinate_report_impl(*args, **kwargs) + + def _context() -> Stage12InvocationContext: return Stage12InvocationContext( request_id="request-1", @@ -279,6 +318,7 @@ def add_simulation(self, simulation, frames=None): content_sha256=digest, size_bytes=len(payload), ), + output_plan_sha256=stage12_output_plan_sha256(_output_plan()), row_identity=row_identity, bundle=simulation.bundle, ) @@ -345,8 +385,14 @@ def test_calculator_uses_the_current_dataset_selection_contract(monkeypatch) -> dataset = object() received: dict[str, object] = {} + class ModelVersion: + def resolve_entity_variables(self, _simulation): + return {entity: list(frame.columns) for entity, frame in _frames().items()} + class Model: spm_config = None + extra_variables = {} + tax_benefit_model_version = ModelVersion() output_dataset = type( "OutputDataset", (), @@ -355,6 +401,7 @@ class Model: 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( @@ -396,7 +443,9 @@ def build_simulation(params, *, dataset, dataset_selection, **kwargs): monkeypatch.setattr(simulation_runtime, "_load_dataset", load_dataset) monkeypatch.setattr(simulation_runtime, "_build_simulation", build_simulation) - result = worker.calculate_simulation_frames(_simulation(SimulationRole.BASELINE)) + result = worker.calculate_simulation_frames( + _planned_simulation(SimulationRole.BASELINE) + ) assert "data" not in received["params"] assert "data_version" not in received["params"] @@ -406,12 +455,13 @@ def build_simulation(params, *, dataset, dataset_selection, **kwargs): assert received["build_dataset"] is dataset assert received["build_selection"] is selection assert received["ensured"] is True + assert received["extras_at_ensure"] == {"household": [], "person": []} assert set(result.frames) == {"household", "person"} def test_single_worker_accepts_one_policy_and_persists_one_artifact() -> None: store = FakeStore() - simulation = _simulation(SimulationRole.BASELINE) + simulation = _planned_simulation(SimulationRole.BASELINE) store.children[simulation.simulation_execution_id] = _child(simulation) artifacts = FakeArtifacts() @@ -435,7 +485,7 @@ def test_single_worker_accepts_one_policy_and_persists_one_artifact() -> None: def test_single_worker_retains_detached_calculation_provenance() -> None: store = FakeStore() - simulation = _simulation(SimulationRole.BASELINE) + simulation = _planned_simulation(SimulationRole.BASELINE) store.children[simulation.simulation_execution_id] = _child(simulation) artifacts = FakeArtifacts() provenance = {"receipt": {"version": 1}} @@ -457,7 +507,7 @@ def test_single_worker_retains_detached_calculation_provenance() -> None: def test_single_worker_exposes_only_a_bounded_failure() -> None: store = FakeStore() - simulation = _simulation(SimulationRole.BASELINE) + simulation = _planned_simulation(SimulationRole.BASELINE) store.children[simulation.simulation_execution_id] = _child(simulation) def fail(_simulation): @@ -482,6 +532,28 @@ def fail(_simulation): assert child.error_summary == "RuntimeError" +def test_single_worker_rejects_frames_that_do_not_satisfy_the_output_plan() -> None: + store = FakeStore() + simulation = _planned_simulation(SimulationRole.BASELINE) + store.children[simulation.simulation_execution_id] = _child(simulation) + frames = _frames() + frames["person"] = frames["person"].drop(columns=["age"]) + + with pytest.raises(RuntimeError, match="Stage 12 simulation execution failed"): + run_single_simulation( + simulation.model_dump(mode="json"), + _context().model_dump(mode="json"), + required_country="us", + store=store, + artifacts=FakeArtifacts(), + calculator=lambda _: frames, + ) + + assert store.children[simulation.simulation_execution_id].status is ( + ComparisonRunLifecycleStatus.FAILED + ) + + def test_aggregate_combines_detached_spm_receipts() -> None: selection = { "forecast_content_sha256": "f" * 64, @@ -552,8 +624,71 @@ def receipt(): assert len(result["spm_provenance"]["reform"]) == 1 +def test_aggregate_stand_ins_preserve_policy_and_cliff_options(monkeypatch) -> None: + from policyengine.outputs import labor_supply_response_is_active + from policyengine_simulation_executor import simulation_output_builder + + report = _report() + options = {"include_cliffs": True} + report = report.model_copy( + update={ + "baseline": report.baseline.model_copy(update={"options": options}), + "reform": report.reform.model_copy( + update={ + "options": options, + "policy": { + "gov.simulation.labor_supply_responses.elasticities.income": 0.1 + }, + } + ), + } + ) + frames = { + entity: pd.DataFrame({f"{entity}_id": [1]}) + for entity in ( + "person", + "marital_unit", + "family", + "spm_unit", + "tax_unit", + "household", + ) + } + observed = {} + + class CapturingBuilder: + def __init__(self, **values): + observed.update(values) + + def serialize(self): + return {"captured": True} + + monkeypatch.setattr( + simulation_output_builder, + "SimulationOutputBuilder", + CapturingBuilder, + ) + artifacts = FakeArtifacts() + + result = build_aggregate_report( + report=report, + baseline_frames=frames, + reform_frames=frames, + baseline_descriptor=artifacts.add_simulation(report.baseline), + reform_descriptor=artifacts.add_simulation(report.reform), + ) + + assert result["result"] == {"captured": True} + assert observed["simulation_params"]["include_cliffs"] is True + assert labor_supply_response_is_active( + observed["baseline"], + observed["reform"], + country_code="us", + ) + + def test_single_worker_rejects_combined_baseline_and_reform_input() -> None: - payload = _simulation(SimulationRole.BASELINE).model_dump(mode="json") + payload = _planned_simulation(SimulationRole.BASELINE).model_dump(mode="json") payload["baseline"] = {} payload["reform"] = {} @@ -607,23 +742,29 @@ def __init__( *, fail_role=None, incompatible=False, + wrong_plan_digest=False, + missing_column=False, production_result=None, ): self.artifacts = artifacts self.fail_role = fail_role self.incompatible = incompatible + self.wrong_plan_digest = wrong_plan_digest + self.missing_column = missing_column self.executor = ThreadPoolExecutor(max_workers=2) self.intervals = {} self.events = [] self.environments = [] self.production_result = production_result self.production_wait_timeouts = [] + self.output_plans = [] def spawn(self, *, simulation, environment, **_): - parsed = SimulationExecutionInput.model_validate(simulation) + parsed = PlannedSimulationExecutionInput.model_validate(simulation) role = parsed.role.value self.events.append(f"spawn:{role}") self.environments.append(environment) + self.output_plans.append(parsed.output_plan) return self._submit(parsed, role) @@ -636,7 +777,7 @@ def restore(self, invocation_id): self.production_wait_timeouts, ) role = invocation_id.removeprefix("call-") - parsed = _simulation(SimulationRole(role)) + parsed = _planned_simulation(SimulationRole(role)) self.events.append(f"restore:{role}") return self._submit(parsed, role) @@ -648,6 +789,10 @@ def run(): if role == self.fail_role: raise RuntimeError("child failed with sensitive values") frames = _frames(100.0 if role == "baseline" else 120.0) + if self.missing_column and role == "reform": + frames["household"] = frames["household"].drop( + columns=["household_net_income"] + ) descriptor = self.artifacts.add_simulation(parsed, frames) if self.incompatible and role == "reform": descriptor = descriptor.model_copy( @@ -659,6 +804,10 @@ def run(): ) } ) + if self.wrong_plan_digest and role == "reform": + descriptor = descriptor.model_copy( + update={"output_plan_sha256": "f" * 64} + ) self.intervals[role] = (started, time.monotonic()) return descriptor.model_dump(mode="json") @@ -685,7 +834,13 @@ def aggregate(**values): } ) parent = _parent().model_copy(update={"environment": "production"}) - result = coordinate_report( + resolver_calls = [] + + def resolver(report): + resolver_calls.append(report.evaluation_id) + return _output_plan() + + result = coordinate_report_impl( _report().model_dump(mode="json"), context.model_dump(mode="json"), parent.model_dump(mode="json"), @@ -695,10 +850,13 @@ def aggregate(**values): artifacts=artifacts, invoker=invoker, aggregator=aggregate, + output_plan_resolver=resolver, ) assert invoker.events == ["spawn:baseline", "spawn:reform"] assert invoker.environments == ["main", "main"] + assert resolver_calls == [EVALUATION_ID] + assert invoker.output_plans == [_output_plan(), _output_plan()] latest_start = max(interval[0] for interval in invoker.intervals.values()) earliest_end = min(interval[1] for interval in invoker.intervals.values()) assert latest_start < earliest_end @@ -905,6 +1063,36 @@ def test_coordinator_never_writes_partial_aggregate_when_a_child_fails() -> None assert failed_child.error_summary == "RuntimeError" +@pytest.mark.parametrize( + "invoker", + [ + lambda artifacts: ConcurrentInvoker(artifacts, wrong_plan_digest=True), + lambda artifacts: ConcurrentInvoker(artifacts, missing_column=True), + ], +) +def test_coordinator_rejects_artifacts_that_do_not_satisfy_the_shared_plan( + invoker, +) -> None: + store = FakeStore() + artifacts = FakeArtifacts() + + with pytest.raises(RuntimeError, match="Stage 12 report coordination failed"): + coordinate_report( + _report().model_dump(mode="json"), + _context().model_dump(mode="json"), + _parent().model_dump(mode="json"), + application_name=_context().modal_application, + coordinator_invocation_id="coordinator-1", + store=store, + artifacts=artifacts, + invoker=invoker(artifacts), + aggregator=lambda **_: {"must": "not run"}, + ) + + assert artifacts.aggregate_writes == [] + assert store.parent.status is ComparisonRunLifecycleStatus.FAILED + + def test_coordinator_records_failure_when_child_persistence_fails() -> None: class FailingChildStore(FakeStore): def create_or_resolve_simulation(self, record): @@ -1027,7 +1215,7 @@ def __init__(self): self.events = [] def spawn(self, *, simulation, **_): - role = SimulationExecutionInput.model_validate(simulation).role.value + role = PlannedSimulationExecutionInput.model_validate(simulation).role.value self.events.append(f"spawn:{role}") return TimeoutCall(f"call-{role}")