Skip to content
Open
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
2 changes: 1 addition & 1 deletion backend/chainlit/data/sql_alchemy.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,7 +332,7 @@ async def list_threads(
)

search_keyword = filters.search.lower() if filters.search else None
feedback_value = int(filters.feedback) if filters.feedback else None
feedback_value = int(filters.feedback) if filters.feedback is not None else None

filtered_threads = []
for thread in all_user_threads:
Expand Down
113 changes: 112 additions & 1 deletion backend/tests/data/test_sql_alchemy.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import json
import uuid
from pathlib import Path
from typing import Literal

import pytest
from sqlalchemy import text
Expand All @@ -10,6 +11,7 @@
from chainlit.data.sql_alchemy import SQLAlchemyDataLayer
from chainlit.data.storage_clients.base import BaseStorageClient
from chainlit.element import Text
from chainlit.types import Feedback, Pagination, ThreadFilter


@pytest.fixture
Expand Down Expand Up @@ -98,7 +100,8 @@ async def data_layer(mock_storage_client: BaseStorageClient, tmp_path: Path):
"page" INT,
"language" TEXT,
"forId" UUID,
"mime" TEXT
"mime" TEXT,
"props" JSONB DEFAULT '{}'
);
"""
)
Expand Down Expand Up @@ -219,6 +222,114 @@ async def test_delete_thread(test_user: User, data_layer: SQLAlchemyDataLayer):
assert thread is None


@pytest.fixture
async def feedback_threads_user_id(
mock_chainlit_context, test_user: User, data_layer: SQLAlchemyDataLayer
) -> str:
user = await data_layer.create_user(test_user)
other_user = await data_layer.create_user(User(identifier="other_user"))
assert user is not None
assert other_user is not None

threads: list[tuple[str, int, Literal[0, 1] | None, str]] = [
("negative_new", 6, 0, "Matching answer"),
("positive", 5, 1, "Matching answer"),
("unrated", 4, None, "Matching answer"),
("negative_old", 3, 0, "Matching answer"),
("negative_other", 2, 0, "Different answer"),
("other_user", 7, 0, "Matching answer"),
]
async with mock_chainlit_context:
for thread_id, day, feedback, output in threads:
await data_layer.update_thread(
thread_id,
user_id=other_user.id if thread_id == "other_user" else user.id,
)
step_id = f"{thread_id}_step"
await data_layer.execute_sql(
"""
INSERT INTO steps (
"id", "threadId", "name", "type", "disableFeedback",
"streaming", "output", "createdAt"
) VALUES (
:id, :thread_id, :name, :type, :disable_feedback,
:streaming, :output, :created_at
)
""",
{
"id": step_id,
"thread_id": thread_id,
"name": "Assistant",
"type": "assistant_message",
"disable_feedback": False,
"streaming": False,
"output": output,
"created_at": f"2026-01-{day:02d}T00:00:00Z",
},
)
assert await data_layer.get_step(step_id) is not None
if feedback is not None:
await data_layer.upsert_feedback(
Feedback(forId=step_id, threadId=thread_id, value=feedback)
)
return user.id


@pytest.mark.parametrize(
("feedback", "expected_ids"),
[
(
None,
["negative_new", "positive", "unrated", "negative_old", "negative_other"],
),
(0, ["negative_new", "negative_old", "negative_other"]),
(1, ["positive"]),
],
ids=["all-feedback", "thumbs-down", "thumbs-up"],
)
@pytest.mark.parametrize("search", [None, "MATCH"])
async def test_list_threads_feedback_filter(
data_layer: SQLAlchemyDataLayer,
feedback_threads_user_id: str,
feedback: Literal[0, 1] | None,
expected_ids: list[str],
search: str | None,
):
result = await data_layer.list_threads(
Pagination(first=10),
ThreadFilter(userId=feedback_threads_user_id, feedback=feedback, search=search),
)
if search:
expected_ids = [tid for tid in expected_ids if tid != "negative_other"]
assert [thread["id"] for thread in result.data] == expected_ids
assert result.pageInfo.hasNextPage is False
assert result.pageInfo.startCursor == expected_ids[0]
assert result.pageInfo.endCursor == expected_ids[-1]


async def test_list_threads_negative_feedback_pagination(
data_layer: SQLAlchemyDataLayer, feedback_threads_user_id: str
):
filters = ThreadFilter(userId=feedback_threads_user_id, feedback=0)
cursor = None
expected_ids = ["negative_new", "negative_old", "negative_other"]
for index, thread_id in enumerate(expected_ids):
result = await data_layer.list_threads(
Pagination(first=1, cursor=cursor), filters
)
assert [thread["id"] for thread in result.data] == [thread_id]
assert result.pageInfo.startCursor == thread_id
assert result.pageInfo.endCursor == thread_id
assert result.pageInfo.hasNextPage is (index < len(expected_ids) - 1)
cursor = result.pageInfo.endCursor

result = await data_layer.list_threads(Pagination(first=1, cursor=cursor), filters)
assert result.data == []
assert result.pageInfo.startCursor is None
assert result.pageInfo.endCursor is None
assert result.pageInfo.hasNextPage is False


async def _get_thread_metadata_raw(
data_layer: SQLAlchemyDataLayer, thread_id: str
) -> str | None:
Expand Down