From 81e7090d705dcf792977916210670343efac486e Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Tue, 28 Jul 2026 17:28:29 -0400 Subject: [PATCH] fix(filter): coerce every operand against its column's type in one place MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The DSL had two operand paths with different rules. Comparison operators parsed datetimes and coerced numbers; bare equality bound whatever it was handed. SQLAlchemy types a bind from the operand rather than the column and the psycopg dialect renders that type as an explicit cast, so a mismatch compiled into valid-looking SQL and failed at execute time: "operator does not exist: timestamp with time zone = character varying". A matrix of column type x operand type x operator against a live database found 462 combinations, of which 55 built cleanly and then failed. The most plausible was a filter someone would write first try: {"created_at": "2026-01-01"} 500 {"created_at": {"gte": "2026-01-01"}} worked _coerce_operand now handles every operand, whatever the operator, keyed on the column's real type: JSONB takes an object, boolean takes only true/false, datetime parses strings, numeric goes through _coerce_numeric, text requires a string, and a column with no python_type (pgvector) is not filterable. eq/ne/gt/in cannot drift apart because they share the one call; `in` coerces element-wise, since a single element's type decides the cast rendered for that parameter. The matrix is now clean. This is a net deletion: the separate datetime, numeric, in-datetime and boolean branches, plus _require_bindable_operand, all collapse into it. Two more execute-time failures fixed on the way. `contains` was keyed on column_name == "h_metadata", so Document's equally-JSONB internal_metadata fell through to ILIKE and produced `jsonb ~~* text`; it now keys on the column type. And {"source_ids": "abc"} was `jsonb = character varying`. Closed-set columns are validated against the Literal that defines them, so declaring a new level or sync state updates filter validation with no change here. {"level": "banana"} was silently matching nothing. Empty IN is now always applied rather than skipped. Unifying the branches inherited a guard that had only ever wrapped the datetime path, which dropped the condition entirely and widened the query to every row — fail-open on an empty allowlist, which session scoping relies on to fail closed (see extract_session_allowlist). Caught by an existing test that asserts returned rows; the fuzz and the type matrix only check for errors, so neither would have seen it. Co-Authored-By: Claude Opus 5 (1M context) --- src/utils/filter.py | 262 ++++++++++++++++++------------------- tests/utils/test_filter.py | 64 ++++++++- 2 files changed, 194 insertions(+), 132 deletions(-) diff --git a/src/utils/filter.py b/src/utils/filter.py index d28de80c..a89a613d 100644 --- a/src/utils/filter.py +++ b/src/utils/filter.py @@ -2,7 +2,7 @@ import datetime from collections.abc import Callable, Sequence from decimal import Decimal from logging import getLogger -from typing import Any, TypeVar +from typing import Any, TypeVar, get_args from typing import cast as typing_cast from sqlalchemy import ColumnElement, Select, and_, case, cast, literal, not_, or_ @@ -11,6 +11,7 @@ from sqlalchemy.types import Numeric from ..exceptions import FilterError from .formatting import ILIKE_ESCAPE_CHAR, escape_ilike_pattern, parse_datetime_iso +from .types import DocumentLevel, VectorSyncState logger = getLogger(__name__) @@ -68,15 +69,14 @@ MAX_SESSION_ALLOWLIST_ENTRIES = 1000 # Values that can be bound to a non-JSONB column. Anything else (dict, list, # bytes, arbitrary objects) compiles into a valid statement and then fails in # psycopg at execute time as an unhandled 500, so it is rejected up front. -SCALAR_OPERAND_TYPES = ( - str, - bool, - int, - float, - Decimal, - datetime.datetime, - datetime.date, -) +# Columns whose values come from a closed set. Derived from the Literal types +# themselves, so adding a level (e.g. "abduction") or a sync state updates +# filter validation with no change here — an unlisted value is a 422 rather +# than a filter that silently matches nothing. +ENUM_COLUMN_VALUES: dict[str, frozenset[str]] = { + "level": frozenset(get_args(DocumentLevel)), + "sync_state": frozenset(get_args(VectorSyncState)), +} def _coerce_numeric(op_value: Any) -> float | Decimal: @@ -114,48 +114,112 @@ def _coerce_numeric(op_value: Any) -> float | Decimal: raise ValueError(f"not a number: {op_value!r}") from None -def _require_bindable_operand( - column_name: str, op_value: Any, operator: str = "", *, is_bool_column: bool = False -) -> None: - """Reject an operand that cannot be bound to a scalar column. +def _column_python_type(column: Any) -> type | None: + """Return a column's Python type, or None when it doesn't declare one. - For ``in``, each element is checked: a dict nested in the list is bound the - same way a bare dict operand would be, and fails identically. + pgvector's Vector raises NotImplementedError rather than returning a type, + so this must not be called bare. + """ + try: + return typing_cast("type | None", column.type.python_type) + except (AttributeError, NotImplementedError): + return None - A boolean column takes only ``true``/``false``. SQLAlchemy types the bind - from the operand rather than the column, so a string or int operand renders - as ``is_active = %(param)s::VARCHAR`` and Postgres rejects it at execute - time ("operator does not exist: boolean = character varying"). + +def _coerce_operand( + column: Any, column_name: str, value: Any, operator: str = "" +) -> Any: + """Return ``value`` ready to bind against ``column``, or raise FilterError. + + SQLAlchemy types a bind from the *operand*, not the column, and the psycopg + dialect renders that type as an explicit cast. So an operand whose type + doesn't match its column compiles into valid-looking SQL and then fails at + execute time — ``operator does not exist: text = integer``. Postgres will + not implicitly bridge these, so the mismatch has to be caught here. + + Every operand passes through this one function, whatever the operator, so + ``eq``/``ne``/``gt``/``in`` cannot drift apart by construction. Callers + handle None (a null check) and ``*`` (a wildcard) before calling. Args: - column_name: Internal column name, for the error message. - op_value: The operand to check. - operator: The comparison operator, when the operand came from one. - is_bool_column: Whether the target column is boolean. + column: SQLAlchemy column object. + column_name: Internal column name, for error messages. + value: The operand to coerce. + operator: The comparison operator, or "" for bare equality. + + Returns: + The operand, coerced where a lossless coercion exists. Raises: - FilterError: If a value is not bindable to the column. + FilterError: If the operand cannot be bound to this column. """ - values: Sequence[Any] = ( - typing_cast("Sequence[Any]", op_value) - if operator == "in" and isinstance(op_value, list | tuple | set) - else (op_value,) - ) - for value in values: - if value is None: - continue - if is_bool_column: - if isinstance(value, bool): - continue - raise FilterError( - f"Invalid value for column '{column_name}': expected true or false, got {type(value).__name__}" - ) - if isinstance(value, SCALAR_OPERAND_TYPES): - continue + # JSONB keeps containment semantics: the operand is a JSON document, not a + # scalar to compare. `jsonb >= 5` and `jsonb @> 'text'` have no operator. + if isinstance(column.type, JSONB): + if operator in ("", "contains") and isinstance(value, dict | list): + return typing_cast("Any", value) raise FilterError( - f"Invalid value for column '{column_name}': expected a scalar, got {type(value).__name__}" + f"Invalid filter for column '{column_name}': a JSONB column takes an object, optionally under 'contains'" ) + python_type = _column_python_type(column) + if python_type is None: + raise FilterError(f"Column '{column_name}' cannot be filtered on") + + # contains/icontains build an ILIKE pattern, so the operand is stringified + # and its own type doesn't matter — but the column must be text, or + # Postgres has no `~~` operator for it. + if operator in ("contains", "icontains"): + if python_type is not str: + raise FilterError( + f"Operator '{operator}' requires a text column, but '{column_name}' is {python_type.__name__}" + ) + return value + + # bool is checked before the numeric branch: it subclasses int, so a boolean + # column would otherwise have `true` coerced to 1, which Postgres rejects + # against a boolean column ("operator does not exist: boolean <> integer"). + if python_type is bool: + if isinstance(value, bool): + return value + raise FilterError( + f"Invalid value for column '{column_name}': expected true or false, got {type(value).__name__}" + ) + + if issubclass(python_type, datetime.datetime | datetime.date): + if isinstance(value, datetime.datetime | datetime.date): + return value + if isinstance(value, str): + validated = _validate_datetime_string(value) + if validated is None: + raise FilterError(f"Invalid datetime value: {value}") + return validated + raise FilterError( + f"Invalid value for column '{column_name}': expected a datetime, got {type(value).__name__}" + ) + + if issubclass(python_type, int | float | Decimal): + try: + return _coerce_numeric(value) + except (TypeError, ValueError): + raise FilterError( + f"Invalid numeric value: {value}. Expected a number, got {type(value).__name__}" + ) from None + + if python_type is str: + if not isinstance(value, str): + raise FilterError( + f"Invalid value for column '{column_name}': expected a string, got {type(value).__name__}" + ) + allowed = ENUM_COLUMN_VALUES.get(column_name) + if allowed is not None and value not in allowed: + raise FilterError( + f"Invalid value for column '{column_name}': {value!r}. Expected one of {sorted(allowed)}" + ) + return value + + raise FilterError(f"Column '{column_name}' cannot be filtered on") + def extract_session_allowlist( filters: dict[str, Any] | None, @@ -462,15 +526,9 @@ def _build_field_condition( return column == value else: if column_name in JSONB_COLUMNS: - return column.contains(value) + return column.contains(_coerce_operand(column, column_name, value)) else: - _require_bindable_operand( - column_name, - value, - is_bool_column=hasattr(column.type, "python_type") - and column.type.python_type is bool, - ) - return column == value + return column == _coerce_operand(column, column_name, value) def _safe_numeric_cast( @@ -671,26 +729,6 @@ def _build_comparison_conditions( """ conditions: list[ColumnElement[bool]] = [] - # Check if this is a datetime column - is_datetime_column = hasattr(column.type, "python_type") and issubclass( - column.type.python_type, datetime.datetime - ) - - # bool subclasses int, so a boolean column would otherwise be treated as - # numeric and have `true` coerced to 1 — which Postgres rejects against a - # boolean column ("operator does not exist: boolean <> integer"). - is_bool_column = ( - hasattr(column.type, "python_type") and column.type.python_type is bool - ) - - # Numeric coercion applies only to actually-numeric columns. On a text - # column, `ne` is a string comparison, not a failed float parse. - is_numeric_column = ( - not is_bool_column - and hasattr(column.type, "python_type") - and issubclass(column.type.python_type, int | float | Decimal) - ) - for operator, op_value in comparisons.items(): # Validate that the operator is supported if operator not in COMPARISON_OPERATORS: @@ -711,37 +749,14 @@ def _build_comparison_conditions( conditions.append(column.is_not(None)) continue - # Every operand bound to a scalar column must itself be a scalar. JSONB - # columns are exempt: a dict there is a containment match. - if not isinstance(column.type, JSONB): - _require_bindable_operand( - column_name, op_value, operator, is_bool_column=is_bool_column - ) - condition = None - # For datetime columns, cast string values to timestamp - if is_datetime_column and isinstance(op_value, str): - # Validate datetime string to prevent SQL injection - validated_datetime = _validate_datetime_string(op_value) - if validated_datetime is None: - # Raise error if datetime validation fails - raise FilterError(f"Invalid datetime value: {op_value}") - - # Use the validated datetime object directly instead of string interpolation - casted_value = validated_datetime - else: - # On a numeric column, a numeric operator's value must cast to a - # number. On a text column, `ne` is a string comparison. - if operator in NUMERIC_OPERATORS and is_numeric_column: - try: - casted_value = _coerce_numeric(op_value) - except (TypeError, ValueError): - raise FilterError( - f"Invalid numeric value: {op_value}. Expected a number, got {type(op_value).__name__}" - ) from None - else: - casted_value = op_value + # `in` coerces element-wise below; every other operator has one operand. + casted_value = ( + op_value + if operator == "in" + else _coerce_operand(column, column_name, op_value, operator) + ) if operator == "gte": condition = column >= casted_value @@ -759,43 +774,28 @@ def _build_comparison_conditions( if "*" in op_value: continue else: - if is_datetime_column: - # Validate and cast each datetime string value - casted_values: list[str | datetime.datetime] = [] - for val in op_value: - if isinstance(val, str): - validated_datetime = _validate_datetime_string(val) - if validated_datetime is None: - raise FilterError( - f"Invalid datetime value in list: {val}" - ) - casted_values.append(validated_datetime) - else: - casted_values.append(val) - if casted_values: - condition = column.in_(casted_values) - elif is_numeric_column: - # Same reason as the scalar path: a plain int renders an - # ::INTEGER cast, so one out-of-range element fails the - # whole IN at execute time. - try: - condition = column.in_( - [_coerce_numeric(val) for val in op_value] - ) - except (TypeError, ValueError): - raise FilterError( - f"Invalid numeric value in list for column '{column_name}': {op_value}" - ) from None - else: - condition = column.in_(list(op_value)) + # Element-wise: one bad element poisons the whole IN, since + # its type decides the cast rendered for that parameter. + # An empty list is applied, not skipped: `in: []` must match + # nothing. Dropping the condition would widen the query to + # every row, and session scoping relies on an empty + # allowlist failing closed (see extract_session_allowlist). + condition = column.in_( + [ + _coerce_operand(column, column_name, val, operator) + for val in op_value + ] + ) else: raise FilterError( f"Invalid value for 'in' operator: {op_value}. Expected an iterable (list, tuple, set), got {type(op_value).__name__}" ) elif operator == "contains": - if column_name == "h_metadata": - # For JSONB columns, use JSONB contains - condition = column.contains(op_value) + if isinstance(column.type, JSONB): + # Keyed on the column type, not the name: internal_metadata is + # equally JSONB and was falling through to ILIKE, which + # Postgres rejects as `jsonb ~~* text`. + condition = column.contains(casted_value) else: # For text columns, use ILIKE with escaped pattern escaped_value = escape_ilike_pattern(str(op_value)) diff --git a/tests/utils/test_filter.py b/tests/utils/test_filter.py index 5ff9279e..f1464827 100644 --- a/tests/utils/test_filter.py +++ b/tests/utils/test_filter.py @@ -3,7 +3,7 @@ from collections.abc import Sequence from datetime import datetime from decimal import Decimal -from typing import Any, cast +from typing import Any, cast, get_args import pytest from sqlalchemy import select @@ -14,6 +14,7 @@ from sqlalchemy.dialects.postgresql import psycopg as psycopg_dialect from src.exceptions import FilterError from src.models import Document, Message, Peer, Session from src.utils.filter import apply_filter +from src.utils.types import DocumentLevel def test_unknown_operator_dict_on_scalar_column_raises(): @@ -146,6 +147,67 @@ def test_null_operand_on_non_ne_operator_raises(): apply_filter(select(Message), Message, {"token_count": {"gt": None}}) +def test_enum_column_rejects_an_unknown_value(): + """An invalid level silently matched nothing, which reads as "no results" + rather than "you sent a value that cannot exist".""" + with pytest.raises(FilterError): + apply_filter(select(Document), Document, {"level": "banana"}) + + +@pytest.mark.parametrize("level", get_args(DocumentLevel)) +def test_enum_column_accepts_every_declared_level(level: str): + """Derived from the Literal, so a new level (e.g. "abduction") is covered + here the moment it is declared — no second list to keep in sync.""" + stmt = apply_filter(select(Document), Document, {"level": level}) + assert stmt.whereclause is not None + + +def test_empty_in_list_matches_nothing_rather_than_everything(): + """Dropping an empty IN would widen the query to every row. Session scoping + relies on an empty allowlist failing closed.""" + stmt = apply_filter(select(Document), Document, {"session_id": {"in": []}}) + assert stmt.whereclause is not None + assert "IN" in str(stmt.whereclause) + + +@pytest.mark.parametrize( + "filters", + [ + {"created_at": "2026-01-01"}, + {"created_at": {"gte": "2026-01-01"}}, + {"token_count": "5"}, + {"token_count": {"gt": "5"}}, + ], +) +def test_equality_and_comparison_paths_coerce_alike(filters: dict[str, Any]): + """The two paths had different rules: comparison operators parsed datetimes + and coerced numbers, bare equality bound the raw string and 500'd on + `timestamp with time zone = character varying`.""" + stmt = apply_filter(select(Message), Message, filters) + assert stmt.whereclause is not None + + +@pytest.mark.parametrize( + ("model", "filters"), + [ + (Message, {"token_count": {"contains": 5}}), # integer ~~* text + (Message, {"created_at": {"contains": "x"}}), # timestamptz ~~* text + (Document, {"metadata": {"gte": 5}}), # jsonb >= integer + (Document, {"metadata": {"contains": "x"}}), # jsonb ~~* text + (Document, {"source_ids": "abc"}), # jsonb = character varying + (Document, {"session_id": 5}), # text = integer + (Document, {"session_id": True}), # text = boolean + (Message, {"created_at": 5}), # timestamptz = integer + (Document, {"embedding": 5}), # no python_type at all + ], +) +def test_incompatible_operand_types_are_rejected(model: Any, filters: dict[str, Any]): + """Each of these compiled cleanly and failed in Postgres as + `operator does not exist: `.""" + with pytest.raises(FilterError): + apply_filter(select(model), model, filters) + + # --- Invariants over the whole DSL ------------------------------------------- # # The filter body is arbitrary client JSON. Enumerating bad shapes one at a time