diff --git a/src/spatialdata/_core/query/relational_query.py b/src/spatialdata/_core/query/relational_query.py index 7ef7c1a0..98d92a0a 100644 --- a/src/spatialdata/_core/query/relational_query.py +++ b/src/spatialdata/_core/query/relational_query.py @@ -236,6 +236,9 @@ def _get_masked_element( mask_values = left_index[mask] else: mask_values = left_index + elif mask_values is not None: + order_mask = np.isin(element_indices, mask_values) + mask_values = np.asarray(element_indices)[order_mask] if isinstance(element, DaskDataFrame): return element.map_partitions(lambda df: df.loc[mask_values], meta=element) @@ -252,6 +255,13 @@ def _right_exclusive_join_spatialelement_table( match_rows: Literal["left", "no", "right"], filter_label_pixels: bool | None = None, ) -> tuple[dict[str, Any], AnnData | None]: + if match_rows == "left": + warnings.warn( + "Matching rows 'left' is not supported for 'right_exclusive' join; it will be treated as 'no'.", + UserWarning, + stacklevel=2, + ) + match_rows = "no" regions, region_column_name, instance_key = get_table_keys(table) if isinstance(regions, str): regions = [regions] @@ -297,7 +307,12 @@ def _right_join_spatialelement_table( filter_label_pixels: bool | None = None, ) -> tuple[dict[str, Any], AnnData]: if match_rows == "left": - warnings.warn("Matching rows 'left' is not supported for 'right' join.", UserWarning, stacklevel=2) + warnings.warn( + "Matching rows 'left' is not supported for 'right' join; it will be treated as 'no'.", + UserWarning, + stacklevel=2, + ) + match_rows = "no" regions, region_column_name, instance_key = get_table_keys(table) if isinstance(regions, str): regions = [regions] @@ -383,6 +398,14 @@ def _inner_join_spatialelement_table( if joined_indices is not None: joined_indices = joined_indices.dropna() if any(joined_indices.isna()) else joined_indices + # `groupby(region)` above collects the matching table rows grouped by region, which does not + # preserve the original `table.obs` row order when a table annotates multiple interleaved + # regions. For `match_rows="no"` there is no element-driven ordering to honor, and for + # `match_rows="right"` the table's own row order takes priority (only `match_rows="left"` lets the + # element's row order override it), so in both cases restore the original table row order, as + # would be expected for a semi-join. + if match_rows in ("no", "right"): + joined_indices = joined_indices.sort_values() joined_table = table[joined_indices.tolist(), :].copy() if joined_indices is not None else None @@ -401,6 +424,13 @@ def _left_exclusive_join_spatialelement_table( match_rows: Literal["left", "no", "right"], filter_label_pixels: bool | None = None, ) -> tuple[dict[str, Any], AnnData | None]: + if match_rows == "right": + warnings.warn( + "Matching rows 'right' is not supported for 'left_exclusive' join; it will be treated as 'no'.", + UserWarning, + stacklevel=2, + ) + match_rows = "no" regions, region_column_name, instance_key = get_table_keys(table) if isinstance(regions, str): regions = [regions] @@ -411,8 +441,7 @@ def _left_exclusive_join_spatialelement_table( group_df = groups_df.get_group(name) table_instance_key_column = group_df[instance_key] if element_type in ["points", "shapes"]: - mask = np.full(len(element), True, dtype=bool) - mask[table_instance_key_column.values] = False + mask = ~np.isin(element.index, table_instance_key_column.values) masked_element = element.loc[mask, :] if mask.sum() != 0 else None element_dict[element_type][name] = masked_element else: @@ -438,7 +467,12 @@ def _left_join_spatialelement_table( filter_label_pixels: bool | None = None, ) -> tuple[dict[str, Any], AnnData]: if match_rows == "right": - warnings.warn("Matching rows 'right' is not supported for 'left' join.", UserWarning, stacklevel=2) + warnings.warn( + "Matching rows 'right' is not supported for 'left' join; it will be treated as 'no'.", + UserWarning, + stacklevel=2, + ) + match_rows = "no" regions, region_column_name, instance_key = get_table_keys(table) if isinstance(regions, str): regions = [regions] @@ -469,6 +503,12 @@ def _left_join_spatialelement_table( # if nan were present, the dtype would have been changed to float if joined_indices.dtype == float: joined_indices = joined_indices.astype(int) + # `groupby(region)` above collects the matching table rows grouped by region, which does not + # preserve the original `table.obs` row order when a table annotates multiple interleaved + # regions. For `match_rows="no"` there is no element-driven ordering to honor, so + # restore the original table row order, as would be expected for a semi-join. + if match_rows == "no": + joined_indices = joined_indices.sort_values() joined_table = table[joined_indices.tolist(), :].copy() if joined_indices is not None else None _inplace_fix_subset_categorical_obs(subset_adata=joined_table, original_adata=table) if joined_table is not None: diff --git a/tests/core/query/test_relational_query.py b/tests/core/query/test_relational_query.py index 4f87098a..21e04a37 100644 --- a/tests/core/query/test_relational_query.py +++ b/tests/core/query/test_relational_query.py @@ -1,5 +1,9 @@ from __future__ import annotations +import re +import warnings +from dataclasses import dataclass, field + import annsel as an import numpy as np import pandas as pd @@ -940,6 +944,189 @@ def test_filter_table_categorical_bug(shapes): shapes.filter_by_coordinate_system("global") +@dataclass +class _JoinOutcome: + """Expected outcome of a `join_spatialelement_table()` call, for a given `how`/`match_rows` pair.""" + + # expected `joined_table.obs["label"]`, in order; `None` when no table is expected to be returned + table_order: list[str] | None = None + # whether a "Matching rows '<...>' is not supported for '<...>' join." UserWarning is expected to be emitted + warns: bool = False + # expected values of `element_dict[name].index` for element name in {"a", "b"}; `None` for a element name whose + # returned join result is expected to be `None` (e.g. fully excluded, or not returned by this join type) + element_index: dict[str, list[int] | None] = field(default_factory=dict) + + +def _make_interleaved_regions_sdata() -> tuple[SpatialData, dict[str, dict[str, _JoinOutcome]]]: + from geopandas import GeoDataFrame + from shapely.geometry import Point + + from spatialdata.models import ShapesModel + + def circles(indices): + # `indices` gives both the number of circles and the (non-default) row order of the element. + gdf = GeoDataFrame( + {"geometry": [Point(i, i) for i in range(len(indices))], "radius": [1.0] * len(indices)}, + index=pd.Index(indices), + ) + return ShapesModel.parse(gdf) + + # assumptions/comments: + # - no duplicate values in the index of each spatial element (duplicate values are tested elsewhere) + # - no duplicate values for the instance_key column of the table (duplicate values are tested elsewhere) + # + # edge cases being tested: + # - instance_id values are non-monotonic + # - the index in each spatial element is non-monotonic + # - we also set the index of the table obs to random values; these should be ignored (in the code we call .index on + # a region_key column, but the index is freshly reset by a nearby call of .reset_index() inside the join + # machinery) + # - "b3" and "c7" are unmatched table rows: "b3" refers to a missing instance in a + # spatial element that exists, while "c7" refers to a region with no spatial element + obs = pd.DataFrame( + { + "region": pd.Categorical(["b", "b", "a", "b", "a", "a", "b", "c"]), + "instance_id": [2, 1, 2, 3, 1, 0, 0, 7], + "label": ["b2", "b1", "a2", "b3", "a1", "a0", "b0", "c7"], + }, + index=np.random.default_rng(0).integers(0, 3, size=8).astype(str), + ) + # "a" additionally has unmatched instance ids 5, 4, and "b" has 4, 6. + # These test that unmatched element rows are preserved in element order by + # "left" and "left_exclusive" joins. + shapes = {"a": circles([2, 1, 0, 5, 4]), "b": circles([1, 2, 0, 4, 6])} + # to make understanding easier, you may want to refer to the figure on joins from the docs: + # https://spatialdata.scverse.org/en/stable/tutorials/notebooks/notebooks/examples/tables.html + expected = { + "left": { + "no": _JoinOutcome( + table_order=["b2", "b1", "a2", "a1", "a0", "b0"], + element_index={"a": [2, 1, 0, 5, 4], "b": [1, 2, 0, 4, 6]}, + ), + "left": _JoinOutcome( + table_order=["a2", "a1", "a0", "b1", "b2", "b0"], + element_index={"a": [2, 1, 0, 5, 4], "b": [1, 2, 0, 4, 6]}, + ), + "right": _JoinOutcome( + table_order=["b2", "b1", "a2", "a1", "a0", "b0"], + warns=True, + element_index={"a": [2, 1, 0, 5, 4], "b": [1, 2, 0, 4, 6]}, + ), + }, + "left_exclusive": { + # by design, "left_exclusive" never returns a table (only filtered elements), regardless of + # match_rows or whether anything was actually excluded. + "no": _JoinOutcome(table_order=None, element_index={"a": [5, 4], "b": [4, 6]}), + "left": _JoinOutcome(table_order=None, element_index={"a": [5, 4], "b": [4, 6]}), + "right": _JoinOutcome(table_order=None, warns=True, element_index={"a": [5, 4], "b": [4, 6]}), + }, + "inner": { + "no": _JoinOutcome( + table_order=["b2", "b1", "a2", "a1", "a0", "b0"], + element_index={"a": [2, 1, 0], "b": [1, 2, 0]}, + ), + "left": _JoinOutcome( + table_order=["a2", "a1", "a0", "b1", "b2", "b0"], element_index={"a": [2, 1, 0], "b": [1, 2, 0]} + ), + "right": _JoinOutcome( + table_order=["b2", "b1", "a2", "a1", "a0", "b0"], element_index={"a": [2, 1, 0], "b": [2, 1, 0]} + ), + }, + "right": { + "no": _JoinOutcome( + table_order=["b2", "b1", "a2", "b3", "a1", "a0", "b0", "c7"], + element_index={"a": [2, 1, 0], "b": [1, 2, 0]}, + ), + "left": _JoinOutcome( + table_order=["b2", "b1", "a2", "b3", "a1", "a0", "b0", "c7"], + warns=True, + element_index={"a": [2, 1, 0], "b": [1, 2, 0]}, + ), + "right": _JoinOutcome( + table_order=["b2", "b1", "a2", "b3", "a1", "a0", "b0", "c7"], + element_index={"a": [2, 1, 0], "b": [2, 1, 0]}, + ), + }, + "right_exclusive": { + "no": _JoinOutcome(table_order=["b3", "c7"], element_index={"a": None, "b": None}), + "left": _JoinOutcome(table_order=["b3", "c7"], warns=True, element_index={"a": None, "b": None}), + "right": _JoinOutcome(table_order=["b3", "c7"], element_index={"a": None, "b": None}), + }, + } + + table = TableModel.parse( + AnnData(X=np.zeros((len(obs), 1)), obs=obs), + region=["a", "b", "c"], + region_key="region", + instance_key="instance_id", + ) + sdata = SpatialData(shapes=shapes, tables={"table": table}) + return sdata, expected + + +@pytest.mark.parametrize("match_rows", ["no", "left", "right"]) +@pytest.mark.parametrize( + "how", + [ + "left", + "left_exclusive", + "inner", + "right", + pytest.param( + "right_exclusive", + marks=pytest.mark.xfail( + reason="known bug (see https://github.com/scverse/spatialdata/issues/1162): 'right_exclusive' join " + "drops unmatched table rows belonging to a region with no queried spatial element (e.g. 'c7')", + strict=True, + ), + ), + ], +) +def test_join_preserves_row_order_multiple_interleaved_regions(how, match_rows): + # generalization to all the join types of the bug reported in https://github.com/scverse/spatialdata/issues/1162 + # covering all `how` values of `join_spatialelement_table`, crossed with all values of `match_rows`, and checking + # whether the row orders of the returned spatial elements and table are correct and if the "match_rows not + # supported" UserWarning is (or isn't) actually raised (see `_make_interleaved_regions_sdata` and `_JoinOutcome`). + sdata, expected_by_how_and_match_rows = _make_interleaved_regions_sdata() + outcome = expected_by_how_and_match_rows[how][match_rows] + + with warnings.catch_warnings(record=True) as record: + warnings.simplefilter("always") + element_dict, joined_table = join_spatialelement_table( + sdata=sdata, + spatial_element_names=["a", "b"], + table_name="table", + how=how, + match_rows=match_rows, + ) + + # other UserWarnings can also fire here (e.g. anndata's "Observation names are not unique", triggered by + # the fixture's duplicated obs_names), so only look for the one this test is actually about. The message + # looks like "Matching rows 'right' is not supported for 'left_exclusive' join; it will be treated as 'no'.", + # with the two quoted values varying by `match_rows` / `how`. + unsupported_match_rows_re = re.compile( + r"Matching rows '[^']+' is not supported for '[^']+' join; it will be treated as 'no'\." + ) + unsupported_match_rows_warnings = [ + w for w in record if issubclass(w.category, UserWarning) and unsupported_match_rows_re.search(str(w.message)) + ] + assert bool(unsupported_match_rows_warnings) == outcome.warns + + if outcome.table_order is None: + assert joined_table is None + else: + assert joined_table is not None + assert list(joined_table.obs["label"]) == outcome.table_order + + for name, expected_index in outcome.element_index.items(): + actual_element = element_dict[name] + if expected_index is None: + assert actual_element is None + else: + assert actual_element is not None + assert list(actual_element.index) == expected_index + + def test_filter_table_non_annotating(full_sdata): obs = pd.DataFrame({"test": ["a", "b", "c"]}, index=list(map(str, range(3)))) adata = AnnData(obs=obs)