refactor: select in dataset_service (SegmentService and remaining cla… (#34547)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Renzo
2026-04-05 00:13:06 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 779e6b8e0b
commit eca0cdc7a9
3 changed files with 92 additions and 129 deletions
@@ -1607,7 +1607,7 @@ class TestDatasetCollectionBindingService:
binding = SimpleNamespace(id="binding-1")
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.order_by.return_value.first.return_value = binding
mock_db.session.scalar.return_value = binding
result = DatasetCollectionBindingService.get_dataset_collection_binding("provider", "model")
@@ -1619,10 +1619,11 @@ class TestDatasetCollectionBindingService:
with (
patch("services.dataset_service.db") as mock_db,
patch("services.dataset_service.select"),
patch("services.dataset_service.DatasetCollectionBinding", return_value=created_binding) as binding_cls,
patch.object(Dataset, "gen_collection_name_by_id", return_value="generated-collection"),
):
mock_db.session.query.return_value.where.return_value.order_by.return_value.first.return_value = None
mock_db.session.scalar.return_value = None
result = DatasetCollectionBindingService.get_dataset_collection_binding("provider", "model", "dataset")
@@ -1638,7 +1639,7 @@ class TestDatasetCollectionBindingService:
def test_get_dataset_collection_binding_by_id_and_type_raises_when_missing(self):
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.order_by.return_value.first.return_value = None
mock_db.session.scalar.return_value = None
with pytest.raises(ValueError, match="Dataset collection binding not found"):
DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type("binding-1")
@@ -1647,7 +1648,7 @@ class TestDatasetCollectionBindingService:
binding = SimpleNamespace(id="binding-1")
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.order_by.return_value.first.return_value = binding
mock_db.session.scalar.return_value = binding
result = DatasetCollectionBindingService.get_dataset_collection_binding_by_id_and_type("binding-1")
@@ -1673,7 +1674,7 @@ class TestDatasetPermissionService:
[{"user_id": "user-1"}, {"user_id": "user-2"}],
)
mock_db.session.query.return_value.where.return_value.delete.assert_called_once()
mock_db.session.execute.assert_called()
mock_db.session.add_all.assert_called_once()
mock_db.session.commit.assert_called_once()
@@ -1744,12 +1745,12 @@ class TestDatasetPermissionService:
with patch("services.dataset_service.db") as mock_db:
DatasetPermissionService.clear_partial_member_list("dataset-1")
mock_db.session.query.return_value.where.return_value.delete.assert_called_once()
mock_db.session.execute.assert_called()
mock_db.session.commit.assert_called_once()
def test_clear_partial_member_list_rolls_back_on_exception(self):
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.delete.side_effect = RuntimeError("boom")
mock_db.session.execute.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError, match="boom"):
DatasetPermissionService.clear_partial_member_list("dataset-1")
@@ -49,7 +49,7 @@ class TestSegmentServiceChildChunks:
patch("services.dataset_service.VectorService") as vector_service,
):
mock_redis.lock.return_value = _make_lock_context()
mock_db.session.query.return_value.where.return_value.scalar.return_value = 2
mock_db.session.scalar.return_value = 2
child_chunk = SegmentService.create_child_chunk("child content", segment, document, dataset)
@@ -75,7 +75,7 @@ class TestSegmentServiceChildChunks:
patch("services.dataset_service.VectorService") as vector_service,
):
mock_redis.lock.return_value = _make_lock_context()
mock_db.session.query.return_value.where.return_value.scalar.return_value = None
mock_db.session.scalar.return_value = None
vector_service.create_child_chunk_vector.side_effect = RuntimeError("vector failed")
with pytest.raises(ChildChunkIndexingError, match="vector failed"):
@@ -247,13 +247,13 @@ class TestSegmentServiceQueries:
child_chunk = _make_child_chunk()
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.first.return_value = child_chunk
mock_db.session.scalar.return_value = child_chunk
result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1")
assert result is child_chunk
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.first.return_value = SimpleNamespace()
mock_db.session.scalar.return_value = SimpleNamespace()
result = SegmentService.get_child_chunk_by_id("child-a", "tenant-1")
assert result is None
@@ -295,13 +295,13 @@ class TestSegmentServiceQueries:
)
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.first.return_value = segment
mock_db.session.scalar.return_value = segment
result = SegmentService.get_segment_by_id("segment-1", "tenant-1")
assert result is segment
with patch("services.dataset_service.db") as mock_db:
mock_db.session.query.return_value.where.return_value.first.return_value = SimpleNamespace()
mock_db.session.scalar.return_value = SimpleNamespace()
result = SegmentService.get_segment_by_id("segment-1", "tenant-1")
assert result is None
@@ -401,11 +401,8 @@ class TestSegmentServiceMutations:
):
mock_redis.lock.return_value = _make_lock_context()
max_position_query = MagicMock()
max_position_query.where.return_value.scalar.return_value = 2
refresh_query = MagicMock()
refresh_query.where.return_value.first.return_value = refreshed_segment
mock_db.session.query.side_effect = [max_position_query, refresh_query]
mock_db.session.scalar.return_value = 2
mock_db.session.get.return_value = refreshed_segment
def add_side_effect(obj):
if obj.__class__.__name__ == "DocumentSegment" and getattr(obj, "id", None) is None:
@@ -461,7 +458,7 @@ class TestSegmentServiceMutations:
):
mock_redis.lock.return_value = _make_lock_context()
model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model
mock_db.session.query.return_value.where.return_value.scalar.return_value = 1
mock_db.session.scalar.return_value = 1
vector_service.create_segments_vector.side_effect = RuntimeError("vector failed")
result = SegmentService.multi_create_segment(segments, document, dataset)
@@ -538,7 +535,7 @@ class TestSegmentServiceMutations:
patch("services.dataset_service.VectorService") as vector_service,
):
mock_redis.get.return_value = None
mock_db.session.query.return_value.where.return_value.first.return_value = refreshed_segment
mock_db.session.get.return_value = refreshed_segment
result = SegmentService.update_segment(args, segment, document, dataset)
@@ -574,13 +571,10 @@ class TestSegmentServiceMutations:
mock_redis.get.return_value = None
model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model_instance
processing_rule_query = MagicMock()
processing_rule_query.where.return_value.first.return_value = processing_rule
summary_query = MagicMock()
summary_query.where.return_value.first.return_value = existing_summary
refreshed_query = MagicMock()
refreshed_query.where.return_value.first.return_value = refreshed_segment
mock_db.session.query.side_effect = [processing_rule_query, summary_query, refreshed_query]
# get calls: processing_rule, then refreshed_segment
mock_db.session.get.side_effect = [processing_rule, refreshed_segment]
# scalar call: existing_summary
mock_db.session.scalar.return_value = existing_summary
result = SegmentService.update_segment(args, segment, document, dataset)
@@ -621,11 +615,8 @@ class TestSegmentServiceMutations:
mock_redis.get.return_value = None
model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model
summary_query = MagicMock()
summary_query.where.return_value.first.return_value = existing_summary
refreshed_query = MagicMock()
refreshed_query.where.return_value.first.return_value = refreshed_segment
mock_db.session.query.side_effect = [summary_query, refreshed_query]
mock_db.session.scalar.return_value = existing_summary
mock_db.session.get.return_value = refreshed_segment
result = SegmentService.update_segment(args, segment, document, dataset)
@@ -664,11 +655,8 @@ class TestSegmentServiceMutations:
mock_redis.get.return_value = None
model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model
summary_query = MagicMock()
summary_query.where.return_value.first.return_value = existing_summary
refreshed_query = MagicMock()
refreshed_query.where.return_value.first.return_value = refreshed_segment
mock_db.session.query.side_effect = [summary_query, refreshed_query]
mock_db.session.scalar.return_value = existing_summary
mock_db.session.get.return_value = refreshed_segment
result = SegmentService.update_segment(args, segment, document, dataset)
@@ -688,7 +676,7 @@ class TestSegmentServiceMutations:
patch("services.dataset_service.delete_segment_from_index_task") as delete_task,
):
mock_redis.get.return_value = None
mock_db.session.query.return_value.where.return_value.all.return_value = [("child-1",), ("child-2",)]
mock_db.session.scalars.return_value.all.return_value = ["child-1", "child-2"]
SegmentService.delete_segment(segment, document, dataset)
@@ -727,15 +715,15 @@ class TestSegmentServiceMutations:
patch("services.dataset_service.delete_segment_from_index_task") as delete_task,
):
segments_query = MagicMock()
segments_query.with_entities.return_value.where.return_value.all.return_value = [
# execute().all() for segments_info (multi-column)
execute_result = MagicMock()
execute_result.all.return_value = [
("node-1", "segment-1", 2),
("node-2", "segment-2", 5),
]
child_query = MagicMock()
child_query.where.return_value.all.return_value = [("child-1",)]
delete_query = MagicMock()
delete_query.where.return_value.delete.return_value = 2
mock_db.session.query.side_effect = [segments_query, child_query, delete_query]
mock_db.session.execute.return_value = execute_result
# scalars() for child_node_ids
mock_db.session.scalars.return_value.all.return_value = ["child-1"]
SegmentService.delete_segments(["segment-1", "segment-2"], document, dataset)
@@ -748,7 +736,6 @@ class TestSegmentServiceMutations:
["segment-1", "segment-2"],
["child-1"],
)
delete_query.where.return_value.delete.assert_called_once()
mock_db.session.commit.assert_called_once()
def test_update_segments_status_enables_only_segments_without_indexing_cache(self):
@@ -868,7 +855,7 @@ class TestSegmentServiceAdditionalRegenerationBranches:
patch("services.dataset_service.VectorService") as vector_service,
):
mock_redis.get.return_value = None
mock_db.session.query.return_value.where.return_value.first.return_value = refreshed_segment
mock_db.session.get.return_value = refreshed_segment
result = SegmentService.update_segment(
SegmentUpdateArgs(content="question", answer="new answer"),
@@ -902,11 +889,8 @@ class TestSegmentServiceAdditionalRegenerationBranches:
):
mock_redis.get.return_value = None
model_manager_cls.for_tenant.return_value.get_model_instance.return_value = embedding_model
summary_query = MagicMock()
summary_query.where.return_value.first.return_value = None
refreshed_query = MagicMock()
refreshed_query.where.return_value.first.return_value = refreshed_segment
mock_db.session.query.side_effect = [summary_query, refreshed_query]
mock_db.session.scalar.return_value = None
mock_db.session.get.return_value = refreshed_segment
result = SegmentService.update_segment(
SegmentUpdateArgs(content="new question", answer="new answer", keywords=["kw-1"]),
@@ -951,13 +935,10 @@ class TestSegmentServiceAdditionalRegenerationBranches:
model_manager_cls.for_tenant.return_value.get_default_model_instance.return_value = embedding_model_instance
update_summary.side_effect = RuntimeError("summary failed")
processing_rule_query = MagicMock()
processing_rule_query.where.return_value.first.return_value = processing_rule
summary_query = MagicMock()
summary_query.where.return_value.first.return_value = existing_summary
refreshed_query = MagicMock()
refreshed_query.where.return_value.first.return_value = refreshed_segment
mock_db.session.query.side_effect = [processing_rule_query, summary_query, refreshed_query]
# get calls: processing_rule, then refreshed_segment
mock_db.session.get.side_effect = [processing_rule, refreshed_segment]
# scalar call: existing_summary
mock_db.session.scalar.return_value = existing_summary
result = SegmentService.update_segment(
SegmentUpdateArgs(content="new parent content", regenerate_child_chunks=True, summary="new summary"),
@@ -1000,7 +981,7 @@ class TestSegmentServiceAdditionalRegenerationBranches:
patch("services.dataset_service.VectorService") as vector_service,
):
mock_redis.get.return_value = None
mock_db.session.query.return_value.where.return_value.first.return_value = refreshed_segment
mock_db.session.get.return_value = refreshed_segment
result = SegmentService.update_segment(
SegmentUpdateArgs(content="same content", regenerate_child_chunks=True),