Skip to content

Commit 3df6fdd

Browse files
committed
Fix CTable.where() applying a short mask to the wrong rows. Closes #607.
1 parent bea6a1c commit 3df6fdd

2 files changed

Lines changed: 75 additions & 4 deletions

File tree

‎src/blosc2/ctable.py‎

Lines changed: 15 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -13439,10 +13439,21 @@ def where( # noqa: C901
1343913439

1344013440
filter_len = len(filter)
1344113441
if filter_len != target_len:
13442-
if filter_len == self.nrows:
13443-
physical = blosc2.zeros(target_len, dtype=np.bool_)
13444-
physical[self._valid_rows] = filter[:]
13445-
filter = physical
13442+
n_live = self.nrows
13443+
if filter_len <= n_live:
13444+
# A mask no longer than the live-row count is logical: entry i
13445+
# selects the i-th live row of this view. A short one simply
13446+
# leaves the trailing live rows unselected. Padding it out to
13447+
# the physical length instead would align it with the
13448+
# underlying column, so it would pick up rows outside the view.
13449+
if filter_len < n_live:
13450+
logical = blosc2.zeros(n_live, dtype=np.bool_)
13451+
logical[:filter_len] = filter[:]
13452+
filter = logical
13453+
if n_live != target_len:
13454+
physical = blosc2.zeros(target_len, dtype=np.bool_)
13455+
physical[self._valid_rows] = filter[:]
13456+
filter = physical
1344613457
filter_intersected = True
1344713458
elif filter_len > target_len:
1344813459
filter = filter[:target_len]

‎tests/ctable/test_where_expressions.py‎

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,3 +73,63 @@ def test_where_string_expression_must_be_boolean():
7373

7474
with pytest.raises(TypeError, match="Expected boolean"):
7575
t.where("value * category")
76+
77+
78+
@dataclass
79+
class IdRow:
80+
id: int = blosc2.field(blosc2.int64(), default=0)
81+
82+
83+
def _id_table(n=6):
84+
t = blosc2.CTable(IdRow, expected_size=n)
85+
arr = np.empty(n, dtype=[("id", "<i8")])
86+
arr["id"] = np.arange(n)
87+
t.extend(arr, validate=False)
88+
return t
89+
90+
91+
def test_where_short_mask_is_relative_to_the_view():
92+
"""A mask shorter than the view selects the view's own rows.
93+
94+
It used to be padded out to the physical column length, which aligned it
95+
with the underlying column and let rows outside the view through.
96+
"""
97+
view = _id_table()[1:4] # ids [1, 2, 3]
98+
99+
# Predicate built from a slice of the view's column: [1, 2] > 1 -> [F, T].
100+
result = view.where(view["id"][0:2] > 1)
101+
102+
np.testing.assert_array_equal(result.id[:], np.array([2], dtype=np.int64))
103+
104+
105+
def test_where_short_mask_skips_deleted_rows():
106+
t = _id_table()
107+
t.delete([0, 2]) # live ids [1, 3, 4, 5]
108+
109+
result = t.where(np.array([True, False]))
110+
111+
np.testing.assert_array_equal(result.id[:], np.array([1], dtype=np.int64))
112+
113+
114+
@pytest.mark.parametrize(
115+
("mask", "expected"),
116+
[
117+
(np.array([], dtype=bool), []),
118+
(np.array([True, True]), [1, 2]),
119+
(np.array([False, True, True]), [2, 3]),
120+
],
121+
)
122+
def test_where_mask_lengths_on_a_sliced_view(mask, expected):
123+
view = _id_table()[1:4] # ids [1, 2, 3]
124+
125+
result = view.where(mask)
126+
127+
np.testing.assert_array_equal(result.id[:], np.array(expected, dtype=np.int64))
128+
129+
130+
def test_where_short_mask_on_a_sorted_view_follows_sort_order():
131+
ordered = _id_table().sort_by("id", ascending=False) # ids [5, 4, 3, 2, 1, 0]
132+
133+
result = ordered.where(np.array([True, False]))
134+
135+
np.testing.assert_array_equal(result.id[:], np.array([5], dtype=np.int64))

0 commit comments

Comments
 (0)