mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: migrate session.query to select API in retrieval_service (#34638)
This commit is contained in:
@@ -119,6 +119,14 @@ class _FakeSummaryQuery:
|
||||
return self._summaries
|
||||
|
||||
|
||||
class _FakeScalarsResult:
|
||||
def __init__(self, data: list) -> None:
|
||||
self._data = data
|
||||
|
||||
def all(self) -> list:
|
||||
return self._data
|
||||
|
||||
|
||||
class _FakeSession:
|
||||
def __init__(self, execute_payloads: list[list], summaries: list) -> None:
|
||||
self._payloads = list(execute_payloads)
|
||||
@@ -128,8 +136,8 @@ class _FakeSession:
|
||||
data = self._payloads.pop(0) if self._payloads else []
|
||||
return _FakeExecuteResult(data)
|
||||
|
||||
def query(self, model):
|
||||
return _FakeSummaryQuery(self._summaries)
|
||||
def scalars(self, stmt):
|
||||
return _FakeScalarsResult(self._summaries)
|
||||
|
||||
|
||||
class _FakeSessionContext:
|
||||
@@ -265,14 +273,14 @@ class TestRetrievalServiceInternals:
|
||||
def test_get_dataset_queries_by_id(self, mock_session_class):
|
||||
expected_dataset = Mock(spec=Dataset)
|
||||
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
|
||||
|
||||
with patch.object(retrieval_service_module, "db", SimpleNamespace(engine=Mock())):
|
||||
result = RetrievalService._get_dataset("dataset-123")
|
||||
|
||||
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.RetrievalService._get_dataset")
|
||||
@@ -1046,12 +1054,8 @@ class TestRetrievalServiceInternals:
|
||||
size=42,
|
||||
)
|
||||
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.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)
|
||||
|
||||
@@ -1076,32 +1080,26 @@ class TestRetrievalServiceInternals:
|
||||
mime_type="image/png",
|
||||
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.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)
|
||||
|
||||
assert result is None
|
||||
|
||||
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.query.return_value = upload_query
|
||||
session.scalar.return_value = None
|
||||
|
||||
result = RetrievalService.get_segment_attachment_info("dataset-id", "tenant-id", "upload-1", session)
|
||||
|
||||
assert result is None
|
||||
|
||||
def test_get_segment_attachment_infos_returns_empty_when_upload_files_missing(self):
|
||||
upload_query = Mock()
|
||||
upload_query.where.return_value.all.return_value = []
|
||||
scalars_result = Mock()
|
||||
scalars_result.all.return_value = []
|
||||
session = Mock()
|
||||
session.query.return_value = upload_query
|
||||
session.scalars.return_value = scalars_result
|
||||
|
||||
result = RetrievalService.get_segment_attachment_infos(["upload-1"], session)
|
||||
|
||||
@@ -1115,12 +1113,12 @@ class TestRetrievalServiceInternals:
|
||||
mime_type="image/png",
|
||||
size=42,
|
||||
)
|
||||
upload_query = Mock()
|
||||
upload_query.where.return_value.all.return_value = [upload_file]
|
||||
binding_query = Mock()
|
||||
binding_query.where.return_value.all.return_value = []
|
||||
upload_scalars = Mock()
|
||||
upload_scalars.all.return_value = [upload_file]
|
||||
binding_scalars = Mock()
|
||||
binding_scalars.all.return_value = []
|
||||
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)
|
||||
|
||||
@@ -1144,12 +1142,12 @@ class TestRetrievalServiceInternals:
|
||||
)
|
||||
binding = SimpleNamespace(attachment_id="upload-1", segment_id="segment-1")
|
||||
|
||||
upload_query = Mock()
|
||||
upload_query.where.return_value.all.return_value = [upload_file_1, upload_file_2]
|
||||
binding_query = Mock()
|
||||
binding_query.where.return_value.all.return_value = [binding]
|
||||
upload_scalars = Mock()
|
||||
upload_scalars.all.return_value = [upload_file_1, upload_file_2]
|
||||
binding_scalars = Mock()
|
||||
binding_scalars.all.return_value = [binding]
|
||||
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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user