Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 7 additions & 5 deletions src/backend/fastapi_app/api_models.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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')")


Expand Down
30 changes: 20 additions & 10 deletions src/backend/fastapi_app/postgres_searcher.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Optional, Union
from typing import Any, Optional, Union

import numpy as np
from openai import AsyncAzureOpenAI, AsyncOpenAI
Expand All @@ -11,6 +11,11 @@


class PostgresSearcher:
FILTER_OPERATORS = {
"brand": {"=", "!="},
"price": {"=", "<", "<=", ">", ">="},
}

def __init__(
self,
db_session: AsyncSession,
Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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()

Expand Down
76 changes: 67 additions & 9 deletions tests/test_postgres_searcher.py
Original file line number Diff line number Diff line change
@@ -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(
Expand Down
Loading