From 92978f6d9fa02289ecd9fc6e1bbd52d243fe2b08 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Fri, 21 Aug 2026 06:50:14 +0000 Subject: [PATCH] test: migrate console dataset segment sessions and ORM models to SQLite (#40515) Co-authored-by: Byron.wang --- .../datasets/test_datasets_segments.py | 574 ++++++++++-------- 1 file changed, 313 insertions(+), 261 deletions(-) diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py index b2485d55d0f..9b258aec8bc 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets_segments.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock, call, patch import pytest from flask import Flask +from sqlalchemy.orm import Session from werkzeug.exceptions import Forbidden, NotFound import services @@ -30,8 +31,9 @@ from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.rag.index_processor.constant.index_type import IndexStructureType from fields.segment_fields import segment_response_with_summary, segment_responses_with_summaries from libs.datetime_utils import naive_utc_now -from models.dataset import ChildChunk, DocumentSegment -from models.enums import SegmentStatus, SegmentType +from models.account import Account, TenantAccountRole +from models.dataset import ChildChunk, Dataset, Document, DocumentSegment +from models.enums import PermissionEnum, SegmentStatus, SegmentType from models.model import UploadFile from services.errors.chunk import ChildChunkDeleteIndexError as ChildChunkDeleteIndexServiceError from services.errors.chunk import ChildChunkIndexingError as ChildChunkIndexingServiceError @@ -79,6 +81,82 @@ def _child_chunk(): return child_chunk +def _account() -> Account: + account = Account(name="Dataset Editor", email="dataset-editor@example.com") + account.id = "u1" + account.role = TenantAccountRole.OWNER + return account + + +def _dataset( + *, + dataset_id: str = "ds-1", + tenant_id: str = "tenant-1", + indexing_technique: str = "economy", + embedding_model_provider: str | None = None, + embedding_model: str | None = None, +) -> Dataset: + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + name="Dataset", + description="", + provider="vendor", + permission=PermissionEnum.ONLY_ME, + indexing_technique=indexing_technique, + embedding_model_provider=embedding_model_provider, + embedding_model=embedding_model, + created_by="u1", + ) + + +def _document( + *, + document_id: str = "doc-1", + dataset_id: str = "ds-1", + tenant_id: str = "tenant-1", + doc_form: IndexStructureType = IndexStructureType.PARAGRAPH_INDEX, +) -> Document: + return Document( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=1, + data_source_type="upload_file", + batch="batch-1", + name="Document", + created_from="api", + created_by="u1", + doc_form=doc_form, + ) + + +def _upload_file(*, name: str = "test.csv", file_id: str = "test-file-id") -> UploadFile: + upload_file = UploadFile( + tenant_id="tenant-1", + storage_type="opendal", + key="test-key", + name=name, + size=0, + extension=name.rsplit(".", maxsplit=1)[-1], + mime_type="text/csv", + created_by_role="account", + created_by="u1", + created_at=datetime.now(), + used=False, + ) + upload_file.id = file_id + return upload_file + + +class SQLiteControllerTest: + session: Session + + @pytest.fixture(autouse=True) + def _use_sqlite_session(self, sqlite_session: Session) -> None: + self.session = sqlite_session + + def _segment_response_dict(): return { "id": "seg-1", @@ -111,33 +189,22 @@ def _segment_response_dict(): } -def _bind_dataset_document(dataset, document, dataset_id: str = "ds-1", document_id: str = "doc-1"): - dataset.id = dataset_id - dataset.tenant_id = "tenant-1" - document.id = document_id - document.dataset_id = dataset_id - document.tenant_id = "tenant-1" - return document - - -def test_segment_response_with_summary(): +def test_segment_response_with_summary(sqlite_session: Session): segment = _segment() - session = MagicMock() with ( patch.object(DocumentSegment, "get_child_chunks", autospec=True, return_value=[]) as get_child_chunks, patch.object(DocumentSegment, "get_attachments", autospec=True, return_value=[]) as get_attachments, ): - result = segment_response_with_summary(segment, "summary", session=session) + result = segment_response_with_summary(segment, "summary", session=sqlite_session) assert result.summary == "summary" assert result.id == segment.id - get_child_chunks.assert_called_once_with(segment, session=session, include_full_doc=False) - get_attachments.assert_called_once_with(segment, session=session) + get_child_chunks.assert_called_once_with(segment, session=sqlite_session, include_full_doc=False) + get_attachments.assert_called_once_with(segment, session=sqlite_session) -def test_segment_responses_with_summaries_reuses_caller_session(): +def test_segment_responses_with_summaries_reuses_caller_session(sqlite_session: Session): segments = [_segment(), _segment()] segments[1].id = "seg-2" - session = MagicMock() expected_responses = [MagicMock(), MagicMock()] with patch( @@ -146,27 +213,24 @@ def test_segment_responses_with_summaries_reuses_caller_session(): responses = segment_responses_with_summaries( segments, {"seg-1": "summary-1", "seg-2": None}, - session=session, + session=sqlite_session, ) assert responses == expected_responses assert serialize_segment.call_args_list == [ - call(segments[0], "summary-1", session=session), - call(segments[1], None, session=session), + call(segments[0], "summary-1", session=sqlite_session), + call(segments[1], None, session=sqlite_session), ] -class TestDatasetDocumentSegmentListApi: +class TestDatasetDocumentSegmentListApi(SQLiteControllerTest): def test_get_success(self, app: Flask): api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - dataset = MagicMock() - document = MagicMock() - user = MagicMock() + dataset = _dataset() + document = _document() + user = _account() segment = _segment() - session = MagicMock() - session.get.return_value = None - session.execute.return_value.all.return_value = [] pagination = MagicMock() pagination.items = [segment] pagination.total = 1 @@ -182,25 +246,25 @@ class TestDatasetDocumentSegmentListApi: patch("controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination), patch("services.summary_index_service.SummaryIndexService.get_segments_summaries", return_value={}), ): - response, status = method(api, session, "tenant-1", user, "ds-1", "doc-1") + response, status = method(api, self.session, "tenant-1", user, "ds-1", "doc-1") assert status == 200 def test_get_dataset_not_found(self, app: Flask): api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - user = MagicMock() + user = _account() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1") def test_get_permission_denied(self, app: Flask): api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - dataset = MagicMock() - user = MagicMock() + dataset = _dataset() + user = _account() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -210,19 +274,16 @@ class TestDatasetDocumentSegmentListApi: ), ): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1") -class TestDatasetDocumentSegmentApi: +class TestDatasetDocumentSegmentApi(SQLiteControllerTest): def test_patch_success(self, app: Flask): api = DatasetDocumentSegmentApi() method = unwrap(api.patch) - user = MagicMock() - user.is_dataset_editor = True - dataset = MagicMock() - dataset.indexing_technique = "economy" - document = MagicMock() - document.id = "doc-1" + user = _account() + dataset = _dataset() + document = _document() with ( app.test_request_context("/?segment_id=s1&segment_id=s2"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -237,19 +298,16 @@ class TestDatasetDocumentSegmentApi: return_value=None, ), ): - response, status = method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "enable") + response, status = method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "enable") assert status == 200 assert response["result"] == "success" def test_patch_document_indexing_in_progress(self, app: Flask): api = DatasetDocumentSegmentApi() method = unwrap(api.patch) - user = MagicMock() - user.is_dataset_editor = True - dataset = MagicMock() - dataset.indexing_technique = "economy" - document = MagicMock() - document.id = "doc-1" + user = _account() + dataset = _dataset() + document = _document() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -265,16 +323,16 @@ class TestDatasetDocumentSegmentApi: patch("controllers.console.datasets.datasets_segments.redis_client.get", return_value=b"running"), ): with pytest.raises(InvalidActionError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "disable") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "disable") def test_patch_llm_bad_request(self, app: Flask): api = DatasetDocumentSegmentApi() method = unwrap(api.patch) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( + user = _account() + dataset = _dataset( indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock(id="doc-1") + document = _document() with ( app.test_request_context("/?segment_id=s1"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -293,16 +351,16 @@ class TestDatasetDocumentSegmentApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "enable") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "enable") def test_patch_provider_token_not_init(self, app: Flask): api = DatasetDocumentSegmentApi() method = unwrap(api.patch) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( + user = _account() + dataset = _dataset( indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock(id="doc-1") + document = _document() with ( app.test_request_context("/?segment_id=s1"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -321,25 +379,18 @@ class TestDatasetDocumentSegmentApi: ), ): with pytest.raises(ProviderNotInitializeError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "enable") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "enable") -class TestDatasetDocumentSegmentAddApi: +class TestDatasetDocumentSegmentAddApi(SQLiteControllerTest): def test_post_success(self, app: Flask): api = DatasetDocumentSegmentAddApi() method = unwrap(api.post) payload = {"content": "hello"} - user = MagicMock() - user.is_dataset_editor = True - dataset = MagicMock() - dataset.indexing_technique = "economy" - document = MagicMock() - document.doc_form = IndexStructureType.PARAGRAPH_INDEX - _bind_dataset_document(dataset, document) + user = _account() + dataset = _dataset() + document = _document() segment = _segment() - session = MagicMock() - session.get.return_value = None - session.execute.return_value.all.return_value = [] with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -360,7 +411,7 @@ class TestDatasetDocumentSegmentAddApi: ), ): response, status = method( - api, SegmentCreatePayload(content="test content"), session, "tenant-1", user, "ds-1", "doc-1" + api, SegmentCreatePayload(content="test content"), self.session, "tenant-1", user, "ds-1", "doc-1" ) assert status == 200 assert response["data"]["id"] == "seg-1" @@ -369,11 +420,11 @@ class TestDatasetDocumentSegmentAddApi: api = DatasetDocumentSegmentAddApi() method = unwrap(api.post) payload = {"content": "x"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( + user = _account() + dataset = _dataset( indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock() + document = _document() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -386,18 +437,18 @@ class TestDatasetDocumentSegmentAddApi: ): with pytest.raises(ProviderNotInitializeError): method( - api, SegmentCreatePayload(content="test content"), MagicMock(), "tenant-1", user, "ds-1", "doc-1" + api, SegmentCreatePayload(content="test content"), self.session, "tenant-1", user, "ds-1", "doc-1" ) def test_post_provider_token_not_init(self, app: Flask): api = DatasetDocumentSegmentAddApi() method = unwrap(api.post) payload = {"content": "x"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( + user = _account() + dataset = _dataset( indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock() + document = _document() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -410,26 +461,19 @@ class TestDatasetDocumentSegmentAddApi: ): with pytest.raises(ProviderNotInitializeError): method( - api, SegmentCreatePayload(content="test content"), MagicMock(), "tenant-1", user, "ds-1", "doc-1" + api, SegmentCreatePayload(content="test content"), self.session, "tenant-1", user, "ds-1", "doc-1" ) -class TestDatasetDocumentSegmentUpdateApi: +class TestDatasetDocumentSegmentUpdateApi(SQLiteControllerTest): def test_patch_success(self, app: Flask): api = DatasetDocumentSegmentUpdateApi() method = unwrap(api.patch) payload = {"content": "updated"} - user = MagicMock() - user.is_dataset_editor = True - dataset = MagicMock() - dataset.indexing_technique = "economy" - document = MagicMock() - document.doc_form = IndexStructureType.PARAGRAPH_INDEX - _bind_dataset_document(dataset, document) + user = _account() + dataset = _dataset() + document = _document() segment = _segment() - session = MagicMock() - session.get.return_value = None - session.execute.return_value.all.return_value = [] with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -453,7 +497,14 @@ class TestDatasetDocumentSegmentUpdateApi: ), ): response, status = method( - api, SegmentUpdatePayload(content="test content"), session, "tenant-1", user, "ds-1", "doc-1", "seg-1" + api, + SegmentUpdatePayload(content="test content"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", ) assert status == 200 assert "data" in response @@ -462,9 +513,9 @@ class TestDatasetDocumentSegmentUpdateApi: api = DatasetDocumentSegmentUpdateApi() method = unwrap(api.patch) payload = {"content": "updated"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock(id="ds-1", tenant_id="tenant-1", indexing_technique="economy") - document = MagicMock(id="doc-1", dataset_id="other-dataset", tenant_id="tenant-1") + user = _account() + dataset = _dataset() + document = _document(dataset_id="other-dataset") with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -483,7 +534,7 @@ class TestDatasetDocumentSegmentUpdateApi: method( api, SegmentUpdatePayload(content="test content"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -495,10 +546,9 @@ class TestDatasetDocumentSegmentUpdateApi: api = DatasetDocumentSegmentUpdateApi() method = unwrap(api.patch) payload = {"content": "updated"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock(indexing_technique="economy") - document = MagicMock() - _bind_dataset_document(dataset, document) + user = _account() + dataset = _dataset() + document = _document() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -520,7 +570,7 @@ class TestDatasetDocumentSegmentUpdateApi: method( api, SegmentUpdatePayload(content="test content"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -532,11 +582,11 @@ class TestDatasetDocumentSegmentUpdateApi: api = DatasetDocumentSegmentUpdateApi() method = unwrap(api.patch) payload = {"content": "x"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock( + user = _account() + dataset = _dataset( indexing_technique="high_quality", embedding_model_provider="openai", embedding_model="text-embed" ) - document = MagicMock() + document = _document() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -559,7 +609,7 @@ class TestDatasetDocumentSegmentUpdateApi: method( api, SegmentUpdatePayload(content="test content"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -568,28 +618,17 @@ class TestDatasetDocumentSegmentUpdateApi: ) -class TestDatasetDocumentSegmentBatchImportApi: +class TestDatasetDocumentSegmentBatchImportApi(SQLiteControllerTest): def test_post_success(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = UploadFile( - tenant_id="tenant-id", - storage_type="opendal", - key="test-key", - name="test.csv", - size=0, - extension="txt", - mime_type="text/plain", - created_by_role="account", - created_by="account-id", - created_at=datetime.now(), - used=False, - ) - user = MagicMock(id="u1") - dataset = MagicMock(id="ds-1", tenant_id="tenant-1") - session = MagicMock() - session.scalar.return_value = upload_file + upload_file = _upload_file() + self.session.add(upload_file) + self.session.commit() + user = _account() + dataset = _dataset() + document = _document() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -599,7 +638,7 @@ class TestDatasetDocumentSegmentBatchImportApi: ) as get_dataset_for_tenant, patch( "controllers.console.datasets.datasets_segments.DatasetRefService.get_document_by_ref", - return_value=MagicMock(), + return_value=document, ) as get_document_by_ref, patch("controllers.console.datasets.datasets_segments.redis_client.setnx", return_value=True), patch( @@ -608,11 +647,17 @@ class TestDatasetDocumentSegmentBatchImportApi: ), ): response, status = method( - api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + api, + BatchImportPayload(upload_file_id="test-file-id"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) assert status == 200 assert response["job_status"] == "waiting" - get_dataset_for_tenant.assert_called_once_with("ds-1", "tenant-1", session=session) + get_dataset_for_tenant.assert_called_once_with("ds-1", "tenant-1", session=self.session) document_ref = get_document_by_ref.call_args.args[0] assert document_ref.dataset.tenant_id == "tenant-1" assert document_ref.dataset.dataset_id == "ds-1" @@ -622,9 +667,7 @@ class TestDatasetDocumentSegmentBatchImportApi: api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - user = MagicMock(id="u1") - session = MagicMock() - session.scalar.return_value = None + user = _account() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -637,7 +680,7 @@ class TestDatasetDocumentSegmentBatchImportApi: method( api, BatchImportPayload(upload_file_id="test-file-id"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -648,15 +691,13 @@ class TestDatasetDocumentSegmentBatchImportApi: api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - user = MagicMock(id="u1") - session = MagicMock() - session.scalar.return_value = None + user = _account() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( "controllers.console.datasets.datasets_segments.DatasetService.get_dataset_for_tenant", - return_value=MagicMock(), + return_value=_dataset(), ), patch( "controllers.console.datasets.datasets_segments.DatasetRefService.get_document_by_ref", @@ -667,7 +708,7 @@ class TestDatasetDocumentSegmentBatchImportApi: method( api, BatchImportPayload(upload_file_id="test-file-id"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -678,78 +719,92 @@ class TestDatasetDocumentSegmentBatchImportApi: api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - user = MagicMock(id="u1") - session = MagicMock() - session.scalar.return_value = None + user = _account() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( "controllers.console.datasets.datasets_segments.DatasetService.get_dataset_for_tenant", - return_value=MagicMock(), + return_value=_dataset(), ), patch( "controllers.console.datasets.datasets_segments.DatasetRefService.get_document_by_ref", - return_value=MagicMock(), + return_value=_document(), ), ): with pytest.raises(NotFound): method( - api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + api, + BatchImportPayload(upload_file_id="test-file-id"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) def test_post_invalid_file_type(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = MagicMock() - upload_file.name = "test.txt" - user = MagicMock(id="u1") - session = MagicMock() - session.scalar.return_value = upload_file + upload_file = _upload_file(name="test.txt") + self.session.add(upload_file) + self.session.commit() + user = _account() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( "controllers.console.datasets.datasets_segments.DatasetService.get_dataset_for_tenant", - return_value=MagicMock(), + return_value=_dataset(), ), patch( "controllers.console.datasets.datasets_segments.DatasetRefService.get_document_by_ref", - return_value=MagicMock(), + return_value=_document(), ), ): with pytest.raises(ValueError): method( - api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + api, + BatchImportPayload(upload_file_id="test-file-id"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) def test_post_async_task_failure(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) payload = {"upload_file_id": "file-1"} - upload_file = MagicMock() - upload_file.name = "test.csv" - user = MagicMock(id="u1") - session = MagicMock() - session.scalar.return_value = upload_file + upload_file = _upload_file() + self.session.add(upload_file) + self.session.commit() + user = _account() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), patch( "controllers.console.datasets.datasets_segments.DatasetService.get_dataset_for_tenant", - return_value=MagicMock(), + return_value=_dataset(), ), patch( "controllers.console.datasets.datasets_segments.DatasetRefService.get_document_by_ref", - return_value=MagicMock(), + return_value=_document(), ), patch( "controllers.console.datasets.datasets_segments.redis_client.setnx", side_effect=Exception("redis down") ), ): response, status = method( - api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + api, + BatchImportPayload(upload_file_id="test-file-id"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) assert status == 500 assert "error" in response @@ -765,7 +820,7 @@ class TestDatasetDocumentSegmentBatchImportApi: method(api, job_id="job-1") -class TestChildChunkAddApi: +class TestChildChunkAddApi(SQLiteControllerTest): def test_patch_documents_batch_update_payload(self): patch_method = cast(Any, ChildChunkAddApi.patch) api_doc = cast(dict[str, Any], patch_method.__apidoc__) @@ -775,8 +830,8 @@ class TestChildChunkAddApi: def test_get_uses_default_pagination_for_malformed_ints(self, app: Flask): api = ChildChunkAddApi() method = unwrap(api.get) - dataset = MagicMock() - document = _bind_dataset_document(dataset, MagicMock()) + dataset = _dataset() + document = _document() pagination = MagicMock(items=[], total=0, pages=0) with ( app.test_request_context("/?page=bad&limit="), @@ -788,32 +843,29 @@ class TestChildChunkAddApi: patch("controllers.console.datasets.datasets_segments.DocumentService.get_document", return_value=document), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_segment_by_ref", - return_value=MagicMock(), + return_value=_segment(), ), patch( "controllers.console.datasets.datasets_segments.SegmentService.get_child_chunks", return_value=pagination, ) as get_child_chunks, ): - response, status = method(api, MagicMock(), "tenant-1", "ds-1", "doc-1", "seg-1") + response, status = method(api, self.session, "tenant-1", "ds-1", "doc-1", "seg-1") assert status == 200 assert response["page"] == 1 assert response["limit"] == 20 session = get_child_chunks.call_args.kwargs["session"] - assert isinstance(session, MagicMock) + assert session is self.session assert get_child_chunks.call_args.args == ("seg-1", "doc-1", "ds-1", 1, 20, None) def test_post_success(self, app: Flask): api = ChildChunkAddApi() method = unwrap(api.post) payload = {"content": "child"} - user = MagicMock() - user.is_dataset_editor = True - dataset = MagicMock() - dataset.indexing_technique = "economy" - document = MagicMock() - _bind_dataset_document(dataset, document) - segment = MagicMock() + user = _account() + dataset = _dataset() + document = _document() + segment = _segment() child_chunk = _child_chunk() with ( app.test_request_context("/", json=payload), @@ -833,7 +885,14 @@ class TestChildChunkAddApi: ), ): response, status = method( - api, ChildChunkCreatePayload(content="child"), MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1" + api, + ChildChunkCreatePayload(content="child"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", + "seg-1", ) assert status == 200 assert response["data"]["id"] == "cc-1" @@ -842,11 +901,10 @@ class TestChildChunkAddApi: api = ChildChunkAddApi() method = unwrap(api.post) payload = {"content": "child"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock(indexing_technique="economy") - document = MagicMock() - _bind_dataset_document(dataset, document) - segment = MagicMock() + user = _account() + dataset = _dataset() + document = _document() + segment = _segment() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -868,7 +926,7 @@ class TestChildChunkAddApi: method( api, ChildChunkCreatePayload(content="child"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -880,10 +938,9 @@ class TestChildChunkAddApi: api = ChildChunkAddApi() method = unwrap(api.post) payload = {"content": "child"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock(indexing_technique="economy") - document = MagicMock() - _bind_dataset_document(dataset, document) + user = _account() + dataset = _dataset() + document = _document() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -898,7 +955,7 @@ class TestChildChunkAddApi: method( api, ChildChunkCreatePayload(content="child"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -907,17 +964,15 @@ class TestChildChunkAddApi: ) -class TestChildChunkUpdateApi: +class TestChildChunkUpdateApi(SQLiteControllerTest): def test_delete_success(self, app: Flask): api = ChildChunkUpdateApi() method = unwrap(api.delete) - user = MagicMock() - user.is_dataset_editor = True - dataset = MagicMock() - document = MagicMock() - _bind_dataset_document(dataset, document) - segment = MagicMock() - child_chunk = MagicMock() + user = _account() + dataset = _dataset() + document = _document() + segment = _segment() + child_chunk = _child_chunk() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -937,19 +992,18 @@ class TestChildChunkUpdateApi: "controllers.console.datasets.datasets_segments.SegmentService.delete_child_chunk", return_value=None ), ): - response, status = method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + response, status = method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") assert status == 204 assert response == "" def test_delete_child_chunk_index_error(self, app: Flask): api = ChildChunkUpdateApi() method = unwrap(api.delete) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() - document = MagicMock() - _bind_dataset_document(dataset, document) - segment = MagicMock() - child_chunk = MagicMock() + user = _account() + dataset = _dataset() + document = _document() + segment = _segment() + child_chunk = _child_chunk() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -971,16 +1025,15 @@ class TestChildChunkUpdateApi: ), ): with pytest.raises(ChildChunkDeleteIndexError): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") def test_delete_child_chunk_not_found(self, app: Flask): api = ChildChunkUpdateApi() method = unwrap(api.delete) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() - document = MagicMock() - _bind_dataset_document(dataset, document) - segment = MagicMock() + user = _account() + dataset = _dataset() + document = _document() + segment = _segment() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=dataset), @@ -998,17 +1051,16 @@ class TestChildChunkUpdateApi: ), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1", "seg-1", "cc-1") def test_patch_child_chunk_not_found(self, app: Flask): api = ChildChunkUpdateApi() method = unwrap(api.patch) payload = {"content": "updated child"} - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() - document = MagicMock() - _bind_dataset_document(dataset, document) - segment = MagicMock() + user = _account() + dataset = _dataset() + document = _document() + segment = _segment() with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -1030,7 +1082,7 @@ class TestChildChunkUpdateApi: method( api, ChildChunkUpdatePayload(content="updated child"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -1040,17 +1092,14 @@ class TestChildChunkUpdateApi: ) -class TestSegmentListAdvancedCases: +class TestSegmentListAdvancedCases(SQLiteControllerTest): def test_segment_list_with_keyword_filter(self, app: Flask): api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - dataset = MagicMock() - document = MagicMock() - user = MagicMock() + dataset = _dataset() + document = _document() + user = _account() segment = _segment() - session = MagicMock() - session.get.return_value = None - session.execute.return_value.all.return_value = [] pagination = MagicMock(items=[segment], total=1, pages=1) with ( app.test_request_context("/?keyword=test"), @@ -1063,7 +1112,7 @@ class TestSegmentListAdvancedCases: patch("controllers.console.datasets.datasets_segments.paginate_query", return_value=pagination), patch("services.summary_index_service.SummaryIndexService.get_segments_summaries", return_value={}), ): - result = method(api, session, "tenant-1", user, "ds-1", "doc-1") + result = method(api, self.session, "tenant-1", user, "ds-1", "doc-1") if isinstance(result, tuple): response, status = result else: @@ -1074,9 +1123,12 @@ class TestSegmentListAdvancedCases: def test_segment_list_postgres_keyword_filter_handles_scalar_keywords(self, app: Flask): api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - dataset = MagicMock() - document = MagicMock() - user = MagicMock() + dataset = _dataset(dataset_id="22222222-2222-2222-2222-222222222222") + document = _document( + document_id="33333333-3333-3333-3333-333333333333", + dataset_id=dataset.id, + ) + user = _account() pagination = MagicMock(items=[], total=0, pages=0) with ( app.test_request_context("/?keyword=test"), @@ -1096,7 +1148,7 @@ class TestSegmentListAdvancedCases: ): method( api, - MagicMock(), + self.session, "11111111-1111-1111-1111-111111111111", user, "22222222-2222-2222-2222-222222222222", @@ -1111,41 +1163,39 @@ class TestSegmentListAdvancedCases: """Test segment list with permission denied""" api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - user = MagicMock() + user = _account() with ( app.test_request_context("/"), - patch( - "controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=MagicMock() - ), + patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=_dataset()), patch( "controllers.console.datasets.datasets_segments.DatasetService.check_dataset_permission", side_effect=services.errors.account.NoPermissionError("No permission"), ), ): with pytest.raises(Forbidden): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1") def test_segment_list_dataset_not_found(self, app: Flask): """Test segment list with dataset not found""" api = DatasetDocumentSegmentListApi() method = unwrap(api.get) - user = MagicMock() + user = _account() with ( app.test_request_context("/"), patch("controllers.console.datasets.datasets_segments.DatasetService.get_dataset", return_value=None), ): with pytest.raises(NotFound): - method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1") + method(api, self.session, "tenant-1", user, "ds-1", "doc-1") -class TestSegmentOperationCases: +class TestSegmentOperationCases(SQLiteControllerTest): def test_segment_add_with_provider_token_error(self, app: Flask): """Test segment add with provider token not initialized""" api = DatasetDocumentSegmentAddApi() method = unwrap(api.post) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() - document = MagicMock() + user = _account() + dataset = _dataset() + document = _document() payload = {"content": "new content", "answer": None} with ( app.test_request_context("/", json=payload), @@ -1163,15 +1213,21 @@ class TestSegmentOperationCases: ): with pytest.raises(ProviderTokenNotInitError): method( - api, SegmentCreatePayload(content="test content"), MagicMock(), "tenant-1", user, "ds-1", "doc-1" + api, + SegmentCreatePayload(content="test content"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) def test_batch_import_with_document_not_found(self, app: Flask): """Test batch import with document not found""" api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() + user = _account() + dataset = _dataset() payload = {"upload_file_id": "file-1"} with ( app.test_request_context("/", json=payload), @@ -1189,7 +1245,7 @@ class TestSegmentOperationCases: method( api, BatchImportPayload(upload_file_id="test-file-id"), - MagicMock(), + self.session, "tenant-1", user, "ds-1", @@ -1200,13 +1256,10 @@ class TestSegmentOperationCases: """Test batch import with invalid file type""" api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() - document = MagicMock() - upload_file = None + user = _account() + dataset = _dataset() + document = _document() payload = {"upload_file_id": "file-1"} - session = MagicMock() - session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -1221,32 +1274,25 @@ class TestSegmentOperationCases: ): with pytest.raises(NotFound): method( - api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + api, + BatchImportPayload(upload_file_id="test-file-id"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) def test_batch_import_with_async_task_failure(self, app: Flask): api = DatasetDocumentSegmentBatchImportApi() method = unwrap(api.post) - user = MagicMock(is_dataset_editor=True) - dataset = MagicMock() - document = MagicMock() - upload_file = UploadFile( - tenant_id="tenant-id", - storage_type="opendal", - key="test-key", - name="test.csv", - size=0, - extension="csv", - mime_type="text/csv", - created_by_role="account", - created_by="account-id", - created_at=datetime.now(), - used=False, - ) - upload_file.id = "file-1" + user = _account() + dataset = _dataset() + document = _document() + upload_file = _upload_file() + self.session.add(upload_file) + self.session.commit() payload = {"upload_file_id": "file-1"} - session = MagicMock() - session.scalar.return_value = upload_file with ( app.test_request_context("/", json=payload), patch.object(type(console_ns), "payload", payload), @@ -1264,7 +1310,13 @@ class TestSegmentOperationCases: ), ): response, status = method( - api, BatchImportPayload(upload_file_id="test-file-id"), session, "tenant-1", user, "ds-1", "doc-1" + api, + BatchImportPayload(upload_file_id="test-file-id"), + self.session, + "tenant-1", + user, + "ds-1", + "doc-1", ) assert status == 500 assert "error" in response