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