diff --git a/src/crud.py b/src/crud.py index 22f7539a..b0860211 100644 --- a/src/crud.py +++ b/src/crud.py @@ -17,6 +17,7 @@ from . import models, schemas from .exceptions import ( ResourceNotFoundException, ) +from .utils.filter import apply_filter load_dotenv(override=True) @@ -85,18 +86,17 @@ async def get_or_create_workspace( async def get_all_workspaces( - filter: dict[str, Any] | None = None, + filters: dict[str, Any] | None = None, ) -> Select[tuple[models.Workspace]]: """ Get all workspaces. Args: db: Database session - filter: Filter the workspaces by a dictionary of metadata + filters: Filter the workspaces by a dictionary of metadata """ stmt = select(models.Workspace) - if filter is not None: - stmt = stmt.where(models.Workspace.h_metadata.contains(filter)) + stmt = apply_filter(stmt, models.Workspace, filters) stmt: Select[tuple[models.Workspace]] = stmt.order_by(models.Workspace.created_at) return stmt @@ -239,12 +239,11 @@ async def get_peer( async def get_peers( workspace_name: str, - filter: dict[str, str] | None = None, + filters: dict[str, str] | None = None, ) -> Select[tuple[models.Peer]]: stmt = select(models.Peer).where(models.Peer.workspace_name == workspace_name) - if filter is not None: - stmt = stmt.where(models.Peer.h_metadata.contains(filter)) + stmt = apply_filter(stmt, models.Peer, filters) stmt = stmt.order_by(models.Peer.created_at) @@ -291,8 +290,7 @@ async def update_peer( async def get_sessions_for_peer( workspace_name: str, peer_name: str, - is_active: bool | None = None, - filter: dict[str, Any] | None = None, + filters: dict[str, Any] | None = None, ) -> Select[tuple[models.Session]]: """ Get all sessions for a peer through the session_peers relationship. @@ -300,8 +298,7 @@ async def get_sessions_for_peer( Args: workspace_name: Name of the workspace peer_name: Name of the peer - is_active: Filter by active status (True/False/None for all) - filter: Filter sessions by metadata + filters: Filter sessions by metadata Returns: SQLAlchemy Select statement @@ -317,11 +314,7 @@ async def get_sessions_for_peer( .where(models.Session.workspace_name == workspace_name) ) - if is_active is not None: - stmt = stmt.where(models.Session.is_active == is_active) - - if filter is not None: - stmt = stmt.where(models.Session.h_metadata.contains(filter)) + stmt = apply_filter(stmt, models.Session, filters) stmt: Select[tuple[models.Session]] = stmt.order_by(models.Session.created_at) @@ -335,19 +328,14 @@ async def get_sessions_for_peer( async def get_sessions( workspace_name: str, - is_active: bool | None = None, - filter: dict[str, Any] | None = None, + filters: dict[str, Any] | None = None, ) -> Select[tuple[models.Session]]: """ Get all sessions in a workspace. """ stmt = select(models.Session).where(models.Session.workspace_name == workspace_name) - if is_active: - stmt = stmt.where(models.Session.is_active.is_(True)) - - if filter is not None: - stmt = stmt.where(models.Session.h_metadata.contains(filter)) + stmt = apply_filter(stmt, models.Session, filters) stmt = stmt.order_by(models.Session.created_at) @@ -1261,7 +1249,7 @@ async def get_messages( workspace_name: str, session_name: str, reverse: bool | None = False, - filter: dict[str, Any] | None = None, + filters: dict[str, Any] | None = None, token_limit: int | None = None, message_count_limit: int | None = None, ) -> Select[tuple[models.Message]]: @@ -1275,7 +1263,7 @@ async def get_messages( workspace_name: Name of the workspace session_name: Name of the session reverse: Whether to reverse the order of messages - filter: Filter to apply to the messages + filters: Filter to apply to the messages token_limit: Maximum number of tokens to include in the messages message_count_limit: Maximum number of messages to include @@ -1288,13 +1276,10 @@ async def get_messages( models.Message.session_name == session_name, ] - # Add metadata filter if provided - if filter is not None: - base_conditions.append(models.Message.h_metadata.contains(filter)) - # Apply message count limit first (takes precedence over token limit) if message_count_limit is not None: stmt = select(models.Message).where(*base_conditions) + stmt = apply_filter(stmt, models.Message, filters) # For message count limit, we want the most recent N messages # So we order by id desc to get most recent, then apply limit stmt = stmt.order_by(models.Message.id.desc()).limit(message_count_limit) @@ -1304,7 +1289,6 @@ async def get_messages( stmt = stmt.order_by(models.Message.id.desc()) else: stmt = stmt.order_by(models.Message.id.asc()) - elif token_limit is not None: # Apply token limit logic # Create a subquery that calculates running sum of tokens for most recent messages @@ -1325,16 +1309,17 @@ async def get_messages( .join(token_subquery, models.Message.id == token_subquery.c.id) .where(token_subquery.c.running_token_sum <= token_limit) ) + stmt = apply_filter(stmt, models.Message, filters) # Apply final ordering based on reverse parameter if reverse: stmt = stmt.order_by(models.Message.id.desc()) else: stmt = stmt.order_by(models.Message.id.asc()) - else: # Default case - no limits applied stmt = select(models.Message).where(*base_conditions) + stmt = apply_filter(stmt, models.Message, filters) if reverse: stmt = stmt.order_by(models.Message.id.desc()) else: @@ -1398,7 +1383,7 @@ async def get_messages_for_peer( workspace_name: str, peer_name: str, reverse: bool | None = False, - filter: dict[str, Any] | None = None, + filters: dict[str, Any] | None = None, ) -> Select[tuple[models.Message]]: stmt = ( select(models.Message) @@ -1407,8 +1392,7 @@ async def get_messages_for_peer( .where(models.Message.session_name.is_(None)) ) - if filter is not None: - stmt = stmt.where(models.Message.h_metadata.contains(filter)) + stmt = apply_filter(stmt, models.Message, filters) if reverse: stmt = stmt.order_by(models.Message.id.desc()) @@ -1531,7 +1515,7 @@ async def query_documents( peer_name: str, collection_name: str, query: str, - filter: dict[str, Any] | None = None, + filters: dict[str, Any] | None = None, max_distance: float | None = None, top_k: int = 5, ) -> Sequence[models.Document]: @@ -1548,8 +1532,7 @@ async def query_documents( stmt = stmt.where( models.Document.embedding.cosine_distance(embedding_query) < max_distance ) - if filter is not None: - stmt = stmt.where(models.Document.internal_metadata.contains(filter)) + stmt = apply_filter(stmt, models.Document, filters) stmt = stmt.limit(top_k).order_by( models.Document.embedding.cosine_distance(embedding_query) ) diff --git a/src/exceptions.py b/src/exceptions.py index e24799ca..80b4e71d 100644 --- a/src/exceptions.py +++ b/src/exceptions.py @@ -63,3 +63,11 @@ class DisabledException(HonchoException): status_code = 405 detail = "Feature is disabled" + + +@final +class FilterError(HonchoException): + """Exception raised when a filter is misconfigured or invalid.""" + + status_code = 422 + detail = "Invalid filter configuration" diff --git a/src/routers/messages.py b/src/routers/messages.py index 32ff396f..b909b7b1 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -3,7 +3,7 @@ from typing import Any from fastapi import APIRouter, BackgroundTasks, Body, Depends, Path, Query from fastapi_pagination import Page -from fastapi_pagination.ext.sqlalchemy import paginate +from fastapi_pagination.ext.sqlalchemy import apaginate from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.sql import insert @@ -321,20 +321,20 @@ async def get_messages( ): """Get all messages for a session""" try: - filter = None + filters = None if options and hasattr(options, "filter"): - filter = options.filter - if filter == {}: - filter = None + filters = options.filter + if filters == {}: + filters = None messages_query = await crud.get_messages( workspace_name=workspace_id, session_name=session_id, - filter=filter, + filters=filters, reverse=reverse, ) - return await paginate(db, messages_query) + return await apaginate(db, messages_query) except ValueError as e: logger.warning(f"Failed to get messages for session {session_id}: {str(e)}") raise ResourceNotFoundException("Session not found") from e diff --git a/src/routers/peers.py b/src/routers/peers.py index dccfe86e..7592ae0d 100644 --- a/src/routers/peers.py +++ b/src/routers/peers.py @@ -11,7 +11,7 @@ from fastapi import ( from fastapi.exceptions import HTTPException from fastapi.responses import StreamingResponse from fastapi_pagination import Page -from fastapi_pagination.ext.sqlalchemy import paginate +from fastapi_pagination.ext.sqlalchemy import apaginate from mirascope.llm import Stream from sqlalchemy.ext.asyncio import AsyncSession @@ -51,9 +51,9 @@ async def get_peers( if filter_param == {}: filter_param = None - return await paginate( + return await apaginate( db, - await crud.get_peers(workspace_name=workspace_id, filter=filter_param), + await crud.get_peers(workspace_name=workspace_id, filters=filter_param), ) @@ -128,23 +128,18 @@ async def get_sessions_for_peer( ): """Get All Sessions for a Peer""" filter_param = None - is_active = True - if options: - if hasattr(options, "filter"): - filter_param = options.filter - if filter_param == {}: - filter_param = None - if hasattr(options, "is_active"): - is_active = options.is_active + if options and hasattr(options, "filter"): + filter_param = options.filter + if filter_param == {}: + filter_param = None - return await paginate( + return await apaginate( db, await crud.get_sessions_for_peer( workspace_name=workspace_id, peer_name=peer_id, - is_active=is_active, - filter=filter_param, + filters=filter_param, ), ) @@ -274,20 +269,20 @@ async def get_messages_for_peer( ): """Get all messages for a peer""" try: - filter = None + filters = None if options and hasattr(options, "filter"): - filter = options.filter - if filter == {}: - filter = None + filters = options.filter + if filters == {}: + filters = None messages_query = await crud.get_messages_for_peer( workspace_name=workspace_id, peer_name=peer_id, - filter=filter, + filters=filters, reverse=reverse, ) - return await paginate(db, messages_query) + return await apaginate(db, messages_query) except ValueError as e: logger.warning(f"Failed to get messages for peer {peer_id}: {str(e)}") raise ResourceNotFoundException("Peer not found") from e @@ -337,4 +332,4 @@ async def search_peer( """Search a Peer""" stmt = await crud.search(query, workspace_name=workspace_id, peer_name=peer_id) - return await paginate(db, stmt) + return await apaginate(db, stmt) diff --git a/src/routers/sessions.py b/src/routers/sessions.py index 92226015..720514ee 100644 --- a/src/routers/sessions.py +++ b/src/routers/sessions.py @@ -2,7 +2,7 @@ import logging from fastapi import APIRouter, Body, Depends, Path, Query, Response from fastapi_pagination import Page -from fastapi_pagination.ext.sqlalchemy import paginate +from fastapi_pagination.ext.sqlalchemy import apaginate from sqlalchemy.ext.asyncio import AsyncSession from src import crud, schemas @@ -84,23 +84,14 @@ async def get_sessions( ): """Get All Sessions in a Workspace""" filter_param = None - is_active_param = False # Default from schema - if options: - if hasattr(options, "filter") and options.filter: - filter_param = options.filter - if filter_param == {}: # Explicitly check for empty dict - filter_param = None - if hasattr(options, "is_active"): # Check if is_active is present - is_active_param = options.is_active + if options and hasattr(options, "filter") and options.filter: + filter_param = options.filter + if filter_param == {}: # Explicitly check for empty dict + filter_param = None - return await paginate( - db, - await crud.get_sessions( - workspace_name=workspace_id, - is_active=is_active_param, - filter=filter_param, - ), + return await apaginate( + db, await crud.get_sessions(workspace_name=workspace_id, filters=filter_param) ) @@ -361,7 +352,7 @@ async def get_session_peers( peers_query = await crud.get_peers_from_session( workspace_name=workspace_id, session_name=session_id ) - return await paginate(db, peers_query) + return await apaginate(db, peers_query) except ValueError as e: logger.warning(f"Failed to get peers from session {session_id}: {str(e)}") raise ResourceNotFoundException("Session not found") from e @@ -468,4 +459,4 @@ async def search_session( query, workspace_name=workspace_id, session_name=session_id ) - return await paginate(db, stmt) + return await apaginate(db, stmt) diff --git a/src/routers/workspaces.py b/src/routers/workspaces.py index fcb0e4ff..1e9632de 100644 --- a/src/routers/workspaces.py +++ b/src/routers/workspaces.py @@ -2,7 +2,7 @@ import logging from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query from fastapi_pagination import Page -from fastapi_pagination.ext.sqlalchemy import paginate +from fastapi_pagination.ext.sqlalchemy import apaginate from sqlalchemy.ext.asyncio import AsyncSession from src import crud, schemas @@ -65,9 +65,9 @@ async def get_all_workspaces( if filter_param == {}: filter_param = None - return await paginate( + return await apaginate( db, - await crud.get_all_workspaces(filter=filter_param), + await crud.get_all_workspaces(filters=filter_param), ) @@ -104,7 +104,7 @@ async def search_workspace( """Search a Workspace""" stmt = await crud.search(query, workspace_name=workspace_id) - return await paginate(db, stmt) + return await apaginate(db, stmt) @router.get( diff --git a/src/schemas.py b/src/schemas.py index dba1d965..7cf189d5 100644 --- a/src/schemas.py +++ b/src/schemas.py @@ -169,7 +169,6 @@ class SessionCreate(SessionBase): class SessionGet(SessionBase): filter: dict[str, Any] | None = None - is_active: bool = False class SessionUpdate(SessionBase): @@ -248,6 +247,7 @@ class DialecticResponse(BaseModel): class SessionCounts(BaseModel): """Counts for a specific session in queue processing.""" + completed: int in_progress: int pending: int @@ -255,6 +255,7 @@ class SessionCounts(BaseModel): class QueueCounts(BaseModel): """Aggregated counts for queue processing status.""" + total: int completed: int in_progress: int @@ -264,6 +265,7 @@ class QueueCounts(BaseModel): class QueueStatusRow(BaseModel): """Represents a row from the queue status SQL query result.""" + session_id: str | None total: int completed: int @@ -277,6 +279,7 @@ class QueueStatusRow(BaseModel): class PeerConfigResult(BaseModel): """Result from querying peer configuration data.""" + peer_name: str peer_configuration: dict[str, Any] session_peer_configuration: dict[str, Any] @@ -284,11 +287,13 @@ class PeerConfigResult(BaseModel): class SessionPeerData(BaseModel): """Data for managing session peer relationships.""" + peer_names: dict[str, SessionPeerConfig] class MessageBulkData(BaseModel): """Data for bulk message operations.""" + messages: list[MessageCreate] session_name: str workspace_name: str diff --git a/src/utils/filter.py b/src/utils/filter.py new file mode 100644 index 00000000..a3982196 --- /dev/null +++ b/src/utils/filter.py @@ -0,0 +1,568 @@ +import datetime +from collections.abc import Callable +from logging import getLogger +from typing import Any, TypeVar + +from sqlalchemy import ColumnElement, Select, and_, case, cast, literal, not_, or_ +from sqlalchemy.types import Numeric + +from ..exceptions import FilterError + +logger = getLogger(__name__) + +# Type variable for SQLAlchemy model classes +T = TypeVar("T") + +# Module-level constants for comparison operators +COMPARISON_OPERATORS = { + "gte", + "lte", + "gt", + "lt", + "ne", + "in", + "contains", + "icontains", +} + +NUMERIC_OPERATORS = {"gte", "lte", "gt", "lt", "ne"} + +ALLOWED_EXTERNAL_TO_INTERNAL_COLUMN_MAPPING = { + "id": "name", + "created_at": "created_at", + "is_active": "is_active", + "workspace_id": "workspace_name", + "session_id": "session_name", + "peer_id": "peer_name", + "metadata": "h_metadata", +} + +ALLOWED_EXTERNAL_TO_INTERNAL_COLUMN_MAPPING_MESSAGES = { + "workspace_id": "workspace_name", + "session_id": "session_name", + "peer_id": "peer_name", + "token_count": "token_count", + "created_at": "created_at", + "metadata": "h_metadata", +} + + +def apply_filter( + stmt: Select[tuple[T]], model_class: type[T], filters: dict[str, Any] | None = None +) -> Select[tuple[T]]: + """ + Apply advanced filter to a SQL statement based on filter dictionary. + + Supports logical operators (AND, OR, NOT), comparison operators + (gte, lte, gt, lt, ne, contains, icontains, in), and wildcard character (*). + + Note that the filter refers to column names from the user perspective: + that means all `*_name` fields are actually `*_id` fields and `h_metadata` + is actually `metadata`. + + Examples: + # Simple filters (backward compatible) + {"peer_id": "alice", "metadata": {"type": "user"}} + + # Logical operators + {"AND": [{"peer_id": "alice"}, {"created_at": {"gte": "2024-01-01"}}]} + {"OR": [{"peer_id": "alice"}, {"peer_id": "bob"}]} + {"NOT": [{"peer_id": "alice"}]} + + # Comparison operators + {"created_at": {"gte": "2024-01-01", "lte": "2024-12-31"}} + {"peer_id": {"in": ["alice", "bob"]}} + + # Wildcards (matches everything for that field) + {"peer_id": "*"} + + Args: + stmt: SQLAlchemy Select statement to modify + model_class: SQLAlchemy model class for column access + filters: Optional filter dictionary + + Returns: + Modified Select statement with filter applied if provided + + Raises: + FilterError: When the filter contains invalid configuration or values + """ + if filters is None: + return stmt + + conditions = _build_filter_conditions(filters, model_class) + if conditions is not None: + stmt = stmt.where(conditions) + + return stmt + + +def _build_filter_conditions( + filter_dict: dict[str, Any], model_class: type[Any] +) -> ColumnElement[bool] | None: + """ + Recursively build filter conditions from a filter dictionary. + + Args: + filter_dict: Filter dictionary that may contain logical operators + model_class: SQLAlchemy model class for column access + + Returns: + SQLAlchemy condition object or None + """ + conditions: list[ColumnElement[bool]] = [] + + # Handle logical operators + if "AND" in filter_dict: + if not isinstance(filter_dict["AND"], list): + raise FilterError( + f"AND operator must contain a list, got {type(filter_dict['AND']).__name__}" + ) + and_conditions: list[ColumnElement[bool]] = [] + for sub_filter in filter_dict["AND"]: # pyright: ignore + sub_condition = _build_filter_conditions(sub_filter, model_class) # pyright: ignore + if sub_condition is not None: + and_conditions.append(sub_condition) + if and_conditions: + conditions.append(and_(*and_conditions)) + + if "OR" in filter_dict: + if not isinstance(filter_dict["OR"], list): + raise FilterError( + f"OR operator must contain a list, got {type(filter_dict['OR']).__name__}" + ) + or_conditions: list[ColumnElement[bool]] = [] + for sub_filter in filter_dict["OR"]: # pyright: ignore + sub_condition = _build_filter_conditions(sub_filter, model_class) # pyright: ignore + if sub_condition is not None: + or_conditions.append(sub_condition) + if or_conditions: + conditions.append(or_(*or_conditions)) + + if "NOT" in filter_dict: + if filter_dict["NOT"] is None: + raise FilterError("NOT operator cannot be None") + if not isinstance(filter_dict["NOT"], list): + raise FilterError( + f"NOT operator must contain a list, got {type(filter_dict['NOT']).__name__}" + ) + not_conditions: list[ColumnElement[bool]] = [] + for sub_filter in filter_dict["NOT"]: # pyright: ignore + sub_condition = _build_filter_conditions(sub_filter, model_class) # pyright: ignore + if sub_condition is not None: + not_conditions.append( + not_(sub_condition) + ) # Apply NOT to each condition individually + if not_conditions: + conditions.append(and_(*not_conditions)) # Then AND them together + + # Handle field-level conditions (skip logical operator keys) + logical_keys = {"AND", "OR", "NOT"} + for key, value in filter_dict.items(): + if key in logical_keys: + continue + + condition = _build_field_condition(key, value, model_class) + if condition is not None: + conditions.append(condition) + + # Combine all conditions with AND + if len(conditions) == 0: + return None + elif len(conditions) == 1: + return conditions[0] + else: + return and_(*conditions) + + +def _build_field_condition( + key: str, value: Any, model_class: type[Any] +) -> ColumnElement[bool] | None: + """ + Build a condition for a single field. + + Args: + key: Field name + value: Field value or comparison dict + model_class: SQLAlchemy model class + + Returns: + SQLAlchemy condition object or None + """ + if model_class.__name__ == "Message": + column_name = ALLOWED_EXTERNAL_TO_INTERNAL_COLUMN_MAPPING_MESSAGES.get(key) + else: + column_name = ALLOWED_EXTERNAL_TO_INTERNAL_COLUMN_MAPPING.get(key) + + if column_name is None: + raise FilterError( + f"Column '{key}' is not allowed to be filtered on or does not exist on {model_class.__name__}" + ) + + # Check if the column exists on the model + if not hasattr(model_class, column_name): + raise FilterError(f"Column '{key}' does not exist on {model_class.__name__}") + + column = getattr(model_class, column_name) + + # Handle wildcard - matches everything, so no condition needed + if value == "*": + return None + + # Handle comparison operators vs regular values + if isinstance(value, dict): + # Check if this is a comparison operators dict by looking for known operators + is_comparison_dict = any(op_key in COMPARISON_OPERATORS for op_key in value) # pyright: ignore + + if is_comparison_dict: + return _build_comparison_conditions(column, column_name, value) # pyright: ignore + else: + # This is a regular value that happens to be a dict + # For JSONB fields (metadata, configuration), check if it contains nested comparison operators + if column_name in ("h_metadata", "configuration"): + return _build_nested_metadata_conditions(column, value) # pyright: ignore + else: + return column == value + else: + if column_name in ("h_metadata", "configuration"): + return column.contains(value) + else: + return column == value + + +def _safe_numeric_cast( + column_accessor: ColumnElement[Any], op_value: Any +) -> tuple[ColumnElement[Any], Any]: + """ + Safely cast JSONB column accessor to appropriate type for comparison. + + Args: + column_accessor: SQLAlchemy JSONB column accessor (.astext) + op_value: The value to compare against + + Returns: + Tuple of (cast_column_accessor, cast_op_value) for typed comparison + or (column_accessor, str_op_value) for string comparison + """ + try: + if isinstance(op_value, bool): + # For boolean values, compare with the string representation + # PostgreSQL JSONB stores booleans as "true"/"false" strings when extracted with ->> + return column_accessor, str(op_value).lower() + + # For numeric values, use a safer cast that handles empty strings and invalid values + # We use CASE WHEN to handle empty strings and non-numeric values gracefully + safe_cast = case( + (column_accessor == "", literal(None)), # Empty string -> NULL + (column_accessor.is_(None), literal(None)), # NULL -> NULL + else_=cast(column_accessor, Numeric()), + ) + + if isinstance(op_value, int | float): + return safe_cast, op_value + else: + # Try to parse as numeric (handles both strings and other types) + try: + # Try int first, then float + parsed_value = int(op_value) + return safe_cast, parsed_value + except (ValueError, TypeError): + try: + parsed_value = float(op_value) + return safe_cast, parsed_value + except (ValueError, TypeError): + if isinstance(op_value, str): + # If it's not numeric, treat as string comparison (e.g., dates, text) + # This allows date strings like "2024-02-01" to be compared lexicographically + return column_accessor, str(op_value) + else: + raise FilterError( + f"Invalid value for numeric operator: {op_value}. Expected a number, got {type(op_value).__name__}" + ) from None + except Exception as e: + raise FilterError( + f"Failed to process numeric cast for value '{op_value}': {str(e)}" + ) from e + + +def _build_comparison_condition( + column: Any, field_name: str, operator: str, op_value: Any +) -> ColumnElement[bool] | None: + """ + Build a single comparison condition for a JSONB field. + + Args: + column: SQLAlchemy JSONB column object + field_name: Name of the field in the JSONB column + operator: Comparison operator + op_value: Value to compare against + + Returns: + SQLAlchemy condition object or None + """ + # Validate that the operator is supported + if operator not in COMPARISON_OPERATORS: + raise FilterError(f"Unsupported comparison operator: {operator}") + + # Handle wildcard - matches everything, so no condition needed + if op_value == "*": + return None + + field_accessor = column[field_name].astext + + # Mapping of operators to their SQLAlchemy methods + if operator in NUMERIC_OPERATORS: + try: + safe_accessor, safe_value = _safe_numeric_cast(field_accessor, op_value) + operator_map: dict[str, Callable[[Any, Any], ColumnElement[bool]]] = { + "gte": lambda a, v: a >= v, + "lte": lambda a, v: a <= v, + "gt": lambda a, v: a > v, + "lt": lambda a, v: a < v, + "ne": lambda a, v: a != v, + } + return operator_map[operator](safe_accessor, safe_value) + except Exception as e: + raise FilterError( + f"Failed to build numeric comparison condition for operator '{operator}' with value '{op_value}': {str(e)}" + ) from e + elif operator == "in": + if hasattr(op_value, "__iter__") and not isinstance(op_value, str | bytes): + # Handle wildcard in iterable - if present, matches everything, so no condition needed + if "*" in op_value: + return None + return field_accessor.in_([str(v) for v 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 in ("contains", "icontains"): + return field_accessor.ilike(f"%{op_value}%") + + return None + + +def _build_nested_metadata_conditions( + column: Any, metadata_dict: dict[str, Any] +) -> ColumnElement[bool] | None: + """ + Build conditions for nested metadata fields with comparison operators. + + Args: + column: SQLAlchemy JSONB column object + metadata_dict: Dictionary containing nested field conditions + + Returns: + Combined SQLAlchemy condition object or None + """ + conditions: list[ColumnElement[bool]] = [] + + for field_name, field_value in metadata_dict.items(): + if isinstance(field_value, dict) and any( + op in COMPARISON_OPERATORS + for op in field_value # pyright: ignore + ): + # This field has comparison operators + field_conditions: list[ColumnElement[bool]] = [] + for operator, op_value in field_value.items(): # pyright: ignore + condition = _build_comparison_condition( + column, + field_name, + operator, # pyright: ignore + op_value, + ) + if condition is not None: + field_conditions.append(condition) + + if field_conditions: + conditions.append( + field_conditions[0] + if len(field_conditions) == 1 + else and_(*field_conditions) + ) + else: + # Handle wildcard - matches everything, so no condition needed + if field_value == "*": + continue + # Regular field equality - use JSONB contains for nested object matching + conditions.append(column.contains({field_name: field_value})) + + # Combine all field conditions with AND + return _combine_conditions_with_and(conditions) + + +def _combine_conditions_with_and( + conditions: list[ColumnElement[bool]], +) -> ColumnElement[bool] | None: + """ + Combine a list of conditions with AND logic. + + Args: + conditions: List of SQLAlchemy condition objects + + Returns: + Combined condition object or None if no conditions + """ + if not conditions: + return None + elif len(conditions) == 1: + return conditions[0] + else: + return and_(*conditions) + + +def _build_comparison_conditions( + column: Any, column_name: str, comparisons: dict[str, Any] +) -> ColumnElement[bool] | None: + """ + Build comparison conditions for a single column. + + Args: + column: SQLAlchemy column object + column_name: Name of the column + comparisons: Dictionary of comparison operators and values + + Returns: + Combined SQLAlchemy condition object or None + """ + 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 + ) + + for operator, op_value in comparisons.items(): + # Validate that the operator is supported + if operator not in COMPARISON_OPERATORS: + raise FilterError(f"Unsupported comparison operator: {operator}") + + # Handle wildcard - matches everything, so no condition needed + if op_value == "*": + continue + + 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: + # if the operator is a numeric operator, the value must cast to a number + if operator in NUMERIC_OPERATORS: + try: + casted_value = float(op_value) + except ValueError: + raise FilterError( + f"Invalid numeric value: {op_value}. Expected a number, got {type(op_value).__name__}" + ) from None + else: + casted_value = op_value + + if operator == "gte": + condition = column >= casted_value + elif operator == "lte": + condition = column <= casted_value + elif operator == "gt": + condition = column > casted_value + elif operator == "lt": + condition = column < casted_value + elif operator == "ne": + condition = column != casted_value + elif operator == "in": + if hasattr(op_value, "__iter__") and not isinstance(op_value, str | bytes): + # Handle wildcard in iterable - if present, matches everything, so no condition needed + 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) + else: + condition = column.in_(list(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) + else: + # For text columns, use ILIKE + condition = column.ilike(f"%{op_value}%") + elif operator == "icontains": + # Case-insensitive contains for text columns + condition = column.ilike(f"%{op_value}%") + + if condition is not None: + conditions.append(condition) + + # Combine all conditions for this field with AND + if len(conditions) == 0: + return None + elif len(conditions) == 1: + return conditions[0] + else: + return and_(*conditions) + + +def _validate_datetime_string(value: str) -> datetime.datetime | None: + """ + Safely validate and parse a datetime string to prevent SQL injection. + + This function attempts to parse the datetime string using multiple common formats + to ensure it's a valid datetime before allowing it to be used in SQL queries. + + Args: + value: String value to validate as datetime + + Returns: + Parsed datetime object if valid, None if invalid + """ + # Strip whitespace + value = value.strip() + + # Try to parse with various common datetime formats + datetime_formats = [ + "%Y-%m-%d %H:%M:%S", # 2024-01-01 12:00:00 + "%Y-%m-%d %H:%M:%S.%f", # 2024-01-01 12:00:00.123456 + "%Y-%m-%dT%H:%M:%S", # 2024-01-01T12:00:00 (ISO format) + "%Y-%m-%dT%H:%M:%S.%f", # 2024-01-01T12:00:00.123456 + "%Y-%m-%dT%H:%M:%SZ", # 2024-01-01T12:00:00Z (UTC) + "%Y-%m-%dT%H:%M:%S.%fZ", # 2024-01-01T12:00:00.123456Z + "%Y-%m-%d", # 2024-01-01 + ] + + for fmt in datetime_formats: + try: + return datetime.datetime.strptime(value, fmt) + except ValueError: + continue + + # Try fromisoformat as a fallback (Python 3.7+) + try: + return datetime.datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError: + pass + + # Return None for invalid datetime - let the caller handle the error + return None diff --git a/tests/routes/test_messages.py b/tests/routes/test_messages.py index 092fa09e..88f1ccae 100644 --- a/tests/routes/test_messages.py +++ b/tests/routes/test_messages.py @@ -375,12 +375,12 @@ async def test_get_filtered_messages( response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/{test_session.name}/messages/list", - json={"filter": {"key": "value2"}}, + json={"filter": {"metadata": {"key": "value2"}}}, ) assert response.status_code == 200 data = response.json() assert "items" in data - assert len(data["items"]) > 0 + assert len(data["items"]) == 1 assert data["items"][0]["content"] == "Test message 2" assert data["items"][0]["peer_id"] == test_peer.name assert data["items"][0]["session_id"] == test_session.name @@ -427,10 +427,10 @@ async def test_get_filtered_messages_with_complex_filter( db_session.add(test_message3) await db_session.commit() - # Filter by multiple criteria + # Test old-style filter (backward compatibility) response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/{test_session.name}/messages/list", - json={"filter": {"priority": "high", "category": "technical"}}, + json={"filter": {"metadata": {"priority": "high", "category": "technical"}}}, ) assert response.status_code == 200 data = response.json() @@ -438,6 +438,40 @@ async def test_get_filtered_messages_with_complex_filter( # Should return messages 1 and 2 (both have high priority and technical category) assert len(data["items"]) >= 2 + # Test new-style filter with AND operator + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{test_session.name}/messages/list", + json={ + "filter": { + "AND": [ + {"metadata": {"priority": "high"}}, + {"metadata": {"category": "technical"}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + assert "items" in data + assert len(data["items"]) == 2 + + # Test OR filter to get high priority OR question type + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{test_session.name}/messages/list", + json={ + "filter": { + "OR": [ + {"metadata": {"priority": "high"}}, + {"metadata": {"type": "question"}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + assert "items" in data + assert len(data["items"]) == 3 # All messages should match + @pytest.mark.asyncio async def test_update_message( diff --git a/tests/routes/test_peers.py b/tests/routes/test_peers.py index 224097de..1c51f6d1 100644 --- a/tests/routes/test_peers.py +++ b/tests/routes/test_peers.py @@ -110,17 +110,29 @@ def test_get_peers(client: TestClient, sample_data: tuple[Workspace, Peer]): assert "items" in data assert len(data["items"]) > 0 - # Get peers with filter + # Get peers with simple filter (backward compatibility) response = client.post( f"/v2/workspaces/{test_workspace.name}/peers/list", - json={"filter": {"peer_key": "peer_value"}}, + json={"filter": {"metadata": {"peer_key": "peer_value"}}}, ) assert response.status_code == 200 data = response.json() assert "items" in data - assert len(data["items"]) >= 2 + assert len(data["items"]) == 2 assert data["items"][0]["metadata"]["peer_key"] == "peer_value" + # Test new filter with NOT operator + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"NOT": [{"metadata": {"peer_key": "peer_value2"}}]}}, + ) + assert response.status_code == 200 + data = response.json() + assert "items" in data + # Should find peers that don't have peer_key = "peer_value2" + # This includes the 2 peers with "peer_value" + the sample peer with empty metadata + assert len(data["items"]) == 3 + def test_get_peers_with_empty_filter( client: TestClient, sample_data: tuple[Workspace, Peer] @@ -275,33 +287,6 @@ def test_get_sessions_for_peer(client: TestClient, sample_data: tuple[Workspace, assert len(data["items"]) == 1 -def test_get_sessions_for_peer_with_is_active_filter( - client: TestClient, sample_data: tuple[Workspace, Peer] -): - """Test getting sessions for peer with is_active parameter""" - test_workspace, test_peer = sample_data - - # Create and then delete a session to have inactive session - session_name = str(generate_nanoid()) - client.post( - f"/v2/workspaces/{test_workspace.name}/sessions", - json={"id": session_name, "peer_names": {test_peer.name: {}}}, - ) - client.delete(f"/v2/workspaces/{test_workspace.name}/sessions/{session_name}") - - # Test getting inactive sessions - response = client.post( - f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/sessions", - json={"is_active": False}, - ) - assert response.status_code == 200 - data = response.json() - assert "items" in data - # Should find at least our deleted session - inactive_sessions = [s for s in data["items"] if not s["is_active"]] - assert len(inactive_sessions) > 0 - - def test_get_sessions_for_peer_with_empty_filter( client: TestClient, sample_data: tuple[Workspace, Peer] ): diff --git a/tests/routes/test_scoped_api.py b/tests/routes/test_scoped_api.py index 634d9043..7e7fa531 100644 --- a/tests/routes/test_scoped_api.py +++ b/tests/routes/test_scoped_api.py @@ -168,7 +168,7 @@ def test_get_peer_by_name_with_auth( # Use POST /list endpoint to get peers response = auth_client.post( f"/v2/workspaces/{test_workspace.name}/peers/list", - json={"filter": {"name": test_peer.name}}, + json={"filter": {"id": test_peer.name}}, ) # Admin JWT or JWT with matching workspace should be allowed diff --git a/tests/routes/test_sessions.py b/tests/routes/test_sessions.py index 341894bb..0e3e3349 100644 --- a/tests/routes/test_sessions.py +++ b/tests/routes/test_sessions.py @@ -146,7 +146,7 @@ def test_create_session_with_too_many_peers( session_response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/list", - json={"filter": {"name": "test_session"}}, + json={"filter": {"id": "test_session"}}, ) assert session_response.status_code == 200 assert len(session_response.json()["items"]) == 0 @@ -186,42 +186,16 @@ def test_get_sessions(client: TestClient, sample_data: tuple[Workspace, Peer]): assert data["workspace_id"] == test_workspace.name response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/list", - json={"filter": {"test_key": "test_value"}}, + json={"filter": {"metadata": {"test_key": "test_value"}}}, ) assert response.status_code == 200 data = response.json() assert "items" in data - assert len(data["items"]) > 0 + assert len(data["items"]) == 1 assert data["items"][0]["metadata"] == {"test_key": "test_value"} assert data["items"][0]["workspace_id"] == test_workspace.name -def test_get_sessions_with_is_active_filter( - client: TestClient, sample_data: tuple[Workspace, Peer] -): - """Test session listing with is_active parameter""" - test_workspace, test_peer = sample_data - - # Create and then delete a session to have inactive session - session_id = str(generate_nanoid()) - client.post( - f"/v2/workspaces/{test_workspace.name}/sessions", - json={"id": session_id, "peer_names": {test_peer.name: {}}}, - ) - client.delete(f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}") - - # Test getting inactive sessions - response = client.post( - f"/v2/workspaces/{test_workspace.name}/sessions/list", json={"is_active": False} - ) - assert response.status_code == 200 - data = response.json() - assert "items" in data - # Should find at least our deleted session - inactive_sessions = [s for s in data["items"] if not s["is_active"]] - assert len(inactive_sessions) > 0 - - def test_get_sessions_with_empty_filter( client: TestClient, sample_data: tuple[Workspace, Peer] ): @@ -382,7 +356,7 @@ def test_delete_session(client: TestClient, sample_data: tuple[Workspace, Peer]) # Check that session is marked as inactive response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/list", - json={"is_active": False}, + json={"filter": {"is_active": False}}, ) data = response.json() # Find our session in the inactive sessions diff --git a/tests/routes/test_workspaces.py b/tests/routes/test_workspaces.py index 8cf2b4f1..52f8a10f 100644 --- a/tests/routes/test_workspaces.py +++ b/tests/routes/test_workspaces.py @@ -91,12 +91,12 @@ async def test_get_all_workspaces(client: TestClient): response = client.post( "/v2/workspaces/list", - json={"filter": {"test_key": "test_value"}}, + json={"filter": {"metadata": {"test_key": "test_value"}}}, ) assert response.status_code == 200 data = response.json() assert "items" in data - assert len(data["items"]) > 0 + assert len(data["items"]) == 1 assert data["items"][0]["metadata"] == {"test_key": "test_value"} diff --git a/tests/test_advanced_filters.py b/tests/test_advanced_filters.py new file mode 100644 index 00000000..73117cd5 --- /dev/null +++ b/tests/test_advanced_filters.py @@ -0,0 +1,2459 @@ +""" +Tests for advanced filter functionality including logical operators, +comparison operators, and wildcards across multiple models. +""" + +from datetime import datetime, timedelta, timezone +from typing import Any + +import pytest +from fastapi.testclient import TestClient +from nanoid import generate as generate_nanoid + +from src.models import Peer, Workspace + + +@pytest.mark.parametrize( + "filter_config,expected_peer_indices,description", + [ + ( + { + "AND": [ + {"metadata": {"role": "admin"}}, + {"metadata": {"department": "engineering"}}, + ] + }, + [0], # Only peer1 (admin + engineering) + "peers who are admin AND in engineering", + ), + ( + { + "OR": [ + {"metadata": {"role": "admin"}}, + {"metadata": {"department": "engineering"}}, + ] + }, + [ + 0, + 1, + 2, + ], # All three peers match (peer1: admin+eng, peer2: user+eng, peer3: admin+sales) + "peers who are admin OR in engineering", + ), + ( + {"NOT": [{"metadata": {"role": "admin"}}]}, + [1], # Only peer2 (user) is not admin + "peers who are NOT admin", + ), + ], +) +@pytest.mark.asyncio +async def test_logical_operators_and_filters( + client: TestClient, + sample_data: tuple[Workspace, Peer], + filter_config: dict[str, Any], + expected_peer_indices: list[int], + description: str, +): + """Test AND, OR, NOT logical operators in filters""" + test_workspace, _test_peer = sample_data + + # Create multiple peers with different metadata + peer_names = [str(generate_nanoid()) for _ in range(3)] + peer_configs = [ + { + "name": peer_names[0], + "metadata": { + "role": "admin", + "department": "engineering", + "level": "senior", + }, + }, + { + "name": peer_names[1], + "metadata": { + "role": "user", + "department": "engineering", + "level": "junior", + }, + }, + { + "name": peer_names[2], + "metadata": {"role": "admin", "department": "sales", "level": "senior"}, + }, + ] + + # Create all peers + for peer_config in peer_configs: + client.post(f"/v2/workspaces/{test_workspace.name}/peers", json=peer_config) + + # Test the filter configuration, but only consider the peers we created + combined_filter = { + "AND": [ + filter_config, + {"id": {"in": peer_names}}, # Only include the peers we created + ] + } + + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": combined_filter}, + ) + assert response.status_code == 200, f"Failed testing {description}" + + data = response.json() + found_names = [item["id"] for item in data["items"]] + expected_names = [peer_names[i] for i in expected_peer_indices] + + assert len(found_names) == len(expected_names), ( + f"Expected {len(expected_names)} peers for {description}, got {len(found_names)}" + ) + + for expected_name in expected_names: + assert expected_name in found_names, ( + f"Expected peer {expected_name} in results for {description}" + ) + + # Verify unexpected peers are not included + for i, peer_name in enumerate(peer_names): + if i not in expected_peer_indices: + assert peer_name not in found_names, ( + f"Unexpected peer {peer_name} found in results for {description}" + ) + + +@pytest.mark.parametrize( + "filter_config,expected_message_indices,description", + [ + ( + {"metadata": {"score": {"gte": 5}}}, + [0, 1], # Messages with score 10 and 5 + "gte (greater than or equal) operator", + ), + ( + {"metadata": {"score": {"lte": 5}}}, + [1, 2], # Messages with score 5 and 1 + "lte (less than or equal) operator", + ), + ( + {"metadata": {"score": {"gt": 5}}}, + [0], # Only message with score 10 + "gt (greater than) operator", + ), + ( + {"metadata": {"score": {"lt": 5}}}, + [2], # Only message with score 1 + "lt (less than) operator", + ), + ( + {"metadata": {"score": {"ne": 5}}}, + [0, 2], # Messages with score 10 and 1 + "ne (not equal) operator", + ), + ( + {"metadata": {"category": {"in": ["high", "low"]}}}, + [0, 2], # Messages with category "high" and "low" + "in operator", + ), + ], +) +@pytest.mark.asyncio +async def test_comparison_operators_filters( + client: TestClient, + sample_data: tuple[Workspace, Peer], + filter_config: dict[str, Any], + expected_message_indices: list[int], + description: str, +): + """Test comparison operators (gte, lte, gt, lt, ne, in, contains, icontains)""" + test_workspace, test_peer = sample_data + + # Create session with messages containing different metadata + session_id = str(generate_nanoid()) + session_response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + assert session_response.status_code == 200 + + # Create messages with numeric metadata for comparison tests + message_configs = [ + { + "content": "Message with score 10", + "peer_id": test_peer.name, + "metadata": { + "score": 10, + "category": "high", + "tags": ["important", "urgent"], + }, + }, + { + "content": "Message with score 5", + "peer_id": test_peer.name, + "metadata": {"score": 5, "category": "medium", "tags": ["normal"]}, + }, + { + "content": "Message with score 1", + "peer_id": test_peer.name, + "metadata": {"score": 1, "category": "low", "tags": ["minor"]}, + }, + ] + + messages_response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={"messages": message_configs}, + ) + assert messages_response.status_code == 200 + + # Test the filter configuration + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": filter_config}, + ) + assert response.status_code == 200, f"Failed testing {description}" + + data = response.json() + expected_contents = [ + message_configs[i]["content"] for i in expected_message_indices + ] + found_contents = [item["content"] for item in data["items"]] + + assert len(found_contents) == len(expected_contents), ( + f"Expected {len(expected_contents)} messages for {description}, got {len(found_contents)}" + ) + + for expected_content in expected_contents: + assert expected_content in found_contents, ( + f"Expected message '{expected_content}' in results for {description}" + ) + + # Verify unexpected messages are not included + for i, message_config in enumerate(message_configs): + if i not in expected_message_indices: + assert message_config["content"] not in found_contents, ( + f"Unexpected message '{message_config['content']}' found in results for {description}" + ) + + +@pytest.mark.asyncio +async def test_wildcard_filters( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test wildcard (*) filters that match everything for a field""" + test_workspace, _test_peer = sample_data + + # Create peers with different names + peer1_name = str(generate_nanoid()) + peer2_name = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer1_name, "metadata": {"type": "bot"}}, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer2_name, "metadata": {"type": "human"}}, + ) + + # Test wildcard for peer_id field + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "AND": [ + {"id": "*"}, + {"metadata": {"type": "bot"}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + # Should find the bot peer (wildcard doesn't filter anything) + found_names = [item["id"] for item in data["items"]] + assert peer1_name in found_names + + # Test wildcard in comparison operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "id": {"in": ["*"]} # Wildcard in comparison should also match all + } + }, + ) + assert response.status_code == 200 + data = response.json() + # Since wildcard should be ignored, this should return all peers + assert len(data["items"]) >= 3 # At least the 3 peers we know about + + +@pytest.mark.asyncio +async def test_complex_nested_filters( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test complex nested logical operations""" + test_workspace, test_peer = sample_data + + # Create session and messages for complex filtering + session_id = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Create messages with various metadata combinations + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Priority urgent task", + "peer_id": test_peer.name, + "metadata": { + "priority": "urgent", + "status": "open", + "assignee": "alice", + }, + }, + { + "content": "Normal task for bob", + "peer_id": test_peer.name, + "metadata": { + "priority": "normal", + "status": "open", + "assignee": "bob", + }, + }, + { + "content": "Completed urgent task", + "peer_id": test_peer.name, + "metadata": { + "priority": "urgent", + "status": "completed", + "assignee": "alice", + }, + }, + { + "content": "Low priority task", + "peer_id": test_peer.name, + "metadata": { + "priority": "low", + "status": "open", + "assignee": "charlie", + }, + }, + ] + }, + ) + + # Complex filter: (urgent OR normal priority) AND open status AND NOT assigned to charlie + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={ + "filter": { + "AND": [ + { + "OR": [ + {"metadata": {"priority": "urgent"}}, + {"metadata": {"priority": "normal"}}, + ] + }, + {"metadata": {"status": "open"}}, + {"NOT": [{"metadata": {"assignee": "charlie"}}]}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 2 # Should match first two messages + + # Verify the correct messages were returned + contents = [item["content"] for item in data["items"]] + assert "Priority urgent task" in contents + assert "Normal task for bob" in contents + assert "Completed urgent task" not in contents # Wrong status + assert "Low priority task" not in contents # Wrong assignee and priority + + +@pytest.mark.asyncio +async def test_filters_across_different_models( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that filters work consistently across different models""" + test_workspace, test_peer = sample_data + + # Test workspace filters + workspace_name = str(generate_nanoid()) + client.post( + "/v2/workspaces", + json={ + "name": workspace_name, + "metadata": {"environment": "production", "version": "2.0"}, + }, + ) + + response = client.post( + "/v2/workspaces/list", + json={ + "filter": { + "AND": [ + {"metadata": {"environment": "production"}}, + {"metadata": {"version": {"gte": "2.0"}}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert workspace_name in found_names + + # Test session filters with comparison operators + session1_id = str(generate_nanoid()) + session2_id = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session1_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"duration": 30, "type": "meeting"}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session2_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"duration": 60, "type": "workshop"}, + }, + ) + + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={ + "filter": { + "OR": [ + {"metadata": {"duration": {"gte": 45}}}, + {"metadata": {"type": "meeting"}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + found_sessions = [item["id"] for item in data["items"]] + assert session1_id in found_sessions # Matches type=meeting + assert session2_id in found_sessions # Matches duration>=45 + + +@pytest.mark.asyncio +async def test_filter_edge_cases( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test edge cases and error handling for filters""" + test_workspace, test_peer = sample_data + + # Test empty logical operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "AND": [] # Empty AND should not crash + } + }, + ) + assert response.status_code == 200 + + # Test nested empty operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"OR": [{"AND": []}, {"id": test_peer.name}]}}, + ) + assert response.status_code == 200 + + # Test filter with non-existent columns (should be ignored gracefully) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"non_existent_column": "value"}}, + ) + assert response.status_code == 422 + + # Test mixed wildcards and regular values + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "AND": [ + {"id": "*"}, # Wildcard + {"id": test_peer.name}, # Regular value + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + # Should find the specific peer since AND combines conditions + found_names = [item["id"] for item in data["items"]] + assert test_peer.name in found_names + + +@pytest.mark.asyncio +async def test_backward_compatibility( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that old simple filter format still works""" + test_workspace, _test_peer = sample_data + + # Create peer with metadata + peer_name = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer_name, + "metadata": {"role": "admin", "department": "engineering"}, + }, + ) + + # Test old-style simple equality filter + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"role": "admin"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer_name in found_names + + # Test multiple field simple filter (implicit AND) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"role": "admin"}, "id": peer_name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == peer_name + + +@pytest.mark.asyncio +async def test_range_queries_with_dates( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test range queries that might be used with date fields""" + test_workspace, test_peer = sample_data + + # Create sessions with date-like metadata + session1_id = str(generate_nanoid()) + session2_id = str(generate_nanoid()) + session3_id = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session1_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"created_date": "2024-01-15", "priority": 5}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session2_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"created_date": "2024-02-20", "priority": 3}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session3_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"created_date": "2024-03-10", "priority": 8}, + }, + ) + + # Test date range query + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={ + "filter": { + "metadata": {"created_date": {"gte": "2024-02-01", "lte": "2024-02-28"}} + } + }, + ) + assert response.status_code == 200 + data = response.json() + found_sessions = [item["id"] for item in data["items"]] + assert session2_id in found_sessions + assert session1_id not in found_sessions + assert session3_id not in found_sessions + + # Test combining date and numeric filters + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={ + "filter": { + "AND": [ + {"metadata": {"created_date": {"gte": "2024-01-01"}}}, + {"metadata": {"priority": {"gt": 4}}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + found_sessions = [item["id"] for item in data["items"]] + assert session1_id in found_sessions # priority 5 + assert session3_id in found_sessions # priority 8 + assert session2_id not in found_sessions # priority 3 + + +@pytest.mark.asyncio +async def test_all_workspace_columns_filtering(client: TestClient): + """Test filtering by all available workspace columns""" + # Create additional workspaces for testing + workspace1_name = str(generate_nanoid()) + workspace2_name = str(generate_nanoid()) + + client.post( + "/v2/workspaces", + json={ + "name": workspace1_name, + "metadata": {"env": "dev", "version": "1.0", "active": True}, + "configuration": {"max_sessions": 100, "timeout": 30}, + }, + ) + client.post( + "/v2/workspaces", + json={ + "name": workspace2_name, + "metadata": {"env": "prod", "version": "2.0", "active": False}, + "configuration": {"max_sessions": 500, "timeout": 60}, + }, + ) + + # Test filtering by id (maps to name internally) + response = client.post( + "/v2/workspaces/list", + json={"filter": {"id": workspace1_name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == workspace1_name + + # Test filtering by id with comparison operators + response = client.post( + "/v2/workspaces/list", + json={"filter": {"id": {"in": [workspace1_name, workspace2_name]}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert workspace1_name in found_names + assert workspace2_name in found_names + + # Test filtering by metadata (maps to h_metadata internally) + response = client.post( + "/v2/workspaces/list", + json={"filter": {"metadata": {"env": "dev"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert workspace1_name in found_names + assert workspace2_name not in found_names + + # Test filtering by metadata with comparison operators + response = client.post( + "/v2/workspaces/list", + json={"filter": {"metadata": {"version": {"gte": "2.0"}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert workspace2_name in found_names + assert workspace1_name not in found_names + + # Test filtering by created_at (datetime field) + response = client.post( + "/v2/workspaces/list", + json={"filter": {"created_at": {"gte": "2020-01-01"}}}, + ) + assert response.status_code == 200 + # Should return all workspaces created after 2020 + + +@pytest.mark.asyncio +async def test_all_peer_columns_filtering( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test filtering by all available peer columns""" + test_workspace, _test_peer = sample_data + + # Create additional peers for testing + peer1_name = str(generate_nanoid()) + peer2_name = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer1_name, + "metadata": {"role": "user", "level": 1, "active": True}, + "configuration": {"notifications": True, "theme": "dark"}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer2_name, + "metadata": {"role": "admin", "level": 5, "active": False}, + "configuration": {"notifications": False, "theme": "light"}, + }, + ) + + # Test filtering by id (maps to name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"id": peer1_name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == peer1_name + + # Test filtering by workspace_id (maps to workspace_name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"workspace_id": test_workspace.name}}, + ) + assert response.status_code == 200 + data = response.json() + # Should return all peers in the workspace + assert len(data["items"]) >= 3 # test_peer + peer1 + peer2 + + # Test filtering by metadata + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"role": "admin"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer2_name in found_names + assert peer1_name not in found_names + + # Test filtering by metadata with comparison operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"level": {"gte": 3}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer2_name in found_names + assert peer1_name not in found_names + + # Test filtering by created_at + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"created_at": {"gte": "2020-01-01"}}}, + ) + assert response.status_code == 200 + # Should return all peers created after 2020 + + +@pytest.mark.asyncio +async def test_all_session_columns_filtering( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test filtering by all available session columns""" + test_workspace, test_peer = sample_data + + # Create additional sessions for testing + session1_id = str(generate_nanoid()) + session2_id = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session1_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"type": "chat", "priority": 1, "active": True}, + "configuration": {"auto_save": True, "timeout": 30}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session2_id, + "peer_names": {test_peer.name: {}}, + "metadata": {"type": "support", "priority": 5, "active": False}, + "configuration": {"auto_save": False, "timeout": 60}, + }, + ) + + # Test filtering by id (maps to name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"id": session1_id}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == session1_id + + # Test filtering by workspace_id (maps to workspace_name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"workspace_id": test_workspace.name}}, + ) + assert response.status_code == 200 + data = response.json() + # Should return all sessions in the workspace + assert len(data["items"]) >= 2 + + # Test filtering by is_active (boolean field) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"is_active": True}}, + ) + assert response.status_code == 200 + data = response.json() + found_sessions = [item["id"] for item in data["items"]] + assert session1_id in found_sessions + # session2 should not be in results since is_active=False by default in get_sessions + + # Test filtering by metadata + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"metadata": {"type": "chat"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_sessions = [item["id"] for item in data["items"]] + assert session1_id in found_sessions + + # Test filtering by metadata with comparison operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"metadata": {"priority": {"gte": 3}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_sessions = [item["id"] for item in data["items"]] + assert session2_id in found_sessions + assert session1_id not in found_sessions + + # Test filtering by created_at + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"created_at": {"gte": "2020-01-01"}}}, + ) + assert response.status_code == 200 + # Should return all sessions created after 2020 + + +@pytest.mark.asyncio +async def test_all_message_columns_filtering( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test filtering by all available message columns""" + test_workspace, test_peer = sample_data + + # Create session and messages for testing + session_id = str(generate_nanoid()) + session_response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + assert session_response.status_code == 200 + + # Create messages with various data + messages_response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Hello world message", + "peer_id": test_peer.name, + "metadata": {"type": "greeting", "priority": 1, "urgent": True}, + }, + { + "content": "Technical support request", + "peer_id": test_peer.name, + "metadata": {"type": "support", "priority": 5, "urgent": False}, + }, + { + "content": "Follow up message", + "peer_id": test_peer.name, + "metadata": {"type": "followup", "priority": 3, "urgent": True}, + }, + ] + }, + ) + assert messages_response.status_code == 200 + + # Test filtering by session_id (maps to session_name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"session_id": session_id}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 3 # All messages in the session + + # Test filtering by peer_id (maps to peer_name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"peer_id": test_peer.name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 3 # All messages from the peer + + # Test filtering by workspace_id (maps to workspace_name internally) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"workspace_id": test_workspace.name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 3 # All messages in the workspace + + # Test filtering by content (text field) (not allowed) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"content": "Hello world message"}}, + ) + assert response.status_code == 422 + + # Test filtering by content with contains operator (not allowed) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"content": {"contains": "support"}}}, + ) + assert response.status_code == 422 + + # Test filtering by metadata + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"type": "greeting"}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + + # Test filtering by metadata with comparison operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"priority": {"gte": 3}}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 2 # priority 5 and 3 + + # Test filtering by token_count (integer field) - this should exist after message creation + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"token_count": {"gte": 0}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 3 # All messages should have token_count >= 0 + + # Test filtering by created_at + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"created_at": {"gte": "2020-01-01"}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 3 # All messages created after 2020 + + +@pytest.mark.asyncio +async def test_id_field_interpolation_consistency(client: TestClient): + """Test that id field interpolation works consistently across all models""" + # Create test data + workspace_name = str(generate_nanoid()) + peer_name = str(generate_nanoid()) + session_id = str(generate_nanoid()) + + # Create workspace + client.post( + "/v2/workspaces", + json={"name": workspace_name, "metadata": {"test": "value"}}, + ) + + # Create peer + client.post( + f"/v2/workspaces/{workspace_name}/peers", + json={"name": peer_name, "metadata": {"test": "value"}}, + ) + + # Create session + client.post( + f"/v2/workspaces/{workspace_name}/sessions", + json={ + "id": session_id, + "peer_names": {peer_name: {}}, + "metadata": {"test": "value"}, + }, + ) + + # Create message + messages_response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Test message", + "peer_id": peer_name, + "metadata": {"test": "value"}, + } + ] + }, + ) + message_id = messages_response.json()[0]["id"] + + # Test that filtering by "id" returns the expected items for each model + + # Workspace: id should map to name + response = client.post( + "/v2/workspaces/list", + json={"filter": {"id": workspace_name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == workspace_name + + # Peer: id should map to name + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/list", + json={"filter": {"id": peer_name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == peer_name + + # Session: id should map to name + response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/list", + json={"filter": {"id": session_id}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert data["items"][0]["id"] == session_id + + # Message: id is not allowed to be filtered on + response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/{session_id}/messages/list", + json={"filter": {"id": message_id}}, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_foreign_key_field_interpolation( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that foreign key _id fields map to _name fields correctly""" + test_workspace, test_peer = sample_data + + # Create test data + peer_name = str(generate_nanoid()) + session_id = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer_name, "metadata": {"role": "test"}}, + ) + + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={ + "id": session_id, + "peer_names": {peer_name: {}}, + "metadata": {"type": "test"}, + }, + ) + + # Create message + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Test message", + "peer_id": peer_name, + "metadata": {"test": "value"}, + } + ] + }, + ) + + # Test workspace_id filtering for peers (maps to workspace_name) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"workspace_id": test_workspace.name}}, + ) + assert response.status_code == 200 + data = response.json() + peer_names = [item["id"] for item in data["items"]] + assert peer_name in peer_names + assert test_peer.name in peer_names + + # Test workspace_id filtering for sessions (maps to workspace_name) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/list", + json={"filter": {"workspace_id": test_workspace.name}}, + ) + assert response.status_code == 200 + data = response.json() + session_names = [item["id"] for item in data["items"]] + assert session_id in session_names + + # Test session_id filtering for messages (maps to session_name) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"session_id": session_id}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + + # Test peer_id filtering for messages (maps to peer_name) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"peer_id": peer_name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + + # Test workspace_id filtering for messages (maps to workspace_name) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"workspace_id": test_workspace.name}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + + +@pytest.mark.asyncio +async def test_metadata_field_interpolation(client: TestClient): + """Test that metadata field maps to h_metadata column correctly""" + # Create test data with metadata + workspace_name = str(generate_nanoid()) + peer_name = str(generate_nanoid()) + session_id = str(generate_nanoid()) + + client.post( + "/v2/workspaces", + json={ + "name": workspace_name, + "metadata": {"env": "test", "version": "1.0", "active": True}, + }, + ) + + client.post( + f"/v2/workspaces/{workspace_name}/peers", + json={ + "name": peer_name, + "metadata": {"role": "user", "level": 5, "premium": True}, + }, + ) + + client.post( + f"/v2/workspaces/{workspace_name}/sessions", + json={ + "id": session_id, + "peer_names": {peer_name: {}}, + "metadata": {"type": "chat", "duration": 120, "archived": False}, + }, + ) + + client.post( + f"/v2/workspaces/{workspace_name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Test message", + "peer_id": peer_name, + "metadata": { + "sentiment": "positive", + "confidence": 0.95, + "flagged": False, + }, + } + ] + }, + ) + + # Test metadata filtering for all models + + # Workspace metadata + response = client.post( + "/v2/workspaces/list", + json={"filter": {"metadata": {"env": "test"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert workspace_name in found_names + + # Peer metadata + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/list", + json={"filter": {"metadata": {"role": "user"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer_name in found_names + + # Session metadata + response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/list", + json={"filter": {"metadata": {"type": "chat"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert session_id in found_names + + # Message metadata + response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"sentiment": "positive"}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + + # Test metadata with comparison operators + + # Numeric metadata comparison + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/list", + json={"filter": {"metadata": {"level": {"gte": 3}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer_name in found_names + + # Boolean metadata comparison + response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/list", + json={"filter": {"metadata": {"archived": {"ne": True}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert session_id in found_names + + # Float metadata comparison + response = client.post( + f"/v2/workspaces/{workspace_name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"confidence": {"gte": 0.9}}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + + +@pytest.mark.asyncio +async def test_nonexistent_columns_ignored_gracefully( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that filtering by non-existent columns is ignored gracefully""" + test_workspace, _test_peer = sample_data + + # Create test data + peer_name = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer_name, "metadata": {"role": "test"}}, + ) + + # Test filtering by non-existent columns + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"nonexistent_column": "value"}}, + ) + assert response.status_code == 422 + + # Test combining real and non-existent columns + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "AND": [ + {"metadata": {"role": "test"}}, # Real filter + {"fake_column": "fake_value"}, # Non-existent filter + ] + } + }, + ) + assert response.status_code == 422 + + # Test with complex nested filters containing non-existent columns + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "OR": [ + {"nonexistent_field": "value"}, + {"metadata": {"role": "test"}}, + ] + } + }, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_not_logic_correctness( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that NOT logic works correctly - THIS WILL LIKELY FAIL due to broken NOT logic""" + test_workspace, _test_peer = sample_data + + # Create peers with specific metadata + peer1_name = str(generate_nanoid()) + peer2_name = str(generate_nanoid()) + peer3_name = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer1_name, + "metadata": {"role": "admin", "department": "engineering"}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer2_name, + "metadata": {"role": "user", "department": "engineering"}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer3_name, "metadata": {"role": "admin", "department": "sales"}}, + ) + + # Test multiple NOT conditions - this exposes the bug + # User expectation: NOT admin AND NOT engineering = exclude admin users AND exclude engineering users + # Current broken code: NOT(admin AND engineering) = exclude users who are BOTH admin AND engineering + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "NOT": [ + {"metadata": {"role": "admin"}}, + {"metadata": {"department": "engineering"}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + + # This assertion will FAIL with current broken logic + # We expect only peers that are NOT admin AND NOT in engineering + # With broken logic, we get peers that are not (admin AND engineering) = peer2 and peer3 + # With correct logic, we should get only peers that are (NOT admin) AND (NOT engineering) = none of our test peers + assert peer1_name not in found_names # admin + engineering - should be excluded + assert ( + peer2_name not in found_names + ) # user + engineering - should be excluded (engineering) + assert peer3_name not in found_names # admin + sales - should be excluded (admin) + + +@pytest.mark.asyncio +async def test_jsonb_type_casting_edge_cases( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test edge cases in JSONB type casting for numeric comparisons""" + test_workspace, test_peer = sample_data + + session_id = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Create messages with various data types in metadata + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Boolean true", + "peer_id": test_peer.name, + "metadata": {"active": True, "score": 10}, + }, + { + "content": "Boolean false", + "peer_id": test_peer.name, + "metadata": {"active": False, "score": 5.5}, + }, + { + "content": "String number", + "peer_id": test_peer.name, + "metadata": {"active": "true", "score": "15"}, + }, + { + "content": "Large number", + "peer_id": test_peer.name, + "metadata": {"active": True, "score": 999999999}, + }, + ] + }, + ) + + # Test boolean comparisons - PostgreSQL stores booleans as "true"/"false" strings + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"active": {"ne": False}}}}, + ) + assert response.status_code == 200 + data = response.json() + # This might fail if boolean casting isn't handled correctly + assert len(data["items"]) == 3 # All except the false one + + # Test string vs numeric comparison + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"score": {"gt": 10}}}}, + ) + assert response.status_code == 200 + data = response.json() + # Should handle both numeric 15 (from string) and 999999999 + contents = [item["content"] for item in data["items"]] + assert "String number" in contents or "Large number" in contents + + # Test large number handling + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"score": {"gte": 999999999}}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert "Large number" in data["items"][0]["content"] + + +@pytest.mark.asyncio +async def test_real_datetime_column_filtering( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test filtering on actual datetime columns (not just metadata strings)""" + test_workspace, _test_peer = sample_data + + # Create peers at different times (this uses created_at datetime column) + peer1_name = str(generate_nanoid()) + peer2_name = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer1_name, "metadata": {"created": "early"}}, + ) + + # Wait a bit to ensure different timestamps + import time + + time.sleep(0.1) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer2_name, "metadata": {"created": "late"}}, + ) + + # Test datetime filtering with various formats + now = datetime.now(timezone.utc) + one_minute_ago = now - timedelta(minutes=1) + + # Test ISO format + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"created_at": {"gte": one_minute_ago.isoformat()}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer1_name in found_names + assert peer2_name in found_names + + # Test date-only format + today = datetime.now(timezone.utc).date().isoformat() + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"created_at": {"gte": today}}}, + ) + assert response.status_code == 200 + + +@pytest.mark.asyncio +async def test_invalid_datetime_handling( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test handling of invalid datetime strings""" + test_workspace, _test_peer = sample_data + + # Test malicious datetime strings (should be rejected by validation) + malicious_datetimes = [ + "2024-01-01'; DROP TABLE peers; --", + "2024-01-01 OR 1=1", + "'; SELECT * FROM users; --", + "2024-01-01' UNION SELECT password FROM auth", + ] + + for malicious_dt in malicious_datetimes: + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"created_at": {"gte": malicious_dt}}}, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_nested_jsonb_filtering( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test complex nested JSONB object filtering""" + test_workspace, test_peer = sample_data + + session_id = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Create messages with deeply nested metadata + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Complex nested data", + "peer_id": test_peer.name, + "metadata": { + "user": { + "profile": {"level": 5, "premium": True}, + "settings": {"notifications": True}, + }, + "tags": ["important", "urgent"], + "scores": [1, 2, 3, 4, 5], + }, + }, + { + "content": "Simple data", + "peer_id": test_peer.name, + "metadata": { + "user": {"profile": {"level": 1, "premium": False}}, + "tags": ["normal"], + }, + }, + ] + }, + ) + + # Test nested object filtering - this might not work as expected + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"user": {"profile": {"premium": True}}}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 + assert "Complex nested data" in data["items"][0]["content"] + + +@pytest.mark.asyncio +async def test_multiple_operators_same_field( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test multiple comparison operators on the same field""" + test_workspace, test_peer = sample_data + + session_id = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Create messages with scores for range testing + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Low score", + "peer_id": test_peer.name, + "metadata": {"score": 1}, + }, + { + "content": "Mid score", + "peer_id": test_peer.name, + "metadata": {"score": 5}, + }, + { + "content": "High score", + "peer_id": test_peer.name, + "metadata": {"score": 10}, + }, + { + "content": "Very high", + "peer_id": test_peer.name, + "metadata": {"score": 15}, + }, + ] + }, + ) + + # Test range query: score >= 3 AND score <= 8 + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"score": {"gte": 3, "lte": 8}}}}, + ) + assert response.status_code == 200 + data = response.json() + assert len(data["items"]) == 1 # Only "Mid score" with score 5 + assert "Mid score" in data["items"][0]["content"] + + +@pytest.mark.asyncio +async def test_empty_and_null_filter_handling( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test handling of empty and null filters""" + test_workspace, _test_peer = sample_data + + # Test completely empty filter + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", json={"filter": {}} + ) + assert response.status_code == 200 + data = response.json() + # Should return all peers + assert len(data["items"]) >= 1 + + # Test null filter + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", json={"filter": None} + ) + assert response.status_code == 200 + + # Test empty comparison dict + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"role": {}}}}, + ) + assert response.status_code == 200 + + +@pytest.mark.asyncio +async def test_unicode_and_special_characters( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test filtering with unicode and special characters""" + test_workspace, test_peer = sample_data + + # Create peer with unicode metadata + # NOTE: peer names are validated to only contain alphanumerics + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": test_peer.name, + "metadata": { + "description": "Héllo Wörld! 你好世界 🌍", + "tags": ["spëcial", "ünîcode", "émojis🎉"], + }, + }, + ) + + # Test unicode in metadata contains + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"description": {"icontains": "wörld"}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert test_peer.name in found_names + + +@pytest.mark.asyncio +async def test_case_sensitivity_edge_cases( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test case sensitivity in various contexts""" + test_workspace, _test_peer = sample_data + + # Create peers with mixed case data + peer1_name = str(generate_nanoid()) + peer2_name = str(generate_nanoid()) + + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer1_name, + "metadata": {"Role": "Admin", "Department": "ENGINEERING"}, + }, + ) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={ + "name": peer2_name, + "metadata": {"role": "admin", "department": "engineering"}, + }, + ) + + # Test exact case matching (should be case sensitive for JSONB) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"Role": "Admin"}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer1_name in found_names + assert peer2_name not in found_names # Different case + + # Test icontains for case insensitive search + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": {"metadata": {"Department": {"icontains": "engineering"}}}}, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + assert peer1_name in found_names # Should match ENGINEERING + + +@pytest.mark.asyncio +async def test_malformed_filter_structures( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test handling of malformed filter structures""" + test_workspace, _test_peer = sample_data + + malformed_filters = [ + # Invalid logical operator structures + {"AND": "not_a_list"}, + {"OR": {"invalid": "structure"}}, + {"NOT": None}, + # Invalid comparison structures + {"metadata": {"score": {"gte": [1, 2, 3]}}}, # Array instead of single value + # Mixed valid/invalid + {"AND": [{"valid_field": "value"}, "invalid_structure"]}, + ] + + for malformed_filter in malformed_filters: + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": malformed_filter}, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_performance_with_complex_filters( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test performance with very complex nested filters""" + test_workspace, _test_peer = sample_data + + # Create a very complex filter structure + complex_filter = { + "AND": [ + { + "OR": [ + {"metadata": {"type": "A"}}, + {"metadata": {"type": "B"}}, + {"metadata": {"type": "C"}}, + ] + }, + { + "NOT": [ + { + "AND": [ + {"metadata": {"status": "inactive"}}, + {"metadata": {"priority": {"lt": 5}}}, + ] + } + ] + }, + { + "OR": [ + {"metadata": {"score": {"gte": 80}}}, + { + "AND": [ + {"metadata": {"premium": True}}, + {"metadata": {"level": {"in": [1, 2, 3, 4, 5]}}}, + ] + }, + ] + }, + ] + } + + # This should complete without timeout + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": complex_filter}, + ) + assert response.status_code == 200 + # Should return results in reasonable time (this is more of a performance test) + + +@pytest.mark.asyncio +async def test_filter_precedence_and_grouping( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that filter precedence and grouping works as expected""" + test_workspace, _test_peer = sample_data + + # Create test data to verify precedence + peers_data = [ + { + "name": str(generate_nanoid()), + "metadata": {"role": "admin", "dept": "eng", "level": 1}, + }, + { + "name": str(generate_nanoid()), + "metadata": {"role": "user", "dept": "eng", "level": 2}, + }, + { + "name": str(generate_nanoid()), + "metadata": {"role": "admin", "dept": "sales", "level": 3}, + }, + { + "name": str(generate_nanoid()), + "metadata": {"role": "user", "dept": "sales", "level": 4}, + }, + ] + + for peer_data in peers_data: + client.post(f"/v2/workspaces/{test_workspace.name}/peers", json=peer_data) + + # Test: (admin OR user) AND (eng OR high level) + # This should test that grouping works correctly + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "AND": [ + { + "OR": [ + {"metadata": {"role": "admin"}}, + {"metadata": {"role": "user"}}, + ] + }, + { + "OR": [ + {"metadata": {"dept": "eng"}}, + {"metadata": {"level": {"gte": 3}}}, + ] + }, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + found_names = [item["id"] for item in data["items"]] + + # Should match: + # - admin+eng+1: matches (admin OR user) AND (eng OR level>=3) ✓ + # - user+eng+2: matches (admin OR user) AND (eng OR level>=3) ✓ + # - admin+sales+3: matches (admin OR user) AND (eng OR level>=3) ✓ + # - user+sales+4: matches (admin OR user) AND (eng OR level>=3) ✓ + # So all should match + assert len(found_names) == 4 + + +@pytest.mark.asyncio +async def test_jsonb_contains_vs_equality_semantics( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test the difference between JSONB contains and equality""" + test_workspace, test_peer = sample_data + + session_id = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Create messages with different JSONB structures + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": "Exact match", + "peer_id": test_peer.name, + "metadata": {"role": "admin"}, # Simple key-value + }, + { + "content": "Superset", + "peer_id": test_peer.name, + "metadata": { + "role": "admin", + "department": "engineering", + }, # Contains role=admin plus more + }, + { + "content": "Different", + "peer_id": test_peer.name, + "metadata": {"role": "user"}, # Different value + }, + ] + }, + ) + + # Test JSONB contains behavior - should match both exact and superset + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"role": "admin"}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Exact match" in contents + assert "Superset" in contents # JSONB contains should match this + assert "Different" not in contents + + +@pytest.mark.asyncio +async def test_wildcard_edge_cases_comprehensive( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test comprehensive edge cases for wildcard behavior""" + test_workspace, _test_peer = sample_data + + peer_name = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/peers", + json={"name": peer_name, "metadata": {"role": "admin", "level": 5}}, + ) + + # Test wildcard with comparison operators + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": {"metadata": {"level": {"gte": "*"}}} # Wildcard in comparison + }, + ) + assert response.status_code == 200 + data = response.json() + # Wildcard in comparison should be ignored, returning all peers + assert len(data["items"]) >= 2 + + # Test wildcard in array (in operator) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "metadata": {"role": {"in": ["*", "admin"]}} + } # Mixed wildcard and value + }, + ) + assert response.status_code == 200 + data = response.json() + # Should return all peers since "*" in array means match all + assert len(data["items"]) >= 2 + + # Test multiple wildcards + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={ + "filter": { + "AND": [ + {"id": "*"}, + {"metadata": {"role": "*"}}, + {"metadata": {"level": {"gte": 3}}}, # Only this should filter + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + # Should only apply the level filter + found_names = [item["id"] for item in data["items"]] + assert peer_name in found_names + + +@pytest.mark.asyncio +async def test_error_logging_and_debugging( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test that errors are logged appropriately for debugging""" + test_workspace, _test_peer = sample_data + + # Test scenarios that should generate an error + problematic_filters = [ + {"nonexistent_column": "value"}, + {"created_at": {"gte": "invalid-date"}}, + ] + + for filter_dict in problematic_filters: + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/list", + json={"filter": filter_dict}, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_boundary_conditions_numeric( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test boundary conditions for numeric comparisons""" + test_workspace, test_peer = sample_data + + session_id = str(generate_nanoid()) + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Create messages with boundary values + boundary_values = [ + {"content": "Zero", "metadata": {"score": 0}}, + {"content": "Negative", "metadata": {"score": -1}}, + {"content": "Float", "metadata": {"score": 3.14159}}, + {"content": "Large int", "metadata": {"score": 2147483647}}, # Max 32-bit int + { + "content": "Very large", + "metadata": {"score": 9223372036854775807}, + }, # Max 64-bit int + {"content": "Small float", "metadata": {"score": 0.0000001}}, + ] + + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + { + "content": msg["content"], + "peer_id": test_peer.name, + "metadata": msg["metadata"], + } + for msg in boundary_values + ] + }, + ) + + # Test boundary conditions + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"score": {"gte": 0}}}}, + ) + assert response.status_code == 200 + data = response.json() + # Should include zero and all positive values + contents = [item["content"] for item in data["items"]] + assert "Zero" in contents + assert "Negative" not in contents + + # Test floating point precision + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages/list", + json={"filter": {"metadata": {"score": {"gt": 3.14}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Float" in contents # 3.14159 > 3.14 + + +@pytest.mark.asyncio +async def test_float_precision_edge_cases(client: TestClient): + """Test floating point precision edge cases and rounding behavior""" + # Create workspace and peer through API to ensure they're properly committed + workspace_name = str(generate_nanoid()) + peer_name = str(generate_nanoid()) + + # Create workspace + response = client.post("/v2/workspaces", json={"name": workspace_name}) + assert response.status_code == 200 + + # Create peer + response = client.post( + f"/v2/workspaces/{workspace_name}/peers", json={"name": peer_name} + ) + assert response.status_code == 200 + + # Create messages with problematic floating point values using peer endpoint + messages_response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages", + json={ + "messages": [ + { + "content": "Point one plus point two", + "peer_id": peer_name, + "metadata": { + "value": 0.1 + 0.2, # = 0.30000000000000004 + "precise": 0.3, + "calculation": "0.1 + 0.2", + }, + }, + { + "content": "Exact point three", + "peer_id": peer_name, + "metadata": { + "value": 0.3, + "precise": 0.3, + "calculation": "exact", + }, + }, + { + "content": "Very small difference", + "peer_id": peer_name, + "metadata": { + "value": 0.30000000000000001, + "precise": 0.3, + "calculation": "tiny_diff", + }, + }, + { + "content": "Large precise float", + "peer_id": peer_name, + "metadata": { + "value": 999999.999999999, + "precise": 1000000.0, + "calculation": "large", + }, + }, + { + "content": "Scientific notation", + "peer_id": peer_name, + "metadata": { + "value": 1.23e-10, + "precise": 0.000000000123, + "calculation": "scientific", + }, + }, + { + "content": "Repeating decimal", + "peer_id": peer_name, + "metadata": { + "value": 1.0 / 3.0, # 0.3333... + "precise": 0.3333333333333333, + "calculation": "one_third", + }, + }, + ] + }, + ) + assert messages_response.status_code == 200 + + # Test exact equality - this may or may not work due to floating point precision + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"value": 0.3}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # This test reveals if the system handles floating point precision correctly + # "Exact point three" should definitely match + assert "Exact point three" in contents + + # Test near-equality using range queries (proper way to handle float precision) + epsilon = 1e-10 + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={ + "filter": { + "AND": [ + {"metadata": {"value": {"gte": 0.3 - epsilon}}}, + {"metadata": {"value": {"lte": 0.3 + epsilon}}}, + ] + } + }, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # Should match both exact 0.3 and 0.1+0.2 and very small difference + assert "Exact point three" in contents + assert len(contents) >= 2 + + # Test greater than with floating point precision + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"value": {"gt": 0.3}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # 0.1 + 0.2 is actually > 0.3 due to floating point precision + # This test reveals the actual behavior + + # Test very small numbers and scientific notation + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"value": {"lt": 1e-9}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Scientific notation" in contents + + # Test large number precision + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"value": {"gte": 999999.0}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Large precise float" in contents + + # Test repeating decimal precision + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"value": {"gte": 0.333}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Repeating decimal" in contents + + # Test floating point comparison with string representation + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"calculation": "0.1 + 0.2"}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Point one plus point two" in contents + + +@pytest.mark.asyncio +async def test_mixed_type_comparisons(client: TestClient): + """Test comparisons between different data types within JSONB fields""" + # Create workspace and peer through API to ensure they're properly committed + workspace_name = str(generate_nanoid()) + peer_name = str(generate_nanoid()) + + # Create workspace + response = client.post("/v2/workspaces", json={"name": workspace_name}) + assert response.status_code == 200 + + # Create peer + response = client.post( + f"/v2/workspaces/{workspace_name}/peers", json={"name": peer_name} + ) + assert response.status_code == 200 + + # Create messages with mixed data types for the same logical field + messages_response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages", + json={ + "messages": [ + { + "content": "String number five", + "peer_id": peer_name, + "metadata": { + "priority": "5", # String + "score": "10.5", # String float + "active": "true", # String boolean + "count": "0", # String zero + }, + }, + { + "content": "Integer five", + "peer_id": peer_name, + "metadata": { + "priority": 5, # Integer + "score": 10.5, # Float + "active": True, # Boolean + "count": 0, # Integer zero + }, + }, + { + "content": "Float five", + "peer_id": peer_name, + "metadata": { + "priority": 5.0, # Float that equals integer + "score": 10, # Integer that could be float + "active": False, # Boolean false + "count": None, # Null value + }, + }, + { + "content": "String vs numeric comparison", + "peer_id": peer_name, + "metadata": { + "priority": "10", # String > numeric 5? + "score": "2", # String < numeric 10? + "active": "false", # String false vs boolean + "count": "null", # String null vs null + }, + }, + { + "content": "Leading zeros and formats", + "peer_id": peer_name, + "metadata": { + "priority": "05", # Leading zero + "score": "010.50", # Leading zeros in float + "active": "TRUE", # Uppercase boolean + "count": "00", # Leading zeros + }, + }, + { + "content": "Edge case values", + "peer_id": peer_name, + "metadata": { + "priority": "", # Empty string + "score": "NaN", # Not a number string + "active": 1, # Numeric truthy + "count": "infinity", # Infinity string + }, + }, + ] + }, + ) + assert messages_response.status_code == 200 + + # Test string vs numeric equality + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"priority": 5}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # Should match integer 5 and float 5.0, behavior with string "5" depends on JSONB casting + assert "Integer five" in contents + assert "Float five" in contents + + # Test string number comparison with numeric operator + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"priority": {"gte": 5}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # Behavior depends on how JSONB handles string-to-number conversion + # Should definitely include numeric values >= 5 + + # Test string boolean vs actual boolean + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"active": True}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Integer five" in contents # Has boolean True + # String "true" behavior depends on JSONB casting + + # Test explicit string matching + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"priority": "5"}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "String number five" in contents + + # Test numeric comparison with string numbers + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"score": {"gt": 10}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # Should include 10.5 (both string and float versions) + + # Test zero comparisons (string "0" vs integer 0) + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"count": 0}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Integer five" in contents # Has integer 0 + + # Test null vs string "null" + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"count": None}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Float five" in contents # Has actual null + + # Test leading zeros handling + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"priority": "05"}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Leading zeros and formats" in contents + + # Test case sensitivity for string booleans + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"active": "TRUE"}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Leading zeros and formats" in contents + + # Test empty string vs other falsy values + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"priority": ""}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Edge case values" in contents + + # Test special string values + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"score": "NaN"}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + assert "Edge case values" in contents + + # Test mixed type in operator + response = client.post( + f"/v2/workspaces/{workspace_name}/peers/{peer_name}/messages/list", + json={"filter": {"metadata": {"priority": {"in": [5, "5", 5.0]}}}}, + ) + assert response.status_code == 200 + data = response.json() + contents = [item["content"] for item in data["items"]] + # Should match various representations of 5 + assert len(contents) >= 2