dference 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
dference/_query.py ADDED
@@ -0,0 +1,492 @@
1
+ """Server-side view engine: filtering, sorting and paging of a :class:`DiffResult`.
2
+
3
+ The widget never receives the full table. Instead it sends a *query* (status
4
+ selection, search text, column filters, sort, page) and gets back only the rows
5
+ of the requested page. All work happens in polars on the joined frame:
6
+
7
+ 1. the query is turned into one boolean predicate and one sort key,
8
+ 2. the resulting ordered row ids are cached (paging through the same view is
9
+ a cheap slice), and
10
+ 3. only the ``page_size`` rows of the page are gathered and serialised.
11
+
12
+ Filter semantics for compared columns: a row matches if the **left or the
13
+ right** value matches - but only sides that actually exist in that row are
14
+ considered, so a missing side never matches ``is null``.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import datetime as dt
20
+ import math
21
+ from collections import OrderedDict
22
+ from collections.abc import Mapping
23
+ from dataclasses import dataclass, field
24
+ from decimal import Decimal
25
+ from typing import TYPE_CHECKING, Any, Final
26
+
27
+ import polars as pl
28
+
29
+ from ._compare import (
30
+ ROW,
31
+ STATUS,
32
+ STATUS_ORDER,
33
+ ColumnInfo,
34
+ DiffResult,
35
+ Status,
36
+ as_text,
37
+ dcol,
38
+ lcol,
39
+ rcol,
40
+ )
41
+
42
+ if TYPE_CHECKING:
43
+ from polars.datatypes import DataType
44
+
45
+ __all__ = ["ColumnFilter", "Query", "ViewEngine"]
46
+
47
+ MAX_PAGE_SIZE: Final = 1000
48
+ _CACHE_SIZE: Final = 8
49
+
50
+ #: Operators per filter type, as offered by the frontend.
51
+ OPERATORS: Final[Mapping[str, frozenset[str]]] = {
52
+ "string": frozenset(
53
+ {
54
+ "contains",
55
+ "not_contains",
56
+ "equals",
57
+ "starts",
58
+ "ends",
59
+ "is_empty",
60
+ "not_empty",
61
+ "is_null",
62
+ "not_null",
63
+ }
64
+ ),
65
+ "number": frozenset({"between", "eq", "ne", "is_null", "not_null"}),
66
+ "datetime": frozenset({"between", "is_null", "not_null"}),
67
+ "boolean": frozenset({"true", "false", "is_null"}),
68
+ }
69
+
70
+
71
+ # --------------------------------------------------------------------------- #
72
+ # Query model
73
+ # --------------------------------------------------------------------------- #
74
+
75
+
76
+ @dataclass(frozen=True, slots=True)
77
+ class ColumnFilter:
78
+ """Filter on one column (by index into ``DiffResult.columns``)."""
79
+
80
+ column: int
81
+ op: str
82
+ a: str = ""
83
+ b: str = ""
84
+
85
+
86
+ @dataclass(frozen=True, slots=True)
87
+ class Query:
88
+ """Everything that determines which rows are shown and in which order.
89
+
90
+ Hashable, so identical queries share one cached view.
91
+ """
92
+
93
+ statuses: frozenset[Status] = field(default_factory=lambda: frozenset(STATUS_ORDER))
94
+ search: str = ""
95
+ diff_column: int | None = None
96
+ #: with ``diff_column``: rows *equal* in that column instead of differing
97
+ diff_equal: bool = False
98
+ filters: tuple[ColumnFilter, ...] = ()
99
+ sort_column: int | None = None
100
+ sort_by_status: bool = False
101
+ descending: bool = False
102
+
103
+ @classmethod
104
+ def from_message(cls, msg: Mapping[str, Any], columns: tuple[ColumnInfo, ...]) -> Query:
105
+ """Parse and validate a query sent by the frontend.
106
+
107
+ Invalid parts are dropped rather than raising, so a stale or malformed
108
+ message can never break the widget.
109
+ """
110
+ n = len(columns)
111
+
112
+ def col_index(value: Any) -> int | None:
113
+ return value if isinstance(value, int) and 0 <= value < n else None
114
+
115
+ valid = {s.value for s in Status}
116
+ statuses = frozenset(Status(s) for s in msg.get("statuses", STATUS_ORDER) if s in valid)
117
+ diff_column = col_index(msg.get("diff_column"))
118
+ if diff_column is not None and columns[diff_column].kind != "compared":
119
+ diff_column = None
120
+
121
+ filters: list[ColumnFilter] = []
122
+ for raw in msg.get("filters", ()):
123
+ if not isinstance(raw, Mapping):
124
+ continue
125
+ i = col_index(raw.get("column"))
126
+ if i is None or raw.get("op") not in OPERATORS[columns[i].filter_type]:
127
+ continue
128
+ filters.append(
129
+ ColumnFilter(i, str(raw["op"]), str(raw.get("a", "")), str(raw.get("b", "")))
130
+ )
131
+
132
+ sort = msg.get("sort") or {}
133
+ sort_col = sort.get("column") if isinstance(sort, Mapping) else None
134
+ return cls(
135
+ statuses=statuses,
136
+ search=str(msg.get("search", ""))[:500],
137
+ diff_column=diff_column,
138
+ diff_equal=diff_column is not None and msg.get("diff_equal") is True,
139
+ filters=tuple(sorted(filters, key=lambda f: f.column)),
140
+ sort_column=col_index(sort_col),
141
+ sort_by_status=sort_col == "status",
142
+ descending=bool(isinstance(sort, Mapping) and sort.get("descending")),
143
+ )
144
+
145
+
146
+ @dataclass(frozen=True, slots=True)
147
+ class View:
148
+ """Ordered row ids of a query plus the columns that differ within it."""
149
+
150
+ ids: pl.Series
151
+ diff_columns: tuple[int, ...]
152
+
153
+ @property
154
+ def size(self) -> int:
155
+ """Number of rows in the view."""
156
+ return self.ids.len()
157
+
158
+
159
+ # --------------------------------------------------------------------------- #
160
+ # Engine
161
+ # --------------------------------------------------------------------------- #
162
+
163
+
164
+ class ViewEngine:
165
+ """Answers queries against one :class:`DiffResult`."""
166
+
167
+ def __init__(self, result: DiffResult) -> None:
168
+ self.result = result
169
+ self._cache: OrderedDict[Query, View] = OrderedDict()
170
+ # Row presence per side: a side exists unless the row is missing there.
171
+ self._present_l = pl.col(STATUS) != Status.MISSING_LEFT.value
172
+ self._present_r = pl.col(STATUS) != Status.MISSING_RIGHT.value
173
+ self._compared_idx = tuple(i for i, c in enumerate(result.columns) if c.kind == "compared")
174
+
175
+ # ---- public API ---------------------------------------------------------
176
+
177
+ def view(self, query: Query) -> View:
178
+ """Return (and cache) the ordered row ids matching ``query``."""
179
+ if (hit := self._cache.get(query)) is not None:
180
+ self._cache.move_to_end(query)
181
+ return hit
182
+
183
+ lazy = self.result.data.lazy().filter(self._predicate(query))
184
+ ordered = self._sorted(lazy, query).select(ROW)
185
+ flags = lazy.select(
186
+ pl.col(dcol(self.result.columns[i].name)).any() for i in self._compared_idx
187
+ )
188
+ ids_df, flags_df = pl.collect_all([ordered, flags])
189
+
190
+ flag_row = flags_df.row(0) if flags_df.width else ()
191
+ diff_columns = tuple(i for i, hit in zip(self._compared_idx, flag_row, strict=True) if hit)
192
+ view = View(ids=ids_df.get_column(ROW), diff_columns=diff_columns)
193
+
194
+ self._cache[query] = view
195
+ if len(self._cache) > _CACHE_SIZE:
196
+ self._cache.popitem(last=False)
197
+ return view
198
+
199
+ def page(self, query: Query, page: int, page_size: int) -> dict[str, Any]:
200
+ """Serialise one page of ``query`` for the frontend."""
201
+ page_size = max(1, min(int(page_size), MAX_PAGE_SIZE))
202
+ view = self.view(query)
203
+ pages = max(1, math.ceil(view.size / page_size))
204
+ page = max(0, min(int(page), pages - 1))
205
+ ids = view.ids.slice(page * page_size, page_size)
206
+ return {
207
+ "filtered": view.size,
208
+ "page": page,
209
+ "page_size": page_size,
210
+ "diff_columns": list(view.diff_columns),
211
+ "rows": self.serialise(ids),
212
+ }
213
+
214
+ def frame(self, query: Query) -> pl.DataFrame:
215
+ """Wide result frame of the rows in ``query``, in view order."""
216
+ return self.result.frame(rows=self.view(query).ids)
217
+
218
+ def export_csv(self, query: Query) -> bytes:
219
+ """CSV (UTF-8) of the rows in ``query``; nested values are stringified."""
220
+ df = self.frame(query)
221
+ nested = [name for name, dtype in df.schema.items() if dtype.is_nested()]
222
+ if nested:
223
+ df = df.with_columns(as_text(pl.col(n), df.schema[n]) for n in nested)
224
+ return df.write_csv().encode()
225
+
226
+ # ---- serialisation ----------------------------------------------------
227
+
228
+ def serialise(self, ids: pl.Series) -> list[dict[str, Any]]:
229
+ """Turn the given row ids into the compact row format of the frontend.
230
+
231
+ Each row is ``{"id", "s", "d", "l", "r"}``: status, indices of differing
232
+ columns and the left/right values aligned with ``columns`` (``None``
233
+ for a side that does not exist in that row).
234
+ """
235
+ if ids.is_empty():
236
+ return []
237
+ cols = self.result.columns
238
+ # Column name in `data` holding the left/right value of column i.
239
+ lsrc = [
240
+ c.name if c.kind == "key" else lcol(c.name) if c.kind != "right_only" else None
241
+ for c in cols
242
+ ]
243
+ rsrc = [
244
+ c.name if c.kind == "key" else rcol(c.name) if c.kind != "left_only" else None
245
+ for c in cols
246
+ ]
247
+ needed = {ROW, STATUS} | {n for n in (*lsrc, *rsrc) if n}
248
+ needed |= {dcol(cols[i].name) for i in self._compared_idx}
249
+ page = self.result.data.select(pl.col(sorted(needed)).gather(ids)).to_dicts()
250
+
251
+ rows: list[dict[str, Any]] = []
252
+ for rec in page:
253
+ status = rec[STATUS]
254
+ has_l = status != Status.MISSING_LEFT.value
255
+ has_r = status != Status.MISSING_RIGHT.value
256
+ rows.append(
257
+ {
258
+ "id": rec[ROW],
259
+ "s": status,
260
+ "d": [i for i in self._compared_idx if rec[dcol(cols[i].name)]],
261
+ "l": [_json(rec[n]) if n else None for n in lsrc] if has_l else None,
262
+ "r": [_json(rec[n]) if n else None for n in rsrc] if has_r else None,
263
+ }
264
+ )
265
+ return rows
266
+
267
+ # ---- expression building ------------------------------------------------
268
+
269
+ def _predicate(self, q: Query) -> pl.Expr:
270
+ """Combine all parts of the query into one boolean expression."""
271
+ parts: list[pl.Expr] = []
272
+ if len(q.statuses) < len(STATUS_ORDER):
273
+ parts.append(pl.col(STATUS).is_in([s.value for s in q.statuses]))
274
+ if q.diff_column is not None:
275
+ differs = pl.col(dcol(self.result.columns[q.diff_column].name))
276
+ # "equal" needs both sides: a row missing on one side is equal in no column
277
+ parts.append(~differs & self._present_l & self._present_r if q.diff_equal else differs)
278
+ if q.search.strip():
279
+ parts.append(self._search(q.search.strip().lower()))
280
+ parts.extend(self._column_filter(f) for f in q.filters)
281
+ return pl.all_horizontal(parts) if parts else pl.lit(True)
282
+
283
+ def _sides(self, i: int) -> list[tuple[pl.Expr, pl.Expr, DataType]]:
284
+ """``(value, present, dtype)`` for every side column ``i`` can exist on."""
285
+ c = self.result.columns[i]
286
+ if c.kind == "key":
287
+ return [(pl.col(c.name), pl.lit(True), c.dtype)]
288
+ sides: list[tuple[pl.Expr, pl.Expr, DataType]] = []
289
+ if c.dtype_left is not None:
290
+ sides.append((pl.col(lcol(c.name)), self._present_l, c.dtype_left))
291
+ if c.dtype_right is not None:
292
+ sides.append((pl.col(rcol(c.name)), self._present_r, c.dtype_right))
293
+ return sides
294
+
295
+ def _search(self, needle: str) -> pl.Expr:
296
+ """Case-insensitive substring search across all values of both sides.
297
+
298
+ Two optimisations keep this fast on millions of rows: non-text columns
299
+ are only searched if the needle could appear in their text form (e.g.
300
+ digits for numbers), and ASCII needles use polars' Aho-Corasick matcher
301
+ instead of lower-casing every value first.
302
+ """
303
+ exprs = [
304
+ _contains(value, dtype, needle)
305
+ for i in range(len(self.result.columns))
306
+ for value, _present, dtype in self._sides(i)
307
+ if _may_contain(dtype, needle)
308
+ ]
309
+ return pl.any_horizontal(exprs).fill_null(False) if exprs else pl.lit(False)
310
+
311
+ def _column_filter(self, f: ColumnFilter) -> pl.Expr:
312
+ """A row matches if any *existing* side of the column matches."""
313
+ ftype = self.result.columns[f.column].filter_type
314
+ tests = [
315
+ present & _test(value, dtype, ftype, f)
316
+ for value, present, dtype in self._sides(f.column)
317
+ ]
318
+ return pl.any_horizontal(tests).fill_null(False)
319
+
320
+ def _sorted(self, lazy: pl.LazyFrame, q: Query) -> pl.LazyFrame:
321
+ """Apply the sort of ``q`` (stable, nulls last, row id as tie-breaker)."""
322
+ if q.sort_by_status:
323
+ key = pl.col(STATUS).to_physical()
324
+ elif q.sort_column is not None:
325
+ key = self._sort_key(q.sort_column)
326
+ else:
327
+ return lazy.sort(ROW)
328
+ return lazy.sort([key, pl.col(ROW)], descending=[q.descending, False], nulls_last=True)
329
+
330
+ def _sort_key(self, i: int) -> pl.Expr:
331
+ """Value to sort column ``i`` by: the left value, else the right one."""
332
+ c = self.result.columns[i]
333
+ if c.kind == "key":
334
+ return pl.col(c.name)
335
+ if c.kind == "left_only":
336
+ return pl.col(lcol(c.name))
337
+ if c.kind == "right_only":
338
+ return pl.col(rcol(c.name))
339
+ left, right = pl.col(lcol(c.name)), pl.col(rcol(c.name))
340
+ if c.compare_dtype is not None:
341
+ left = (
342
+ as_text(left, c.dtype_left)
343
+ if c.compare_dtype == pl.String
344
+ else left.cast(c.compare_dtype)
345
+ )
346
+ right = (
347
+ as_text(right, c.dtype_right)
348
+ if c.compare_dtype == pl.String
349
+ else right.cast(c.compare_dtype)
350
+ )
351
+ return pl.when(self._present_l).then(left).otherwise(right)
352
+
353
+
354
+ # --------------------------------------------------------------------------- #
355
+ # Filter tests
356
+ # --------------------------------------------------------------------------- #
357
+
358
+
359
+ def _test(value: pl.Expr, dtype: DataType, ftype: str, f: ColumnFilter) -> pl.Expr: # noqa: PLR0911, PLR0912 - flat operator dispatch
360
+ """Boolean expression for one side of one column filter (nulls → False)."""
361
+ if dtype.is_float():
362
+ value = value.fill_nan(None) # NaN counts as missing, as in the comparison
363
+ match f.op:
364
+ case "is_null":
365
+ return value.is_null()
366
+ case "not_null":
367
+ return value.is_not_null()
368
+ case "true" | "false":
369
+ return value.eq_missing(pl.lit(f.op == "true"))
370
+
371
+ if ftype == "number":
372
+ return _number_test(value, f)
373
+ if ftype == "datetime":
374
+ return _date_test(value, f)
375
+
376
+ text = as_text(value, dtype)
377
+ lower, needle = text.str.to_lowercase(), f.a.lower()
378
+ match f.op:
379
+ case "is_empty":
380
+ return value.is_null() | (text == "")
381
+ case "not_empty":
382
+ return value.is_not_null() & (text != "")
383
+ case "equals":
384
+ result = text == f.a
385
+ case "starts":
386
+ result = lower.str.starts_with(needle)
387
+ case "ends":
388
+ result = lower.str.ends_with(needle)
389
+ case "not_contains":
390
+ result = ~lower.str.contains(needle, literal=True)
391
+ case _:
392
+ result = lower.str.contains(needle, literal=True)
393
+ return result.fill_null(False)
394
+
395
+
396
+ _NUMERIC_CHARS: Final = frozenset("0123456789.-+e")
397
+ _TEMPORAL_CHARS: Final = frozenset("0123456789-:t .")
398
+ _BOOL_WORDS: Final = ("true", "false")
399
+
400
+
401
+ def _may_contain(dtype: DataType, needle: str) -> bool:
402
+ """Can the text form of a ``dtype`` value contain ``needle`` at all?"""
403
+ if dtype == pl.Boolean:
404
+ return any(needle in word for word in _BOOL_WORDS)
405
+ if dtype.is_numeric():
406
+ return set(needle) <= _NUMERIC_CHARS | {"n", "a", "i", "f"} # also nan / inf
407
+ if dtype.is_temporal():
408
+ return set(needle) <= _TEMPORAL_CHARS
409
+ return True
410
+
411
+
412
+ def _contains(value: pl.Expr, dtype: DataType, needle: str) -> pl.Expr:
413
+ """Case-insensitive ``needle in value`` for one column (needle is lower-case)."""
414
+ text = as_text(value, dtype)
415
+ if needle.isascii():
416
+ return text.str.contains_any([needle], ascii_case_insensitive=True)
417
+ return text.str.to_lowercase().str.contains(needle, literal=True)
418
+
419
+
420
+ def _number_test(value: pl.Expr, f: ColumnFilter) -> pl.Expr:
421
+ a, b = _to_float(f.a), _to_float(f.b)
422
+ if f.op == "eq":
423
+ return (value == a).fill_null(False) if a is not None else pl.lit(False)
424
+ if f.op == "ne":
425
+ return (value != a).fill_null(False) if a is not None else pl.lit(False)
426
+ cond = value.is_not_null()
427
+ if a is not None:
428
+ cond &= value >= a
429
+ if b is not None:
430
+ cond &= value <= b
431
+ return cond.fill_null(False)
432
+
433
+
434
+ def _date_test(value: pl.Expr, f: ColumnFilter) -> pl.Expr:
435
+ as_date = value.cast(pl.Date, strict=False)
436
+ cond = value.is_not_null()
437
+ if (a := _to_date(f.a)) is not None:
438
+ cond &= as_date >= a
439
+ if (b := _to_date(f.b)) is not None:
440
+ cond &= as_date <= b
441
+ return cond.fill_null(False)
442
+
443
+
444
+ def _to_float(text: str) -> float | None:
445
+ try:
446
+ return float(text) if text.strip() else None
447
+ except ValueError:
448
+ return None
449
+
450
+
451
+ def _to_date(text: str) -> dt.date | None:
452
+ try:
453
+ return dt.date.fromisoformat(text.strip()[:10]) if text.strip() else None
454
+ except ValueError:
455
+ return None
456
+
457
+
458
+ # --------------------------------------------------------------------------- #
459
+ # JSON conversion of single values
460
+ # --------------------------------------------------------------------------- #
461
+
462
+
463
+ #: largest integer a JavaScript number represents exactly (Number.MAX_SAFE_INTEGER)
464
+ _MAX_SAFE_INT: Final = 2**53 - 1
465
+
466
+
467
+ def _json(v: Any) -> Any: # noqa: PLR0911 - a flat dispatch reads best here
468
+ """Make one value JSON-safe for the frontend.
469
+
470
+ Non-finite floats become strings (``"NaN"``, ``"inf"``) so they stay
471
+ distinguishable from null; temporal values use ISO 8601. Integers beyond
472
+ JavaScript's exact range and decimals a float cannot hold exactly become
473
+ strings too - otherwise two different values could show up as equal.
474
+ """
475
+ if v is None or isinstance(v, (bool, str)):
476
+ return v
477
+ if isinstance(v, int):
478
+ return v if abs(v) <= _MAX_SAFE_INT else str(v)
479
+ if isinstance(v, float):
480
+ return v if math.isfinite(v) else str(v)
481
+ if isinstance(v, (dt.datetime, dt.date, dt.time)):
482
+ return v.isoformat()
483
+ if isinstance(v, Decimal):
484
+ if not v.is_finite(): # NaN, sNaN, Infinity: float() rejects sNaN
485
+ return str(v)
486
+ f = float(v)
487
+ return f if Decimal(repr(f)) == v else str(v)
488
+ if isinstance(v, dt.timedelta):
489
+ return str(v)
490
+ if isinstance(v, bytes):
491
+ return v.hex()
492
+ return str(v)