Source code for fsh_lib.filters

import re
from datetime import date, datetime
from enum import StrEnum
from typing import TYPE_CHECKING, Any, Literal

from codegen_database.ext.text_range import TextMultiRangeType
from sqlalchemy import and_, cast, func, or_

from fsh_lib.relative_dates import resolve_relative_date

if TYPE_CHECKING:
    from collections.abc import Callable, Sequence

    from pydantic import BaseModel
    from sqlalchemy import Select
    from sqlalchemy.sql.elements import ColumnElement


[docs] class FilterOperator(StrEnum): EQ = "eq" NEQ = "neq" GT = "gt" GTE = "gte" LT = "lt" LTE = "lte" CONTAINS = "contains" STARTS_WITH = "starts_with" PATH_CONTAINS = "path_contains" RANGE_CONTAINS = "range_contains" IN = "in" IS_NULL = "is_null"
_FILTER_OPS: dict[FilterOperator, Callable[[Any, Any], ColumnElement[bool]]] = { FilterOperator.EQ: lambda col, v: col == v, FilterOperator.NEQ: lambda col, v: col != v, FilterOperator.GT: lambda col, v: col > v, FilterOperator.GTE: lambda col, v: col >= v, FilterOperator.LT: lambda col, v: col < v, FilterOperator.LTE: lambda col, v: col <= v, FilterOperator.CONTAINS: lambda col, v: col.contains(v), FilterOperator.STARTS_WITH: lambda col, v: col.startswith(v), FilterOperator.PATH_CONTAINS: lambda col, v: ( func.index(col, func.text2ltree(func.replace(v, "-", ""))) >= 0 ), FilterOperator.RANGE_CONTAINS: lambda col, v: cast( col, TextMultiRangeType() ).op("@>")(func.textrange(func.record_key(v), func.record_key(v), "[]")), FilterOperator.IN: lambda col, v: col.in_(v), FilterOperator.IS_NULL: ( lambda col, v: col.is_(None) if v else col.is_not(None) ), } _COMBINERS = {"and_": and_, "or_": or_} _TRIGRAM_LEN = 3 _RECORD_KEY_RE = re.compile(r"^[A-Za-z0-9-]+$") def _textrange_clause( col: Any, # noqa: ANN401 needle: str, ) -> ColumnElement[bool] | None: if not _RECORD_KEY_RE.match(needle): return None return cast(col, TextMultiRangeType()).op("@>")( func.textrange(func.record_key(needle), func.record_key(needle), "[]") ) def trigram_clause(col: Any, needle: str) -> ColumnElement[bool]: # noqa: ANN401 if len(needle) < _TRIGRAM_LEN: return col.icontains(needle, autoescape=True) return col.bool_op("%>")(needle) def apply_filters( stmt: Select, node: BaseModel | None, model: type, *, timezone: str = "UTC", ) -> Select: if node is None: return stmt clause = _build_filter_clause(node, model, timezone=timezone) if clause is None: return stmt return stmt.where(clause) def _build_filter_clause( node: BaseModel, model: type, *, timezone: str, ) -> ColumnElement[bool] | None: for attr, combiner in _COMBINERS.items(): children = getattr(node, attr, None) if children is not None: return _combine(children, combiner, model, timezone=timezone) field_name = getattr(node, "field", None) if field_name is None: return None # Normalize the wire ``op`` (a bare string on the Pydantic node) # into the enum once, so the dispatch lookup and comparisons below # are over a single type. op = FilterOperator(getattr(node, "op", FilterOperator.EQ)) value = getattr(node, "value", None) col = getattr(model, field_name) if isinstance(value, dict) and value.get("kind") == "relativeDate": value = resolve_relative_date( value, timezone=timezone, as_datetime=_column_python_type(col) is datetime, ) if op != FilterOperator.IS_NULL and _is_bool_column(col): if op == FilterOperator.IN and isinstance(value, (list, tuple)): value = [_as_bool(item) for item in value] else: value = _as_bool(value) elif op != FilterOperator.IS_NULL: if op == FilterOperator.IN and isinstance(value, (list, tuple)): value = [_coerce_to_column_type(col, item) for item in value] else: value = _coerce_to_column_type(col, value) return _FILTER_OPS[op](col, value) def _column_python_type(col: Any) -> type | None: # noqa: ANN401 try: return col.type.python_type except AttributeError, NotImplementedError: return None def _coerce_to_column_type(col: Any, value: Any) -> Any: # noqa: ANN401 if not isinstance(value, str): return value python_type = _column_python_type(col) if python_type is datetime: try: return datetime.fromisoformat(value) except ValueError: return value if python_type is date: try: return date.fromisoformat(value) except ValueError: return value return value def _is_bool_column(col: Any) -> bool: # noqa: ANN401 return _column_python_type(col) is bool def _as_bool(value: Any) -> Any: # noqa: ANN401 if isinstance(value, str): return value.strip().lower() == "true" return value def apply_search( stmt: Select, model: type, columns: Sequence[str], q: str | None, *, strategy: Literal["trigram", "tsvector"] = "trigram", textrange_columns: Sequence[str] = (), ) -> Select: if q is None or not q.strip() or not columns: return stmt needle = q.strip() clauses: list[ColumnElement[bool]] = [] for name in columns: col = getattr(model, name) if name in textrange_columns: clause = _textrange_clause(col, needle) if clause is not None: clauses.append(clause) continue if strategy == "tsvector": clauses.append( col.bool_op("@@")( func.websearch_to_tsquery("english", needle), ), ) else: clauses.append(trigram_clause(col, needle)) if not clauses: return stmt return stmt.where(or_(*clauses)) def _combine( children: Sequence[BaseModel], combiner: Callable[..., ColumnElement[bool]], model: type, *, timezone: str, ) -> ColumnElement[bool] | None: built = ( _build_filter_clause(child, model, timezone=timezone) for child in children ) clauses = [clause for clause in built if clause is not None] return combiner(*clauses) if clauses else None