diff --git a/AGENTS.md b/AGENTS.md index d96028cf..bd4f88aa 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -33,3 +33,16 @@ round-trip is a documented property of the public surface. Do not cache arbitrary Python objects in public result structures. The `core.Simulation` output must stay serialisable. + +## UK country filtering + +Select Scotland, Wales, and Northern Ireland through the stored household +`region` column: `SCOTLAND`, `WALES`, and `NORTHERN_IRELAND`, respectively. +Never require a `country` column when filtering raw UK household inputs; +`country` is derived by the UK model after loading the data. The public selector +`country/scotland` describes the requested area, not an input column name. + +England must select the nine English `region` values, not `ENGLAND` and not +the complement of the other countries. Preserve these rules in new filters, +refactors, and tests. Read the UK country filtering section in +`docs/engineering/skills/repository-guidance.md` before changing this behavior. diff --git a/changelog.d/564.fixed.md b/changelog.d/564.fixed.md new file mode 100644 index 00000000..5f850edd --- /dev/null +++ b/changelog.d/564.fixed.md @@ -0,0 +1 @@ +Filter UK country simulations using stored household regions, selecting all nine English regions for England and preserving entity relationships and weights without requiring a derived country column. diff --git a/docs/engineering/skills/repository-guidance.md b/docs/engineering/skills/repository-guidance.md index 6400ecbb..10b8a57b 100644 --- a/docs/engineering/skills/repository-guidance.md +++ b/docs/engineering/skills/repository-guidance.md @@ -52,6 +52,35 @@ representative-data runs unless the change specifically needs that coverage. When a test needs country package data, make the dependency explicit and skip cleanly if credentials or local artifacts are unavailable. +## UK country filtering + +Raw UK household data stores `region`. The UK model derives `country` from +that input during calculation. Filtering precedes calculation, so agents and +developers must use `region` for country selection even though the registry's +public identifiers start with `country/`. + +| Public selector | Required input filter | +| --- | --- | +| `country/scotland` | `region == "SCOTLAND"` | +| `country/wales` | `region == "WALES"` | +| `country/northern_ireland` | `region == "NORTHERN_IRELAND"` | +| `country/england` | `region` is one of the nine English values below | + +The English values are `NORTH_EAST`, `NORTH_WEST`, `YORKSHIRE`, `EAST_MIDLANDS`, +`WEST_MIDLANDS`, `EAST_OF_ENGLAND`, `LONDON`, `SOUTH_EAST`, and `SOUTH_WEST`. +There is no stored `ENGLAND` region value. Use explicit membership so unknown +or missing regions do not become English households. + +Keep all four country entries as `RowFilterStrategy` so the simulation API can +combine them with `RegionGroupStrategy`. Do not add a dataset `country` column +or run a preliminary model calculation to work around an incorrect filter. +Preserve household membership, related people/benefit units, and weights. + +Regression tests must exercise Scotland, Wales, Northern Ireland, and England +against household inputs with `region` and **no `country` column**, including +the string and byte representations accepted by the dataset loader. See +`tests/test_uk_regions.py` and the public explanation in `docs/regions.md`. + ## Anti-Patterns - Do not bypass the wrapper layer without a clear reason. diff --git a/docs/regions.md b/docs/regions.md index 6d703c9f..3371b76f 100644 --- a/docs/regions.md +++ b/docs/regions.md @@ -54,6 +54,38 @@ for row in impacts.district_results: at-large districts use `00`). Congressional district regions load the same ACS-local file and scope its rows by `congressional_district_geoid`. +## UK countries + +UK country simulations filter the national dataset's stored household `region` +column. The UK model derives `country` from `region`; the input dataset does not +need a separate `country` column. + +| Country selector | Stored household input selection | +| --- | --- | +| `country/scotland` | `region == "SCOTLAND"` | +| `country/wales` | `region == "WALES"` | +| `country/northern_ireland` | `region == "NORTHERN_IRELAND"` | + +The `country/` prefix identifies a public geographic selector; it does not name +a dataset column. Always filter these countries through `region` before running +the simulation. + +Each country uses `RowFilterStrategy`. England matches a list of nine ITL1 regions: +`NORTH_EAST`, `NORTH_WEST`, `YORKSHIRE`, `EAST_MIDLANDS`, `WEST_MIDLANDS`, +`EAST_OF_ENGLAND`, `LONDON`, `SOUTH_EAST`, and `SOUTH_WEST`. Scotland, Wales and +Northern Ireland each match their corresponding stored value: `SCOTLAND`, +`WALES`, or `NORTHERN_IRELAND`. Unknown regions are excluded +from all four countries. These filters preserve the selected households' +people, benefit units and weights. + +```python +scotland = pe.uk.model.get_region("country/scotland") +england = pe.uk.model.get_region("country/england") +``` + +Pass the registry entry's `scoping_strategy` to `Simulation` when using a +national UK dataset. + ## UK parliamentary constituencies Constituency-level impacts group household output rows by the longwise diff --git a/src/policyengine/core/scoping_strategy.py b/src/policyengine/core/scoping_strategy.py index 3cdd3f1e..cd54cf7d 100644 --- a/src/policyengine/core/scoping_strategy.py +++ b/src/policyengine/core/scoping_strategy.py @@ -24,6 +24,7 @@ from pydantic import BaseModel, Discriminator, Field from policyengine.utils.entity_utils import ( + HouseholdFilterValue, filter_dataset_by_household_ids, filter_dataset_by_household_variable, matching_household_ids, @@ -69,12 +70,13 @@ class RowFilterStrategy(RegionScopingStrategy): """Scoping strategy that filters dataset rows by a household variable. Used for regions where we want to keep only households matching a - specific variable value (e.g., US states or congressional districts). + specific variable value or any value in a list (e.g., the nine English + regions). Additional filters must also match. """ strategy_type: Literal["row_filter"] = "row_filter" variable_name: str - variable_value: Union[str, int, float] + variable_value: HouseholdFilterValue additional_filters: dict[str, Union[str, int, float]] = Field(default_factory=dict) def apply( diff --git a/src/policyengine/countries/uk/regions.py b/src/policyengine/countries/uk/regions.py index e6ea8c14..f74dd1f2 100644 --- a/src/policyengine/countries/uk/regions.py +++ b/src/policyengine/countries/uk/regions.py @@ -39,6 +39,21 @@ "northern_ireland": "Northern Ireland", } +# Stored names from policyengine_uk's Region enum. England spans nine ITL1 +# regions; each other UK country is itself one ITL1 region. +# https://www.ons.gov.uk/methodology/geography/ukgeographies/eurostat +ENGLISH_REGIONS = ( + "NORTH_EAST", + "NORTH_WEST", + "YORKSHIRE", + "EAST_MIDLANDS", + "WEST_MIDLANDS", + "EAST_OF_ENGLAND", + "LONDON", + "SOUTH_EAST", + "SOUTH_WEST", +) + def _load_constituencies_from_csv() -> list[dict]: """Load UK constituency data from CSV. @@ -135,18 +150,19 @@ def build_uk_region_registry( ) ) - # 2. Country regions (filter from national by 'country' variable) + # 2. Countries scope the stored region input; country is a derived variable. for code, name in UK_COUNTRIES.items(): + scoping_strategy = RowFilterStrategy( + variable_name="region", + variable_value=list(ENGLISH_REGIONS) if code == "england" else code.upper(), + ) regions.append( Region( code=f"country/{code}", label=name, region_type="country", parent_code="uk", - scoping_strategy=RowFilterStrategy( - variable_name="country", - variable_value=code.upper(), - ), + scoping_strategy=scoping_strategy, ) ) diff --git a/src/policyengine/utils/entity_utils.py b/src/policyengine/utils/entity_utils.py index 3f8b54e7..d06f662b 100644 --- a/src/policyengine/utils/entity_utils.py +++ b/src/policyengine/utils/entity_utils.py @@ -8,6 +8,8 @@ logger = logging.getLogger(__name__) +HouseholdFilterValue = Union[str, int, float, list[Union[str, int, float]]] + def _resolve_id_column(person_data: pd.DataFrame, entity_name: str) -> str: """Resolve the ID column name for a group entity in person data. @@ -55,10 +57,10 @@ def build_entity_relationships( def _household_mask( household_data: pd.DataFrame, variable_name: str, - variable_value: Union[str, int, float], + variable_value: HouseholdFilterValue, additional_filters: Optional[dict[str, Union[str, int, float]]] = None, ): - """Boolean mask over the household table for a single-value filter. + """Boolean mask matching one value or any value in a list. Local intermediate only — never crosses a function boundary — so no positional-alignment invariant leaks out (callers key on household_id). @@ -85,7 +87,7 @@ def _household_mask( def matching_household_ids( entity_data: dict[str, MicroDataFrame], variable_name: str, - variable_value: Union[str, int, float], + variable_value: HouseholdFilterValue, additional_filters: Optional[dict[str, Union[str, int, float]]] = None, ) -> set: """Return the set of ``household_id`` values matching a household filter. @@ -163,7 +165,7 @@ def filter_dataset_by_household_variable( entity_data: dict[str, MicroDataFrame], group_entities: list[str], variable_name: str, - variable_value: Union[str, int, float], + variable_value: HouseholdFilterValue, additional_filters: Optional[dict[str, Union[str, int, float]]] = None, ) -> dict[str, MicroDataFrame]: """Filter dataset entities to only include households matching variables. @@ -177,7 +179,8 @@ def filter_dataset_by_household_variable( (from YearData.entity_data). group_entities: List of group entity names for this country. variable_name: The household-level variable to filter on. - variable_value: The value to match. Handles both str and bytes encoding. + variable_value: One value or a list of allowed values. String values + match both str and bytes encoding. additional_filters: Optional household-level filters that must also match, keyed by variable name. @@ -201,7 +204,13 @@ def filter_dataset_by_household_variable( ) -def _values_match(values, expected: Union[str, int, float]): +def _values_match(values, expected: HouseholdFilterValue): + if isinstance(expected, list): + allowed = [ + *expected, + *(value.encode() for value in expected if isinstance(value, str)), + ] + return pd.Series(values).isin(allowed).to_numpy(copy=True) if isinstance(expected, str): return (values == expected) | (values == expected.encode()) return values == expected diff --git a/tests/test_scoping_strategy.py b/tests/test_scoping_strategy.py index 3e8d94b9..d6704f7a 100644 --- a/tests/test_scoping_strategy.py +++ b/tests/test_scoping_strategy.py @@ -15,6 +15,63 @@ from policyengine.core.simulation import Simulation +@pytest.mark.parametrize( + "values,selected,expected", + [ + (["LONDON", b"NORTH_WEST", "SCOTLAND"], ["LONDON", "NORTH_WEST"], [0, 1]), + ([6, 34, 36], [6, 36], [0, 2]), + ([6, "6", b"6"], [6], [0]), + ([6, "6", b"6"], ["6"], [1, 2]), + (["LONDON", "SCOTLAND", "WALES"], [], []), + ], +) +def test_row_filter_matches_any_list_value(values, selected, expected): + household = MicroDataFrame( + pd.DataFrame( + { + "household_id": [0, 1, 2], + "household_weight": [10.0, 20.0, 30.0], + "geography": values, + } + ), + weights="household_weight", + ) + person = MicroDataFrame( + pd.DataFrame({"person_id": [0, 1, 2], "person_household_id": [0, 1, 2]}) + ) + strategy = RowFilterStrategy(variable_name="geography", variable_value=selected) + data = {"household": household, "person": person} + + if not expected: + with pytest.raises(ValueError, match="No households found"): + strategy.apply(data, ["household"], 2026) + else: + result = strategy.apply(data, ["household"], 2026) + assert result["household"]["household_id"].tolist() == expected + assert result["person"]["person_id"].tolist() == expected + + +def test_row_filter_list_values_respect_additional_filters(us_entity_data): + strategy = RowFilterStrategy( + variable_name="state_fips", + variable_value=[6, 34], + additional_filters={"congressional_district_geoid": 601}, + ) + + result = strategy.apply( + us_entity_data, + ["household", "tax_unit", "spm_unit", "family", "marital_unit"], + 2026, + ) + + assert result["household"]["household_id"].tolist() == [1] + + +def test_row_filter_scalar_cache_key_is_unchanged(): + strategy = RowFilterStrategy(variable_name="region", variable_value="SCOTLAND") + assert strategy.cache_key == "row_filter:region=SCOTLAND" + + class TestRowFilterStrategy: """Tests for RowFilterStrategy.""" diff --git a/tests/test_uk_regions.py b/tests/test_uk_regions.py index 8327cee2..f0cd3cc9 100644 --- a/tests/test_uk_regions.py +++ b/tests/test_uk_regions.py @@ -2,8 +2,15 @@ from unittest.mock import patch +import pandas as pd +import pytest +from microdf import MicroDataFrame +from pydantic import TypeAdapter + from policyengine.core.scoping_strategy import ( + RegionGroupStrategy, RowFilterStrategy, + ScopingStrategy, ) from policyengine.countries.uk.regions import ( UK_COUNTRIES, @@ -91,7 +98,7 @@ def test__given_uk_registry__then_has_four_country_regions(self): def test__given_england_region__then_filters_from_national(self): """Given: England country region When: Checking its properties - Then: Filters from national with country field + Then: Filters from national using the nine English regions """ # When england = uk_region_registry.get("country/england") @@ -102,22 +109,24 @@ def test__given_england_region__then_filters_from_national(self): assert england.region_type == "country" assert england.parent_code == "uk" assert england.requires_filter - assert england.scoping_strategy.variable_name == "country" - assert england.scoping_strategy.variable_value == "ENGLAND" + assert isinstance(england.scoping_strategy, RowFilterStrategy) + assert england.scoping_strategy.variable_name == "region" + assert len(england.scoping_strategy.variable_value) == 9 assert england.dataset_path is None - def test__given_country_regions__then_have_row_filter_strategy(self): + def test__given_country_regions__then_filter_stored_regions(self): """Given: UK country regions When: Checking their scoping strategies - Then: Each has a RowFilterStrategy with correct variable/value + Then: Each country retains a row filter on stored regions """ - for code, name in UK_COUNTRIES.items(): + for code in UK_COUNTRIES: region = uk_region_registry.get(f"country/{code}") assert region is not None assert region.scoping_strategy is not None assert isinstance(region.scoping_strategy, RowFilterStrategy) - assert region.scoping_strategy.variable_name == "country" - assert region.scoping_strategy.variable_value == code.upper() + assert region.scoping_strategy.variable_name == "region" + if code != "england": + assert region.scoping_strategy.variable_value == code.upper() def test__given_scotland_region__then_filters_from_national(self): """Given: Scotland country region @@ -292,3 +301,153 @@ def test__given_local_authorities_included__then_filters_on_dataset_geography( assert isinstance(local_authority.scoping_strategy, RowFilterStrategy) assert local_authority.scoping_strategy.variable_name == "la_code_oa" assert local_authority.scoping_strategy.variable_value == "LA001" + + +# The pinned UK model's Region enum contains nine English ITL1 regions and +# Scotland, Wales and Northern Ireland. Country is derived, not a stored input. +_REGION_COUNTRIES = [ + ("NORTH_EAST", "england"), + ("NORTH_WEST", "england"), + ("YORKSHIRE", "england"), + ("EAST_MIDLANDS", "england"), + ("WEST_MIDLANDS", "england"), + ("EAST_OF_ENGLAND", "england"), + ("LONDON", "england"), + ("SOUTH_EAST", "england"), + ("SOUTH_WEST", "england"), + ("SCOTLAND", "scotland"), + ("WALES", "wales"), + ("NORTHERN_IRELAND", "northern_ireland"), + ("UNKNOWN", None), +] + + +@pytest.fixture +def region_only_uk_data(): + """Raw UK entities without a derived country column, with distinct weights.""" + household_ids = list(range(1, len(_REGION_COUNTRIES) + 1)) + household = pd.DataFrame( + { + "household_id": household_ids, + "household_weight": [100.0 + hid for hid in household_ids], + "region": [region for region, _ in _REGION_COUNTRIES], + } + ).iloc[::-1] + person = pd.DataFrame( + { + "person_id": [ + hid * 10 + offset for hid in household_ids for offset in (1, 2) + ], + "person_household_id": [hid for hid in household_ids for _ in (1, 2)], + "person_benunit_id": [hid * 100 for hid in household_ids for _ in (1, 2)], + "person_weight": [100.0 + hid for hid in household_ids for _ in (1, 2)], + } + ) + benunit = pd.DataFrame( + { + "benunit_id": [hid * 100 for hid in household_ids], + "benunit_weight": [100.0 + hid for hid in household_ids], + } + ) + return { + name: MicroDataFrame(frame, weights=f"{name}_weight") + for name, frame in ( + ("person", person), + ("benunit", benunit), + ("household", household), + ) + } + + +@pytest.mark.parametrize("country", UK_COUNTRIES) +@pytest.mark.parametrize("encoding", ["str", "bytes", "category"]) +def test_country_scoping_uses_raw_regions_and_preserves_entities( + country, encoding, region_only_uk_data +): + household = region_only_uk_data["household"] + if encoding == "bytes": + household["region"] = household["region"].map(str.encode) + elif encoding == "category": + household["region"] = household["region"].astype("category") + originals = { + entity: pd.DataFrame(frame).copy(deep=True) + for entity, frame in region_only_uk_data.items() + } + expected_households = { + index + for index, (_, expected_country) in enumerate(_REGION_COUNTRIES, start=1) + if expected_country == country + } + + strategy = uk_region_registry.get(f"country/{country}").scoping_strategy + result = strategy.apply(region_only_uk_data, ["benunit", "household"], 2026) + + for entity, id_column, ids in ( + ("household", "household_id", expected_households), + ("person", "person_household_id", expected_households), + ("benunit", "benunit_id", {hid * 100 for hid in expected_households}), + ): + expected = originals[entity][originals[entity][id_column].isin(ids)] + pd.testing.assert_frame_equal( + pd.DataFrame(result[entity]), expected.reset_index(drop=True) + ) + assert result[entity].weights.tolist() == expected[f"{entity}_weight"].tolist() + pd.testing.assert_frame_equal( + pd.DataFrame(region_only_uk_data[entity]), originals[entity] + ) + + +def test_england_scoping_allows_unrepresented_english_regions(region_only_uk_data): + """Missing English regions must not prevent selecting those present.""" + # Only North West (household 2) and Scotland (household 10) are represented. + selected = { + entity: MicroDataFrame( + pd.DataFrame(frame)[frame[column].isin(ids)], weights=f"{entity}_weight" + ) + for entity, frame, column, ids in ( + ("household", region_only_uk_data["household"], "household_id", [2, 10]), + ("person", region_only_uk_data["person"], "person_household_id", [2, 10]), + ("benunit", region_only_uk_data["benunit"], "benunit_id", [200, 1000]), + ) + } + + result = uk_region_registry.get("country/england").scoping_strategy.apply( + selected, ["benunit", "household"], 2026 + ) + + assert result["household"]["household_id"].tolist() == [2] + assert result["person"]["person_id"].tolist() == [21, 22] + assert result["benunit"]["benunit_id"].tolist() == [200] + + +@pytest.mark.parametrize("country", UK_COUNTRIES) +def test_country_scoping_json_round_trip(country, region_only_uk_data): + strategy = uk_region_registry.get(f"country/{country}").scoping_strategy + restored = TypeAdapter(ScopingStrategy).validate_json(strategy.model_dump_json()) + + assert restored == strategy + assert restored.cache_key == strategy.cache_key + expected = strategy.apply(region_only_uk_data, ["benunit", "household"], 2026) + actual = restored.apply(region_only_uk_data, ["benunit", "household"], 2026) + for entity in expected: + pd.testing.assert_frame_equal( + pd.DataFrame(actual[entity]), pd.DataFrame(expected[entity]) + ) + + +def test_country_filters_can_be_combined_without_duplicate_households( + region_only_uk_data, +): + """The worker composes country groups from RowFilterStrategy members.""" + members = [ + uk_region_registry.get(f"country/{country}").scoping_strategy + for country in ("england", "scotland", "england") + ] + assert all(isinstance(member, RowFilterStrategy) for member in members) + + result = RegionGroupStrategy(members=members).apply( + region_only_uk_data, ["benunit", "household"], 2026 + ) + + assert set(result["household"]["household_id"]) == set(range(1, 11)) + assert result["household"]["household_id"].is_unique