refactor: migrate session.query to select API in retrieval_service (#34638)

This commit is contained in:
Renzo 2026-04-06 23:46:30 -05:00 committed by GitHub
parent 1194957fde
commit 72adb5468c
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 41 additions and 47 deletions

View File

@ -240,7 +240,7 @@ class RetrievalService:
@classmethod @classmethod
def _get_dataset(cls, dataset_id: str) -> Dataset | None: def _get_dataset(cls, dataset_id: str) -> Dataset | None:
with Session(db.engine) as session: with Session(db.engine) as session:
return session.query(Dataset).where(Dataset.id == dataset_id).first() return session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1))
@classmethod @classmethod
def keyword_search( def keyword_search(
@ -573,15 +573,13 @@ class RetrievalService:
# Batch query summaries for segments retrieved via summary (only enabled summaries) # Batch query summaries for segments retrieved via summary (only enabled summaries)
if summary_segment_ids: if summary_segment_ids:
summaries = ( summaries = session.scalars(
session.query(DocumentSegmentSummary) select(DocumentSegmentSummary).where(
.filter(
DocumentSegmentSummary.chunk_id.in_(list(summary_segment_ids)), DocumentSegmentSummary.chunk_id.in_(list(summary_segment_ids)),
DocumentSegmentSummary.status == "completed", DocumentSegmentSummary.status == "completed",
DocumentSegmentSummary.enabled == True, # Only retrieve enabled summaries DocumentSegmentSummary.enabled.is_(True), # Only retrieve enabled summaries
) )
.all() ).all()
)
for summary in summaries: for summary in summaries:
if summary.summary_content: if summary.summary_content:
segment_summary_map[summary.chunk_id] = summary.summary_content segment_summary_map[summary.chunk_id] = summary.summary_content
@ -851,12 +849,12 @@ class RetrievalService:
def get_segment_attachment_info( def get_segment_attachment_info(
cls, dataset_id: str, tenant_id: str, attachment_id: str, session: Session cls, dataset_id: str, tenant_id: str, attachment_id: str, session: Session
) -> SegmentAttachmentResult | None: ) -> SegmentAttachmentResult | None:
upload_file = session.query(UploadFile).where(UploadFile.id == attachment_id).first() upload_file = session.scalar(select(UploadFile).where(UploadFile.id == attachment_id).limit(1))
if upload_file: if upload_file:
attachment_binding = ( attachment_binding = session.scalar(
session.query(SegmentAttachmentBinding) select(SegmentAttachmentBinding)
.where(SegmentAttachmentBinding.attachment_id == upload_file.id) .where(SegmentAttachmentBinding.attachment_id == upload_file.id)
.first() .limit(1)
) )
if attachment_binding: if attachment_binding:
attachment_info: AttachmentInfoDict = { attachment_info: AttachmentInfoDict = {
@ -875,14 +873,12 @@ class RetrievalService:
cls, attachment_ids: list[str], session: Session cls, attachment_ids: list[str], session: Session
) -> list[SegmentAttachmentInfoResult]: ) -> list[SegmentAttachmentInfoResult]:
attachment_infos: list[SegmentAttachmentInfoResult] = [] attachment_infos: list[SegmentAttachmentInfoResult] = []
upload_files = session.query(UploadFile).where(UploadFile.id.in_(attachment_ids)).all() upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all()
if upload_files: if upload_files:
upload_file_ids = [upload_file.id for upload_file in upload_files] upload_file_ids = [upload_file.id for upload_file in upload_files]
attachment_bindings = ( attachment_bindings = session.scalars(
session.query(SegmentAttachmentBinding) select(SegmentAttachmentBinding).where(SegmentAttachmentBinding.attachment_id.in_(upload_file_ids))
.where(SegmentAttachmentBinding.attachment_id.in_(upload_file_ids)) ).all()
.all()
)
attachment_binding_map = {binding.attachment_id: binding for binding in attachment_bindings} attachment_binding_map = {binding.attachment_id: binding for binding in attachment_bindings}
if attachment_bindings: if attachment_bindings:

View File

@ -119,6 +119,14 @@ class _FakeSummaryQuery:
return self._summaries return self._summaries
class _FakeScalarsResult:
def __init__(self, data: list) -> None:
self._data = data
def all(self) -> list:
return self._data
class _FakeSession: class _FakeSession:
def __init__(self, execute_payloads: list[list], summaries: list) -> None: def __init__(self, execute_payloads: list[list], summaries: list) -> None:
self._payloads = list(execute_payloads) self._payloads = list(execute_payloads)
@ -128,8 +136,8 @@ class _FakeSession:
data = self._payloads.pop(0) if self._payloads else [] data = self._payloads.pop(0) if self._payloads else []
return _FakeExecuteResult(data) return _FakeExecuteResult(data)
def query(self, model): def scalars(self, stmt):
return _FakeSummaryQuery(self._summaries) return _FakeScalarsResult(self._summaries)
class _FakeSessionContext: class _FakeSessionContext:
@ -265,14 +273,14 @@ class TestRetrievalServiceInternals:
def test_get_dataset_queries_by_id(self, mock_session_class): def test_get_dataset_queries_by_id(self, mock_session_class):
expected_dataset = Mock(spec=Dataset) expected_dataset = Mock(spec=Dataset)
mock_session = Mock() mock_session = Mock()
mock_session.query.return_value.where.return_value.first.return_value = expected_dataset mock_session.scalar.return_value = expected_dataset
mock_session_class.return_value.__enter__.return_value = mock_session mock_session_class.return_value.__enter__.return_value = mock_session
with patch.object(retrieval_service_module, "db", SimpleNamespace(engine=Mock())): with patch.object(retrieval_service_module, "db", SimpleNamespace(engine=Mock())):
result = RetrievalService._get_dataset("dataset-123") result = RetrievalService._get_dataset("dataset-123")
assert result == expected_dataset assert result == expected_dataset
mock_session.query.assert_called_once() mock_session.scalar.assert_called_once()
@patch("core.rag.datasource.retrieval_service.Keyword") @patch("core.rag.datasource.retrieval_service.Keyword")
@patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset") @patch("core.rag.datasource.retrieval_service.RetrievalService._get_dataset")
@ -1046,12 +1054,8 @@ class TestRetrievalServiceInternals:
size=42, size=42,
) )
binding = SimpleNamespace(segment_id="segment-1", attachment_id="upload-1") binding = SimpleNamespace(segment_id="segment-1", attachment_id="upload-1")
upload_query = Mock()
upload_query.where.return_value.first.return_value = upload_file
binding_query = Mock()
binding_query.where.return_value.first.return_value = binding
session = Mock() session = Mock()
session.query.side_effect = [upload_query, binding_query] session.scalar.side_effect = [upload_file, binding]
result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session) result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session)
@ -1076,32 +1080,26 @@ class TestRetrievalServiceInternals:
mime_type="image/png", mime_type="image/png",
size=42, size=42,
) )
upload_query = Mock()
upload_query.where.return_value.first.return_value = upload_file
binding_query = Mock()
binding_query.where.return_value.first.return_value = None
session = Mock() session = Mock()
session.query.side_effect = [upload_query, binding_query] session.scalar.side_effect = [upload_file, None]
result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session) result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session)
assert result is None assert result is None
def test_get_segment_attachment_info_returns_none_when_upload_file_missing(self): def test_get_segment_attachment_info_returns_none_when_upload_file_missing(self):
upload_query = Mock()
upload_query.where.return_value.first.return_value = None
session = Mock() session = Mock()
session.query.return_value = upload_query session.scalar.return_value = None
result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session) result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session)
assert result is None assert result is None
def test_get_segment_attachment_infos_returns_empty_when_upload_files_missing(self): def test_get_segment_attachment_infos_returns_empty_when_upload_files_missing(self):
upload_query = Mock() scalars_result = Mock()
upload_query.where.return_value.all.return_value = [] scalars_result.all.return_value = []
session = Mock() session = Mock()
session.query.return_value = upload_query session.scalars.return_value = scalars_result
result = RetrievalService.get_segment_attachment_infos(["upload-1"], session) result = RetrievalService.get_segment_attachment_infos(["upload-1"], session)
@ -1115,12 +1113,12 @@ class TestRetrievalServiceInternals:
mime_type="image/png", mime_type="image/png",
size=42, size=42,
) )
upload_query = Mock() upload_scalars = Mock()
upload_query.where.return_value.all.return_value = [upload_file] upload_scalars.all.return_value = [upload_file]
binding_query = Mock() binding_scalars = Mock()
binding_query.where.return_value.all.return_value = [] binding_scalars.all.return_value = []
session = Mock() session = Mock()
session.query.side_effect = [upload_query, binding_query] session.scalars.side_effect = [upload_scalars, binding_scalars]
result = RetrievalService.get_segment_attachment_infos(["upload-1"], session) result = RetrievalService.get_segment_attachment_infos(["upload-1"], session)
@ -1144,12 +1142,12 @@ class TestRetrievalServiceInternals:
) )
binding = SimpleNamespace(attachment_id="upload-1", segment_id="segment-1") binding = SimpleNamespace(attachment_id="upload-1", segment_id="segment-1")
upload_query = Mock() upload_scalars = Mock()
upload_query.where.return_value.all.return_value = [upload_file_1, upload_file_2] upload_scalars.all.return_value = [upload_file_1, upload_file_2]
binding_query = Mock() binding_scalars = Mock()
binding_query.where.return_value.all.return_value = [binding] binding_scalars.all.return_value = [binding]
session = Mock() session = Mock()
session.query.side_effect = [upload_query, binding_query] session.scalars.side_effect = [upload_scalars, binding_scalars]
result = RetrievalService.get_segment_attachment_infos(["upload-1", "upload-2"], session) result = RetrievalService.get_segment_attachment_infos(["upload-1", "upload-2"], session)