Skip to content

Commit 832f4b2

Browse files
Fix: add validation to tables for relational queries (#1160)
fix: return type of region metadata from join restored to released behavior
1 parent 471ae71 commit 832f4b2

2 files changed

Lines changed: 14 additions & 6 deletions

File tree

src/spatialdata/_core/query/relational_query.py

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -242,6 +242,10 @@ def _get_masked_element(
242242
return element.loc[mask_values, :]
243243

244244

245+
def _region_as_str_if_list_of_len_one(region: list[str]) -> str | list[str]:
246+
return region if len(region) > 1 else region[0]
247+
248+
245249
def _right_exclusive_join_spatialelement_table(
246250
element_dict: dict[str, dict[str, Any]],
247251
table: AnnData,
@@ -279,9 +283,10 @@ def _right_exclusive_join_spatialelement_table(
279283
exclusive_table = table[keep, :] if has_match and keep.any() else None
280284
_inplace_fix_subset_categorical_obs(subset_adata=exclusive_table, original_adata=table)
281285
if exclusive_table is not None:
282-
exclusive_table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] = (
286+
exclusive_table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] = _region_as_str_if_list_of_len_one(
283287
exclusive_table.obs[region_column_name].unique().tolist()
284288
)
289+
TableModel.validate(exclusive_table)
285290
return element_dict, exclusive_table
286291

287292

@@ -383,9 +388,10 @@ def _inner_join_spatialelement_table(
383388

384389
_inplace_fix_subset_categorical_obs(subset_adata=joined_table, original_adata=table)
385390
if joined_table is not None:
386-
joined_table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] = (
391+
joined_table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] = _region_as_str_if_list_of_len_one(
387392
joined_table.obs[region_column_name].unique().tolist()
388393
)
394+
TableModel.validate(joined_table)
389395
return element_dict, joined_table
390396

391397

@@ -466,9 +472,10 @@ def _left_join_spatialelement_table(
466472
joined_table = table[joined_indices.tolist(), :].copy() if joined_indices is not None else None
467473
_inplace_fix_subset_categorical_obs(subset_adata=joined_table, original_adata=table)
468474
if joined_table is not None:
469-
joined_table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] = (
475+
joined_table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] = _region_as_str_if_list_of_len_one(
470476
joined_table.obs[region_column_name].unique().tolist()
471477
)
478+
TableModel.validate(joined_table)
472479
return element_dict, joined_table
473480

474481

tests/core/query/test_relational_query.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -256,13 +256,13 @@ def test_join_updates_spatialdata_attrs(sdata_query_aggregation):
256256
_, table = join_spatialelement_table(
257257
sdata=sdata, spatial_element_names="values_circles", table_name="table", how="left"
258258
)
259-
assert table.uns["spatialdata_attrs"]["region"] == ["values_circles"]
259+
assert table.uns["spatialdata_attrs"]["region"] == "values_circles"
260260

261261
# inner join on a single element
262262
_, table = join_spatialelement_table(
263263
sdata=sdata, spatial_element_names="values_circles", table_name="table", how="inner"
264264
)
265-
assert table.uns["spatialdata_attrs"]["region"] == ["values_circles"]
265+
assert table.uns["spatialdata_attrs"]["region"] == "values_circles"
266266

267267
# right_exclusive join: pass a truncated circles element so some table rows have no match.
268268
# values_circles has 9 instances (0-8); keep only 5 → 4 table rows are exclusive.
@@ -275,7 +275,7 @@ def test_join_updates_spatialdata_attrs(sdata_query_aggregation):
275275
)
276276
assert table is not None
277277
assert table.n_obs == 4
278-
assert table.uns["spatialdata_attrs"]["region"] == ["values_circles"]
278+
assert table.uns["spatialdata_attrs"]["region"] == "values_circles"
279279

280280
# original table metadata must be unchanged
281281
assert set(sdata["table"].uns["spatialdata_attrs"]["region"]) == {"values_circles", "values_polygons"}
@@ -1008,6 +1008,7 @@ def test_labels_table_joins(full_sdata):
10081008
def test_points_table_joins(full_sdata):
10091009
full_sdata["table"].uns["spatialdata_attrs"]["region"] = "points_0"
10101010
full_sdata["table"].obs["region"] = ["points_0"] * 100
1011+
full_sdata["table"].obs["region"] = full_sdata["table"].obs["region"].astype("category")
10111012

10121013
element_dict, table = join_spatialelement_table(
10131014
sdata=full_sdata,

0 commit comments

Comments
 (0)