Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
1 change: 1 addition & 0 deletions changelog.d/564.fixed.md
Original file line number Diff line number Diff line change
@@ -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.
29 changes: 29 additions & 0 deletions docs/engineering/skills/repository-guidance.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
32 changes: 32 additions & 0 deletions docs/regions.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 4 additions & 2 deletions src/policyengine/core/scoping_strategy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
26 changes: 21 additions & 5 deletions src/policyengine/countries/uk/regions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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,
)
)

Expand Down
21 changes: 15 additions & 6 deletions src/policyengine/utils/entity_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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).
Expand All @@ -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.
Expand Down Expand Up @@ -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.
Expand All @@ -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.

Expand All @@ -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
57 changes: 57 additions & 0 deletions tests/test_scoping_strategy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
Loading
Loading