diff --git a/src/backend/fastapi_app/api_models.py b/src/backend/fastapi_app/api_models.py index fefcbdf..68714b6 100644 --- a/src/backend/fastapi_app/api_models.py +++ b/src/backend/fastapi_app/api_models.py @@ -1,5 +1,5 @@ from enum import Enum -from typing import Any, Optional +from typing import Any, Literal, Optional from openai.types.responses import ResponseInputItemParam from pydantic import BaseModel, Field @@ -90,14 +90,16 @@ class Filter(BaseModel): class PriceFilter(Filter): - column: str = Field(default="price", description="The column to filter on (always 'price' for this filter)") - comparison_operator: str = Field(description="The operator for price comparison ('>', '<', '>=', '<=', '=')") + column: Literal["price"] = Field(default="price", description="The column to filter on (always 'price')") + comparison_operator: Literal[">", "<", ">=", "<=", "="] = Field( + description="The operator for price comparison ('>', '<', '>=', '<=', '=')" + ) value: float = Field(description="The price value to compare against (e.g., 30.00)") class BrandFilter(Filter): - column: str = Field(default="brand", description="The column to filter on (always 'brand' for this filter)") - comparison_operator: str = Field(description="The operator for brand comparison ('=' or '!=')") + column: Literal["brand"] = Field(default="brand", description="The column to filter on (always 'brand')") + comparison_operator: Literal["=", "!="] = Field(description="The operator for brand comparison ('=' or '!=')") value: str = Field(description="The brand name to compare against (e.g., 'AirStrider')") diff --git a/src/backend/fastapi_app/postgres_searcher.py b/src/backend/fastapi_app/postgres_searcher.py index aa84eaf..12c45cd 100644 --- a/src/backend/fastapi_app/postgres_searcher.py +++ b/src/backend/fastapi_app/postgres_searcher.py @@ -1,4 +1,4 @@ -from typing import Optional, Union +from typing import Any, Optional, Union import numpy as np from openai import AsyncAzureOpenAI, AsyncOpenAI @@ -11,6 +11,11 @@ class PostgresSearcher: + FILTER_OPERATORS = { + "brand": {"=", "!="}, + "price": {"=", "<", "<=", ">", ">="}, + } + def __init__( self, db_session: AsyncSession, @@ -27,17 +32,22 @@ def __init__( self.embed_dimensions = embed_dimensions self.embedding_column = embedding_column - def build_filter_clause(self, filters: Optional[list[Filter]]) -> tuple[str, str]: + def build_filter_clause(self, filters: Optional[list[Filter]]) -> tuple[str, str, dict[str, Any]]: if filters is None: - return "", "" + return "", "", {} filter_clauses = [] - for filter in filters: - filter_value = f"'{filter.value}'" if isinstance(filter.value, str) else filter.value - filter_clauses.append(f"{filter.column} {filter.comparison_operator} {filter_value}") + filter_params = {} + for index, search_filter in enumerate(filters): + allowed_operators = self.FILTER_OPERATORS.get(search_filter.column) + if allowed_operators is None or search_filter.comparison_operator not in allowed_operators: + raise ValueError(f"Unsupported filter: {search_filter.column} {search_filter.comparison_operator}") + parameter_name = f"filter_{index}" + filter_clauses.append(f"{search_filter.column} {search_filter.comparison_operator} :{parameter_name}") + filter_params[parameter_name] = search_filter.value filter_clause = " AND ".join(filter_clauses) if len(filter_clause) > 0: - return f"WHERE {filter_clause}", f"AND {filter_clause}" - return "", "" + return f"WHERE {filter_clause}", f"AND {filter_clause}", filter_params + return "", "", {} async def search( self, @@ -46,7 +56,7 @@ async def search( top: int = 5, filters: Optional[list[Filter]] = None, ): - filter_clause_where, filter_clause_and = self.build_filter_clause(filters) + filter_clause_where, filter_clause_and, filter_params = self.build_filter_clause(filters) table_name = Item.__tablename__ vector_query = f""" SELECT id, RANK () OVER (ORDER BY {self.embedding_column} <=> :embedding) AS rank @@ -93,7 +103,7 @@ async def search( results = ( await self.db_session.execute( sql, - {"embedding": np.array(query_vector), "query": query_text, "k": 60}, + {"embedding": np.array(query_vector), "query": query_text, "k": 60, **filter_params}, ) ).fetchall() diff --git a/tests/test_postgres_searcher.py b/tests/test_postgres_searcher.py index fff7fdf..617a340 100644 --- a/tests/test_postgres_searcher.py +++ b/tests/test_postgres_searcher.py @@ -1,41 +1,99 @@ +from unittest.mock import AsyncMock, MagicMock + import pytest +from pydantic import ValidationError -from fastapi_app.api_models import Filter, ItemPublic +from fastapi_app.api_models import BrandFilter, Filter, ItemPublic, PriceFilter +from fastapi_app.postgres_searcher import PostgresSearcher from tests.data import test_data def test_postgres_build_filter_clause_without_filters(postgres_searcher): - assert postgres_searcher.build_filter_clause(None) == ("", "") - assert postgres_searcher.build_filter_clause([]) == ("", "") + assert postgres_searcher.build_filter_clause(None) == ("", "", {}) + assert postgres_searcher.build_filter_clause([]) == ("", "", {}) def test_postgres_build_filter_clause_with_filters(postgres_searcher): assert postgres_searcher.build_filter_clause( [ - Filter(column="brand", comparison_operator="=", value="AirStrider"), + BrandFilter(comparison_operator="=", value="AirStrider"), ] ) == ( - "WHERE brand = 'AirStrider'", - "AND brand = 'AirStrider'", + "WHERE brand = :filter_0", + "AND brand = :filter_0", + {"filter_0": "AirStrider"}, ) def test_postgres_build_filter_clause_with_filters_numeric(postgres_searcher): assert postgres_searcher.build_filter_clause( [ - Filter(column="price", comparison_operator="<", value=30), + PriceFilter(comparison_operator="<", value=30), ] ) == ( - "WHERE price < 30", - "AND price < 30", + "WHERE price < :filter_0", + "AND price < :filter_0", + {"filter_0": 30}, + ) + + +def test_postgres_build_filter_clause_parameterizes_injection_payload(postgres_searcher): + payload = "x' OR TRUE --" + + assert postgres_searcher.build_filter_clause([BrandFilter(comparison_operator="=", value=payload)]) == ( + "WHERE brand = :filter_0", + "AND brand = :filter_0", + {"filter_0": payload}, ) +def test_postgres_build_filter_clause_rejects_unsupported_filter(postgres_searcher): + with pytest.raises(ValueError, match="Unsupported filter"): + postgres_searcher.build_filter_clause([Filter(column="description", comparison_operator="=", value="tent")]) + + +def test_filter_models_reject_overridden_columns_and_operators(): + with pytest.raises(ValidationError): + BrandFilter.model_validate({"column": "description", "comparison_operator": "=", "value": "tent"}) + with pytest.raises(ValidationError): + PriceFilter.model_validate({"comparison_operator": "OR TRUE --", "value": 30}) + + @pytest.mark.asyncio async def test_postgres_searcher_search_empty_text_search(postgres_searcher): assert await postgres_searcher.search("", [], 5, None) == [] +@pytest.mark.asyncio +async def test_postgres_searcher_search_binds_filter_value(): + db_session = AsyncMock() + result = MagicMock() + result.fetchall.return_value = [] + db_session.execute.return_value = result + searcher = PostgresSearcher( + db_session=db_session, + openai_embed_client=AsyncMock(), + embed_deployment=None, + embed_model="text-embedding-3-small", + embed_dimensions=1536, + embedding_column="embedding", + ) + payload = "x' OR TRUE --" + + assert ( + await searcher.search( + "", + [], + 5, + [BrandFilter(comparison_operator="=", value=payload)], + ) + == [] + ) + sql, parameters = db_session.execute.await_args.args + assert payload not in str(sql) + assert parameters["filter_0"] == payload + + @pytest.mark.asyncio async def test_postgres_searcher_search(postgres_searcher): assert (await postgres_searcher.search(test_data.name, test_data.embeddings, 5, None))[0].to_dict() == ItemPublic(