Skip to content
Closed
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
16 changes: 7 additions & 9 deletions apps/ml-service/app/recommender/content_based.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,10 +37,10 @@ async def get_content_recommendations(

# Build exclusion clause
exclude_clause = ""
params: list = [preference_vector.tolist(), top_n * 2]
params: list = [preference_vector.tolist()]
if exclude_ids:
exclude_clause = "AND r.id != ALL(%s)"
params.insert(1, exclude_ids)
params.append(exclude_ids)

# Query recipes by cosine similarity to preference vector
query = f"""
Expand All @@ -55,9 +55,8 @@ async def get_content_recommendations(
LIMIT %s
"""

# Re-add vector for the ORDER BY clause
params.append(preference_vector.tolist())
params.append(top_n * 2)
# Re-add vector for the ORDER BY clause.
params.extend([preference_vector.tolist(), top_n * 2])

result = await conn.execute(query, tuple(params))
rows = await result.fetchall()
Expand Down Expand Up @@ -102,10 +101,10 @@ async def _get_nutrition_similarity(
) -> list[dict]:
"""Get recipes similar by nutrition vector."""
exclude_clause = ""
params: list = [preference_vector.tolist(), top_n]
params: list = [preference_vector.tolist()]
if exclude_ids:
exclude_clause = "AND r.id != ALL(%s)"
params.insert(1, exclude_ids)
params.append(exclude_ids)

query = f"""
SELECT
Expand All @@ -118,8 +117,7 @@ async def _get_nutrition_similarity(
ORDER BY r.nutrition_vector <=> %s::vector ASC
LIMIT %s
"""
params.append(preference_vector.tolist())
params.append(top_n)
params.extend([preference_vector.tolist(), top_n])

result = await conn.execute(query, tuple(params))
return await result.fetchall()
100 changes: 100 additions & 0 deletions apps/ml-service/tests/test_recommend.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,41 @@

from app.main import app
from app.models.schemas import RecommendResponse, RecipeScore
from app.recommender.content_based import get_content_recommendations


class _FakeResult:
def __init__(self, *, one=None, many=None):
self._one = one
self._many = many or []

async def fetchone(self):
return self._one

async def fetchall(self):
return self._many


class _ContentRecommendationConn:
def __init__(self, expected_calls):
self.expected_calls = expected_calls
self.calls = []

async def execute(self, query, params):
self.calls.append((query, params))

if "FROM user_taste_profiles" in query:
return _FakeResult(
one={
"preference_vector": [0.1, 0.2, 0.3],
"interaction_count": 8,
}
)

expected_params, rows = self.expected_calls.pop(0)
assert params == expected_params
assert query.count("%s") == len(params)
return _FakeResult(many=rows)


@asynccontextmanager
Expand Down Expand Up @@ -118,3 +153,68 @@ async def test_recommend_validates_top_n(client: AsyncClient):
async def test_recommend_validates_missing_user_id(client: AsyncClient):
response = await client.post("/recommend", json={"top_n": 10})
assert response.status_code == 422


@pytest.mark.asyncio
async def test_content_recommendations_binds_vector_params_in_sql_order():
vector = [0.1, 0.2, 0.3]
conn = _ContentRecommendationConn(
expected_calls=[
(
(vector, vector, 6),
[
{
"recipe_id": "ingredient-match",
"title": "Ingredient Match",
"similarity_score": 0.9,
}
],
),
(
(vector, vector, 3),
[
{
"recipe_id": "nutrition-match",
"title": "Nutrition Match",
"similarity_score": 0.8,
}
],
),
]
)

recommendations = await get_content_recommendations(conn, "user-001", top_n=3)

assert [rec["recipe_id"] for rec in recommendations] == [
"ingredient-match",
"nutrition-match",
]
assert conn.expected_calls == []


@pytest.mark.asyncio
async def test_content_recommendations_binds_exclusions_before_order_vector():
vector = [0.1, 0.2, 0.3]
exclude_ids = ["recipe-to-skip"]
conn = _ContentRecommendationConn(
expected_calls=[
(
(vector, exclude_ids, vector, 4),
[],
),
(
(vector, exclude_ids, vector, 2),
[],
),
]
)

recommendations = await get_content_recommendations(
conn,
"user-001",
top_n=2,
exclude_ids=exclude_ids,
)

assert recommendations == []
assert conn.expected_calls == []
Loading