diff --git a/api/tests/unit_tests/events/event_handlers/test_create_document_index.py b/api/tests/unit_tests/events/event_handlers/test_create_document_index.py index 4b0ca5291aa..fad8932f9e2 100644 --- a/api/tests/unit_tests/events/event_handlers/test_create_document_index.py +++ b/api/tests/unit_tests/events/event_handlers/test_create_document_index.py @@ -2,10 +2,8 @@ import logging from unittest.mock import MagicMock, patch import pytest -from sqlalchemy.engine import Engine -from sqlalchemy.orm import Session, sessionmaker +from sqlalchemy.orm import Session -import core.db.session_factory as session_factory_module from core.indexing_runner import DocumentIsPausedError from events.event_handlers import create_document_index as handler_module from models.dataset import Document @@ -14,13 +12,8 @@ from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus @pytest.fixture def persisted_document( - sqlite_engine: Engine, sqlite_session: Session, - monkeypatch: pytest.MonkeyPatch, ) -> Document: - real_session_maker = sessionmaker(bind=sqlite_engine, expire_on_commit=False) - monkeypatch.setattr(session_factory_module, "_session_maker", real_session_maker) - document = Document( id="doc-1", tenant_id="tenant-1", diff --git a/api/tests/unit_tests/events/event_handlers/test_delete_tool_parameters_cache_when_sync_draft_workflow.py b/api/tests/unit_tests/events/event_handlers/test_delete_tool_parameters_cache_when_sync_draft_workflow.py index 27f162de75e..7b13f337611 100644 --- a/api/tests/unit_tests/events/event_handlers/test_delete_tool_parameters_cache_when_sync_draft_workflow.py +++ b/api/tests/unit_tests/events/event_handlers/test_delete_tool_parameters_cache_when_sync_draft_workflow.py @@ -1,43 +1,64 @@ import logging -from types import SimpleNamespace import pytest from core.tools.errors import ToolProviderNotFoundError from events.event_handlers import delete_tool_parameters_cache_when_sync_draft_workflow as handler_module +from graphon.nodes.tool.entities import ToolEntity, ToolProviderType +from models.model import App, AppMode, IconType +from models.workflow import Workflow def test_missing_tool_provider_does_not_log_error_traceback( monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, ): - app = SimpleNamespace(id="workflow-id", tenant_id="tenant-id") - workflow = SimpleNamespace( - graph_dict={ - "nodes": [ - { - "id": "node-id", - "data": { - "type": "tool", - }, - } - ] - } + app = App( + id="workflow-id", + tenant_id="tenant-id", + name="Workflow app", + description="", + mode=AppMode.WORKFLOW, + icon_type=IconType.EMOJI, + icon="workflow", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, ) - tool_entity = SimpleNamespace( - provider_type=SimpleNamespace(value="mcp"), + workflow = Workflow( + tenant_id=app.tenant_id, + app_id=app.id, + type="workflow", + version="draft", + graph='{"nodes": [{"id": "node-id", "data": {"type": "tool"}}]}', + features="{}", + created_by="account-id", + environment_variables=[], + conversation_variables=[], + ) + tool_entity = ToolEntity( + provider_type=ToolProviderType.MCP, provider_id="my-test-mcp-server", provider_name="my-test-mcp-server", tool_name="echo", + tool_label="Echo", + tool_configurations={}, credential_id=None, ) - monkeypatch.setattr(handler_module, "adapt_node_config_for_graph", lambda node_data: {"data": node_data["data"]}) - monkeypatch.setattr(handler_module.ToolEntity, "model_validate", lambda data: tool_entity) + monkeypatch.setattr( + handler_module, + "adapt_node_config_for_graph", + lambda node_data: { + "data": node_data["data"], + }, + ) + monkeypatch.setattr(handler_module.ToolEntity, "model_validate", lambda _data: tool_entity) monkeypatch.setattr( handler_module.ToolManager, "get_tool_runtime", - lambda **kwargs: (_ for _ in ()).throw(ToolProviderNotFoundError("mcp provider not found")), + lambda **_kwargs: (_ for _ in ()).throw(ToolProviderNotFoundError("mcp provider not found")), ) with caplog.at_level(logging.INFO, logger=handler_module.logger.name): diff --git a/api/tests/unit_tests/events/event_handlers/test_queue_default_plugin_install_when_tenant_created.py b/api/tests/unit_tests/events/event_handlers/test_queue_default_plugin_install_when_tenant_created.py index 18264cf3510..67c8c8e827f 100644 --- a/api/tests/unit_tests/events/event_handlers/test_queue_default_plugin_install_when_tenant_created.py +++ b/api/tests/unit_tests/events/event_handlers/test_queue_default_plugin_install_when_tenant_created.py @@ -1,10 +1,16 @@ import logging -from types import SimpleNamespace from unittest.mock import MagicMock import pytest from events.event_handlers import queue_default_plugin_install_when_tenant_created as handler_module +from models.account import Tenant + + +def _tenant() -> Tenant: + tenant = Tenant(name="Test tenant") + tenant.id = "tenant-1" + return tenant def test_handle_skips_when_no_default_plugins_are_configured(monkeypatch: pytest.MonkeyPatch) -> None: @@ -12,7 +18,7 @@ def test_handle_skips_when_no_default_plugins_are_configured(monkeypatch: pytest monkeypatch.setattr(handler_module.dify_config, "NEW_USER_DEFAULT_PLUGIN_IDS", "") monkeypatch.setattr(handler_module.install_default_plugins_task, "delay", delay) - handler_module.handle(SimpleNamespace(id="tenant-1")) + handler_module.handle(_tenant()) delay.assert_not_called() @@ -26,7 +32,7 @@ def test_handle_queues_configured_plugins(monkeypatch: pytest.MonkeyPatch) -> No monkeypatch.setattr(handler_module.dify_config, "NEW_USER_DEFAULT_PLUGIN_IDS", ",".join(plugins)) monkeypatch.setattr(handler_module.install_default_plugins_task, "delay", delay) - handler_module.handle(SimpleNamespace(id="tenant-1")) + handler_module.handle(_tenant()) delay.assert_called_once_with("tenant-1", plugins) @@ -47,6 +53,6 @@ def test_handle_does_not_fail_tenant_creation_when_queue_is_unavailable( ) with caplog.at_level(logging.ERROR, logger=handler_module.logger.name): - handler_module.handle(SimpleNamespace(id="tenant-1")) + handler_module.handle(_tenant()) assert "Failed to queue default plugin installation for tenant tenant-1" in caplog.text diff --git a/api/tests/unit_tests/events/test_app_event_signals.py b/api/tests/unit_tests/events/test_app_event_signals.py index 6feedf01273..35471eab0de 100644 --- a/api/tests/unit_tests/events/test_app_event_signals.py +++ b/api/tests/unit_tests/events/test_app_event_signals.py @@ -1,11 +1,10 @@ import json from collections.abc import Iterator -from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import patch from uuid import uuid4 import pytest -from sqlalchemy import event +from sqlalchemy import event, select from sqlalchemy.orm import Session from events.app_event import app_was_deleted, app_was_updated @@ -223,57 +222,105 @@ class TestAppWasUpdatedSignal: class TestAppModelConfigWasUpdatedSignal: - def test_requires_caller_session(self) -> None: + def test_requires_caller_session(self, app_model: App) -> None: from events.event_handlers.update_app_dataset_join_when_app_model_config_updated import handle with pytest.raises(TypeError, match="session"): - handle(SimpleNamespace(id="app-1"), app_model_config=None) + handle(app_model, app_model_config=None) - def test_reuses_provided_session_without_committing(self) -> None: + @pytest.mark.parametrize("sqlite_session", [(App, Account, AppModelConfig, AppDatasetJoin)], indirect=True) + def test_reuses_provided_session_without_committing(self, app_model: App, sqlite_session: Session) -> None: from events.event_handlers.update_app_dataset_join_when_app_model_config_updated import handle - session = MagicMock() - session.scalars.return_value.all.return_value = [] - app_model_config = AppModelConfig(app_id="app-1", created_by="user-1", updated_by="user-1") + app_model_config = AppModelConfig(app_id=app_model.id, created_by="user-1", updated_by="user-1") app_model_config.dataset_configs = json.dumps( { "retrieval_model": "multiple", "datasets": {"datasets": [{"dataset": {"id": "dataset-1"}}]}, } ) + sqlite_session.add_all( + [ + app_model_config, + AppDatasetJoin(app_id="other-app", dataset_id="dataset-1"), + ] + ) + sqlite_session.flush() + commits: list[str] = [] - handle(SimpleNamespace(id="app-1"), app_model_config=app_model_config, session=session) + def after_commit(_session: Session) -> None: + commits.append("commit") - added_join = session.add.call_args.args[0] - assert isinstance(added_join, AppDatasetJoin) - assert added_join.app_id == "app-1" - assert added_join.dataset_id == "dataset-1" - session.commit.assert_not_called() + event.listen(sqlite_session, "after_commit", after_commit) + try: + handle(app_model, app_model_config=app_model_config, session=sqlite_session) + finally: + event.remove(sqlite_session, "after_commit", after_commit) + + added_joins = [row for row in sqlite_session.new if isinstance(row, AppDatasetJoin)] + assert len(added_joins) == 1 + assert added_joins[0].app_id == app_model.id + assert added_joins[0].dataset_id == "dataset-1" + assert commits == [] + + sqlite_session.flush() + persisted = sqlite_session.scalars(select(AppDatasetJoin).where(AppDatasetJoin.app_id == app_model.id)).all() + assert [(row.app_id, row.dataset_id) for row in persisted] == [(app_model.id, "dataset-1")] +@pytest.mark.parametrize("sqlite_session", [(App, Account, InstalledApp)], indirect=True) class TestCreateInstalledAppWhenAppCreated: - def test_skips_existing_installation(self) -> None: + def test_skips_existing_installation(self, app_model: App, sqlite_session: Session) -> None: from events.event_handlers.create_installed_app_when_app_created import handle - session = MagicMock() - session.scalar.return_value = "installed-app-1" + existing = InstalledApp( + tenant_id=app_model.tenant_id, + app_id=app_model.id, + app_owner_tenant_id=app_model.tenant_id, + ) + sqlite_session.add(existing) + sqlite_session.commit() - handle(SimpleNamespace(id="app-1", tenant_id="tenant-1"), session=session) + handle(app_model, session=sqlite_session) - session.add.assert_not_called() - session.flush.assert_not_called() + installations = sqlite_session.scalars( + select(InstalledApp).where( + InstalledApp.tenant_id == app_model.tenant_id, + InstalledApp.app_id == app_model.id, + ) + ).all() + assert installations == [existing] - def test_adds_missing_installation_without_committing(self) -> None: + def test_adds_missing_installation_without_committing(self, app_model: App, sqlite_session: Session) -> None: from events.event_handlers.create_installed_app_when_app_created import handle - session = MagicMock() - session.scalar.return_value = None + sqlite_session.add( + InstalledApp( + tenant_id="other-tenant", + app_id=app_model.id, + app_owner_tenant_id="other-tenant", + ) + ) + sqlite_session.commit() + commits: list[str] = [] - handle(SimpleNamespace(id="app-1", tenant_id="tenant-1"), session=session) + def after_commit(_session: Session) -> None: + commits.append("commit") + + event.listen(sqlite_session, "after_commit", after_commit) + try: + handle(app_model, session=sqlite_session) + finally: + event.remove(sqlite_session, "after_commit", after_commit) + + installed_app = sqlite_session.scalar( + select(InstalledApp).where( + InstalledApp.tenant_id == app_model.tenant_id, + InstalledApp.app_id == app_model.id, + ) + ) - installed_app = session.add.call_args.args[0] assert isinstance(installed_app, InstalledApp) - assert installed_app.app_id == "app-1" - assert installed_app.tenant_id == "tenant-1" - session.flush.assert_called_once_with() - session.commit.assert_not_called() + assert installed_app.app_id == app_model.id + assert installed_app.tenant_id == app_model.tenant_id + assert commits == [] diff --git a/api/tests/unit_tests/events/test_update_provider_when_message_created.py b/api/tests/unit_tests/events/test_update_provider_when_message_created.py index 7c55c949fab..248709c6be4 100644 --- a/api/tests/unit_tests/events/test_update_provider_when_message_created.py +++ b/api/tests/unit_tests/events/test_update_provider_when_message_created.py @@ -5,7 +5,6 @@ from uuid import uuid4 import pytest from sqlalchemy import select -from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import AgentAppGenerateEntity, ChatAppGenerateEntity @@ -17,11 +16,10 @@ from models.provider import ProviderType @pytest.fixture -def credit_pool_session_factory(sqlite_engine: Engine) -> Iterator[sessionmaker[Session]]: +def credit_pool_session_factory(sqlite_session_factory: sessionmaker[Session]) -> Iterator[sessionmaker[Session]]: """Bind message-created accounting to fixture-owned SQLite sessions.""" - session_factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False) - with patch("events.event_handlers.update_provider_when_message_created.db.session", session_factory): - yield session_factory + with patch("events.event_handlers.update_provider_when_message_created.db.session", sqlite_session_factory): + yield sqlite_session_factory def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_insufficient( diff --git a/api/tests/unit_tests/fields/test_conversation_variable_fields.py b/api/tests/unit_tests/fields/test_conversation_variable_fields.py index 813aa8e44c7..88bd321e548 100644 --- a/api/tests/unit_tests/fields/test_conversation_variable_fields.py +++ b/api/tests/unit_tests/fields/test_conversation_variable_fields.py @@ -1,7 +1,6 @@ from __future__ import annotations from datetime import UTC, datetime -from types import SimpleNamespace import pytest @@ -30,7 +29,7 @@ def test_conversation_variable_response_normalizes_callable_exposed_type() -> No { "id": "550e8400-e29b-41d4-a716-446655440000", "name": "foo", - "value_type": SimpleNamespace(exposed_type=lambda: SegmentType.STRING.exposed_type()), + "value_type": SegmentType.STRING, } ) diff --git a/api/tests/unit_tests/fields/test_dataset_fields.py b/api/tests/unit_tests/fields/test_dataset_fields.py index caffaae09bf..71d4ef0a735 100644 --- a/api/tests/unit_tests/fields/test_dataset_fields.py +++ b/api/tests/unit_tests/fields/test_dataset_fields.py @@ -1,7 +1,10 @@ -from types import SimpleNamespace -from unittest.mock import Mock +import pytest +from sqlalchemy.orm import Session from fields.dataset_fields import DatasetDetailResponse, dataset_detail_response_source +from models.account import Account +from models.dataset import AppDatasetJoin, Dataset +from models.model import App, AppMode, IconType def _dataset_detail_payload(**overrides): @@ -184,45 +187,59 @@ def test_dataset_detail_expands_missing_weighted_score_nested_fields(): } -def test_dataset_detail_response_source_uses_caller_session_for_database_fields(): - session = Mock() - getter_mocks = { - "get_app_count": Mock(return_value=3), - "get_document_count": Mock(return_value=4), - "get_word_count": Mock(return_value=500), - "get_author_name": Mock(return_value="Ada"), - "get_tags": Mock(return_value=[{"id": "tag-1", "name": "Tag", "type": "knowledge"}]), - "get_doc_form": Mock(return_value="paragraph"), - "get_external_knowledge_info": Mock( - return_value={ - "external_knowledge_id": "knowledge-id", - "external_knowledge_api_id": "api-id", - "external_knowledge_api_name": "api", - "external_knowledge_api_endpoint": "https://example.com", - } - ), - "get_doc_metadata": Mock(return_value=[{"id": "metadata-1", "name": "Metadata", "type": "string"}]), - "get_is_published": Mock(return_value=True), - "get_total_documents": Mock(return_value=4), - "get_total_available_documents": Mock(return_value=2), - } - dataset = SimpleNamespace(**_dataset_detail_payload(), **getter_mocks) +@pytest.mark.parametrize("sqlite_session", [(Dataset, Account, App, AppDatasetJoin)], indirect=True) +def test_dataset_detail_response_source_uses_caller_session_for_database_fields(sqlite_session: Session): + account = Account(name="Ada", email="ada@example.com") + account.id = "account-1" + dataset = Dataset( + id="ds-1", + tenant_id="tenant-1", + name="Dataset", + description="desc", + provider="vendor", + permission="only_me", + data_source_type=None, + indexing_technique="economy", + created_by=account.id, + retrieval_model=_dataset_detail_payload()["retrieval_model_dict"], + summary_index_setting=_dataset_detail_payload()["summary_index_setting"], + built_in_field_enabled=False, + icon_info=_dataset_detail_payload()["icon_info"], + runtime_mode="general", + enable_api=False, + is_multimodal=False, + ) + dataset.embedding_available = True + decoy_app = App( + id="decoy-app", + tenant_id="tenant-1", + name="Decoy app", + description="", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="app", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, + ) + decoy_join = AppDatasetJoin(app_id=decoy_app.id, dataset_id="other-dataset") + sqlite_session.add_all([account, dataset, decoy_app, decoy_join]) + sqlite_session.flush() response = DatasetDetailResponse.model_validate( - dataset_detail_response_source(dataset, session=session), + dataset_detail_response_source(dataset, session=sqlite_session), from_attributes=True, ) - assert response.app_count == 3 - assert response.document_count == 4 - assert response.word_count == 500 + assert response.app_count == 0 + assert response.document_count == 0 + assert response.word_count == 0 assert response.author_name == "Ada" - assert response.tags[0].id == "tag-1" - assert response.doc_form == "paragraph" - assert response.external_knowledge_info.external_knowledge_api_id == "api-id" - assert response.doc_metadata[0].id == "metadata-1" - assert response.is_published is True - assert response.total_documents == 4 - assert response.total_available_documents == 2 - for getter in getter_mocks.values(): - getter.assert_called_once_with(session=session) + assert response.tags == [] + assert response.doc_form is None + assert response.external_knowledge_info.external_knowledge_api_id is None + assert response.doc_metadata == [] + assert response.is_published is False + assert response.total_documents == 0 + assert response.total_available_documents == 0 diff --git a/api/tests/unit_tests/fields/test_document_fields.py b/api/tests/unit_tests/fields/test_document_fields.py index d3483e2ad75..c097de67b65 100644 --- a/api/tests/unit_tests/fields/test_document_fields.py +++ b/api/tests/unit_tests/fields/test_document_fields.py @@ -1,19 +1,59 @@ -from unittest.mock import MagicMock +import json +from datetime import datetime +import pytest +from sqlalchemy.orm import Session + +from extensions.storage.storage_type import StorageType from fields.document_fields import DocumentWithSession +from models.dataset import Document, DocumentSegment +from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom +from models.model import UploadFile -def test_document_with_session_uses_explicit_getters() -> None: - session = MagicMock() - document = MagicMock() - document.get_data_source_detail_dict.return_value = {"source": "detail"} - document.get_hit_count.return_value = 3 - document.get_doc_metadata_details.return_value = [{"name": "author"}] - source = DocumentWithSession(document=document, session=session) +@pytest.mark.parametrize("sqlite_session", [(Document, DocumentSegment, UploadFile)], indirect=True) +def test_document_with_session_uses_explicit_getters(sqlite_session: Session) -> None: + upload = UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key="documents/source.txt", + name="source.txt", + size=12, + extension=".txt", + mime_type="text/plain", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + created_at=datetime(2024, 1, 1), + used=True, + ) + document = Document( + id="document-1", + tenant_id="tenant-1", + dataset_id="dataset-1", + position=1, + data_source_type=DataSourceType.UPLOAD_FILE, + data_source_info=json.dumps({"upload_file_id": upload.id}), + batch="batch-1", + name="source.txt", + created_from=DocumentCreatedFrom.WEB, + created_by="account-1", + doc_metadata=None, + ) + segment = DocumentSegment( + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, + position=1, + content="hello world", + word_count=2, + tokens=2, + created_by="account-1", + hit_count=3, + ) + sqlite_session.add_all([upload, document, segment]) + sqlite_session.flush() + source = DocumentWithSession(document=document, session=sqlite_session) - assert source.data_source_detail_dict == {"source": "detail"} + assert source.data_source_detail_dict["upload_file"]["name"] == "source.txt" assert source.hit_count == 3 - assert source.doc_metadata_details == [{"name": "author"}] - document.get_data_source_detail_dict.assert_called_once_with(session=session) - document.get_hit_count.assert_called_once_with(session=session) - document.get_doc_metadata_details.assert_called_once_with(session=session) + assert source.doc_metadata_details is None diff --git a/api/tests/unit_tests/fields/test_file_fields.py b/api/tests/unit_tests/fields/test_file_fields.py index c35ca5f7b1d..000cd6880b7 100644 --- a/api/tests/unit_tests/fields/test_file_fields.py +++ b/api/tests/unit_tests/fields/test_file_fields.py @@ -1,34 +1,40 @@ from __future__ import annotations from datetime import datetime -from types import SimpleNamespace import pytest from core.workflow.file_reference import build_file_reference +from extensions.storage.storage_type import StorageType from fields import conversation_fields, message_fields from fields.file_fields import FileResponse, FileWithSignedUrl, RemoteFileInfo, UploadConfig from graphon.file import File, FileTransferMethod, FileType +from models.enums import CreatorUserRole +from models.model import UploadFile def test_file_response_serializes_datetime() -> None: created_at = datetime(2024, 1, 1, 12, 0, 0) - file_obj = SimpleNamespace( - id="file-1", + file_obj = UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key="key-1", name="example.txt", size=1024, extension="txt", mime_type="text/plain", + created_by_role=CreatorUserRole.ACCOUNT, created_by="user-1", created_at=created_at, - preview_url="https://preview", source_url="https://source", - original_url="https://origin", - user_id="user-1", - tenant_id="tenant-1", - conversation_id="conv-1", - file_key="key-1", + used=False, ) + file_obj.id = "file-1" + file_obj.preview_url = "https://preview" + file_obj.original_url = "https://origin" + file_obj.user_id = "user-1" + file_obj.conversation_id = "conv-1" + file_obj.file_key = "key-1" serialized = FileResponse.model_validate(file_obj, from_attributes=True).model_dump(mode="json") diff --git a/api/tests/unit_tests/fields/test_snippet_fields.py b/api/tests/unit_tests/fields/test_snippet_fields.py index 97573908576..f3119012c58 100644 --- a/api/tests/unit_tests/fields/test_snippet_fields.py +++ b/api/tests/unit_tests/fields/test_snippet_fields.py @@ -1,13 +1,22 @@ from datetime import UTC, datetime -from types import SimpleNamespace + +import pytest +from sqlalchemy.orm import Session from fields.snippet_fields import SnippetListItemResponse from libs.helper import dump_response +from models import snippet as snippet_module +from models.account import Account +from models.snippet import CustomizedSnippet -def test_snippet_list_fields_include_author_name() -> None: - snippet = SimpleNamespace( +@pytest.mark.parametrize("sqlite_session", [(CustomizedSnippet, Account)], indirect=True) +def test_snippet_list_fields_include_author_name(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> None: + account = Account(name="Alice", email="alice@example.com") + account.id = "account-1" + snippet = CustomizedSnippet( id="snippet-1", + tenant_id="tenant-1", name="Snippet", description="Reusable node", type="node", @@ -15,13 +24,14 @@ def test_snippet_list_fields_include_author_name() -> None: use_count=0, is_published=False, icon_info=None, - tags=[], created_by="account-1", - author_name="Alice", created_at=datetime.fromtimestamp(1704067200, tz=UTC), updated_by="account-1", updated_at=datetime.fromtimestamp(1704067201, tz=UTC), ) + sqlite_session.add_all([account, snippet]) + sqlite_session.flush() + monkeypatch.setattr(snippet_module.db, "session", sqlite_session) result = dump_response(SnippetListItemResponse, snippet) diff --git a/api/tests/unit_tests/models/test_account_models.py b/api/tests/unit_tests/models/test_account_models.py index d2dc7e9ae5f..364e09ddd04 100644 --- a/api/tests/unit_tests/models/test_account_models.py +++ b/api/tests/unit_tests/models/test_account_models.py @@ -12,7 +12,7 @@ This test suite covers: import base64 import secrets from datetime import UTC, datetime -from unittest.mock import MagicMock, patch +from unittest.mock import patch from uuid import uuid4 import pytest @@ -335,36 +335,67 @@ class TestTenantRelationshipIntegrity: # Assert assert tenant_id_none is None - def test_set_current_tenant_with_session_uses_caller_session(self): + @pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin)], indirect=True) + def test_set_current_tenant_with_session_uses_caller_session(self, sqlite_session: Session): account = Account(name="Test User", email="test@example.com") account.id = str(uuid4()) tenant = Tenant(name="Test Tenant") tenant.id = str(uuid4()) - join = MagicMock(role=TenantAccountRole.OWNER) - session = MagicMock(spec=Session) - session.scalar.return_value = join - session.scalars.return_value.one.return_value = tenant + decoy_tenant = Tenant(name="Decoy Tenant") + decoy_tenant.id = str(uuid4()) + sqlite_session.add_all( + [ + account, + tenant, + decoy_tenant, + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole.OWNER, + ), + TenantAccountJoin( + tenant_id=decoy_tenant.id, + account_id=account.id, + role=TenantAccountRole.NORMAL, + ), + ] + ) + sqlite_session.flush() - with patch("models.account.Session") as session_class: - account.set_current_tenant_with_session(tenant, session=session) + account.set_current_tenant_with_session(tenant, session=sqlite_session) - session_class.assert_not_called() assert account.current_tenant is tenant assert account.role == TenantAccountRole.OWNER - def test_set_tenant_id_with_session_uses_caller_session(self): + @pytest.mark.parametrize("sqlite_session", [(Account, Tenant, TenantAccountJoin)], indirect=True) + def test_set_tenant_id_with_session_uses_caller_session(self, sqlite_session: Session): account = Account(name="Test User", email="test@example.com") account.id = str(uuid4()) + decoy_account = Account(name="Decoy User", email="decoy@example.com") + decoy_account.id = str(uuid4()) tenant = Tenant(name="Test Tenant") tenant.id = str(uuid4()) - join = MagicMock(role=TenantAccountRole.ADMIN) - session = MagicMock(spec=Session) - session.execute.return_value.first.return_value = (tenant, join) + sqlite_session.add_all( + [ + account, + decoy_account, + tenant, + TenantAccountJoin( + tenant_id=tenant.id, + account_id=account.id, + role=TenantAccountRole.ADMIN, + ), + TenantAccountJoin( + tenant_id=tenant.id, + account_id=decoy_account.id, + role=TenantAccountRole.OWNER, + ), + ] + ) + sqlite_session.flush() - with patch("models.account.Session") as session_class: - account.set_tenant_id_with_session(tenant.id, session=session) + account.set_tenant_id_with_session(tenant.id, session=sqlite_session) - session_class.assert_not_called() assert account.current_tenant is tenant assert account.role == TenantAccountRole.ADMIN diff --git a/api/tests/unit_tests/models/test_dataset_models.py b/api/tests/unit_tests/models/test_dataset_models.py index e7a3a6f83ae..ff634dcdc50 100644 --- a/api/tests/unit_tests/models/test_dataset_models.py +++ b/api/tests/unit_tests/models/test_dataset_models.py @@ -12,17 +12,17 @@ This test suite covers: import json import pickle from datetime import UTC, datetime -from types import SimpleNamespace -from unittest.mock import Mock, call, patch +from unittest.mock import patch from urllib.parse import parse_qs, urlparse from uuid import uuid4 import pytest -from sqlalchemy.orm import Session +from sqlalchemy.orm import Session, scoped_session, sessionmaker from core.rag.entities import ParentMode from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType from extensions.storage.storage_type import StorageType +from models import dataset as dataset_module from models.account import Account from models.dataset import ( AppDatasetJoin, @@ -36,6 +36,7 @@ from models.dataset import ( Embedding, ExternalKnowledgeApis, ExternalKnowledgeBindings, + SegmentAttachmentBinding, ) from models.enums import ( CreatorUserRole, @@ -45,7 +46,84 @@ from models.enums import ( ProcessRuleMode, SegmentStatus, ) -from models.model import UploadFile +from models.model import App, AppMode, IconType, UploadFile + + +def _make_dataset( + *, + dataset_id: str = "dataset-1", + tenant_id: str = "tenant-1", + created_by: str = "account-1", + provider: str = "vendor", +) -> Dataset: + return Dataset( + id=dataset_id, + tenant_id=tenant_id, + name=f"Dataset {dataset_id}", + data_source_type=DataSourceType.UPLOAD_FILE, + created_by=created_by, + provider=provider, + built_in_field_enabled=False, + ) + + +def _make_document( + *, + document_id: str = "document-1", + dataset_id: str = "dataset-1", + tenant_id: str = "tenant-1", + process_rule_id: str | None = None, + position: int = 1, + word_count: int | None = None, + indexing_status: IndexingStatus = IndexingStatus.WAITING, +) -> Document: + return Document( + id=document_id, + tenant_id=tenant_id, + dataset_id=dataset_id, + position=position, + data_source_type=DataSourceType.UPLOAD_FILE, + dataset_process_rule_id=process_rule_id, + batch="batch-1", + name=f"{document_id}.txt", + created_from=DocumentCreatedFrom.WEB, + created_by="account-1", + word_count=word_count, + indexing_status=indexing_status, + ) + + +def _make_app(*, app_id: str, tenant_id: str = "tenant-1") -> App: + return App( + id=app_id, + tenant_id=tenant_id, + name=f"App {app_id}", + description="", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="app", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, + ) + + +def _make_segments(document: Document, hit_counts: list[int]) -> list[DocumentSegment]: + return [ + DocumentSegment( + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, + position=position, + content=f"Segment {position}", + word_count=2, + tokens=2, + created_by="account-1", + hit_count=hit_count, + ) + for position, hit_count in enumerate(hit_counts, start=1) + ] class TestDatasetModelValidation: @@ -92,104 +170,97 @@ class TestDatasetModelValidation: assert dataset.embedding_model == "text-embedding-ada-002" assert dataset.embedding_model_provider == "openai" - def test_session_aware_dataset_getters_use_caller_session(self): - dataset = Dataset( - tenant_id=str(uuid4()), - name="Test Dataset", - data_source_type=DataSourceType.UPLOAD_FILE, - created_by=str(uuid4()), - id=str(uuid4()), + @pytest.mark.parametrize("sqlite_session", [(Dataset, Account, DatasetProcessRule, Document)], indirect=True) + def test_session_aware_dataset_getters_use_caller_session(self, sqlite_session: Session): + account = Account(name="Ada", email="ada@example.com") + account.id = "account-1" + dataset = _make_dataset(created_by=account.id) + process_rule = DatasetProcessRule( + dataset_id=dataset.id, + mode=ProcessRuleMode.CUSTOM, + rules=json.dumps({"segmentation": {"max_tokens": 100}}), + created_by=account.id, ) - account = Mock() - process_rule = Mock() - session = Mock() - session.get.return_value = account - session.scalar.side_effect = [process_rule, IndexStructureType.PARAGRAPH_INDEX] + document = _make_document(dataset_id=dataset.id) + document.doc_form = IndexStructureType.PARAGRAPH_INDEX + sqlite_session.add_all([dataset, account, process_rule, document]) + sqlite_session.flush() - assert dataset.get_created_by_account(session=session) is account - assert dataset.get_latest_process_rule(session=session) is process_rule - assert dataset.get_doc_form(session=session) == IndexStructureType.PARAGRAPH_INDEX - - session.get.assert_called_once_with(Account, dataset.created_by) - assert session.scalar.call_count == 2 + assert dataset.get_created_by_account(session=sqlite_session) is account + assert dataset.get_latest_process_rule(session=sqlite_session) is process_rule + assert dataset.get_doc_form(session=sqlite_session) == IndexStructureType.PARAGRAPH_INDEX + @pytest.mark.parametrize("sqlite_session", [(Dataset, Document)], indirect=True) def test_get_doc_form_ignores_foreign_tenant_document(self, sqlite_session: Session) -> None: - dataset_id = str(uuid4()) - tenant_id = str(uuid4()) - created_by = str(uuid4()) - dataset = Dataset( - id=dataset_id, - tenant_id=tenant_id, - name="Dataset", - data_source_type=DataSourceType.UPLOAD_FILE, - created_by=created_by, - ) - foreign_document = Document( - tenant_id=str(uuid4()), - dataset_id=dataset_id, - position=1, - data_source_type=DataSourceType.UPLOAD_FILE, - batch="foreign", - name="Foreign", - created_from=DocumentCreatedFrom.WEB, - created_by=created_by, - doc_form=IndexStructureType.PARENT_CHILD_INDEX, + dataset = _make_dataset() + foreign_document = _make_document( + dataset_id=dataset.id, + tenant_id="tenant-2", ) + foreign_document.doc_form = IndexStructureType.PARENT_CHILD_INDEX sqlite_session.add_all([dataset, foreign_document]) + sqlite_session.flush() assert dataset.get_doc_form(session=sqlite_session) is None - def test_get_dataset_keyword_table_uses_caller_session(self): - dataset = Dataset( - tenant_id=str(uuid4()), - name="Test Dataset", - data_source_type=DataSourceType.UPLOAD_FILE, - created_by=str(uuid4()), - id=str(uuid4()), + @pytest.mark.parametrize("sqlite_session", [(Dataset, DatasetKeywordTable)], indirect=True) + def test_get_dataset_keyword_table_uses_caller_session(self, sqlite_session: Session): + dataset = _make_dataset() + keyword_table = DatasetKeywordTable( + dataset_id=dataset.id, + keyword_table=json.dumps({"keyword": ["node-1"]}), ) - keyword_table = Mock() - session = Mock() - session.scalar.return_value = keyword_table + sqlite_session.add_all([dataset, keyword_table]) + sqlite_session.flush() - result = dataset.get_dataset_keyword_table(session=session) + result = dataset.get_dataset_keyword_table(session=sqlite_session) assert result is keyword_table - session.scalar.assert_called_once() - def test_dataset_detail_getters_use_caller_session(self): - dataset = Dataset( - tenant_id=str(uuid4()), - name="Test Dataset", - data_source_type=DataSourceType.UPLOAD_FILE, - created_by=str(uuid4()), - provider="vendor", - built_in_field_enabled=False, - id=str(uuid4()), + @pytest.mark.parametrize("sqlite_session", [(Dataset, Account, App, AppDatasetJoin, Document)], indirect=True) + def test_dataset_detail_getters_use_caller_session(self, sqlite_session: Session): + account = Account(name="Ada", email="ada@example.com") + account.id = "account-1" + dataset = _make_dataset(created_by=account.id) + available_document = _make_document( + document_id="document-1", + dataset_id=dataset.id, + word_count=200, + indexing_status=IndexingStatus.COMPLETED, ) - account = Mock(name="account", name_value="Ada") - account.name = "Ada" - session = Mock() - session.get.return_value = account - session.scalar.side_effect = [2, 1, 3, 2, 500, IndexStructureType.PARAGRAPH_INDEX] - session.scalars.return_value.all.return_value = [] + available_document.enabled = True + available_document.archived = False + available_document.doc_form = IndexStructureType.PARAGRAPH_INDEX + waiting_document = _make_document( + document_id="document-2", + dataset_id=dataset.id, + position=2, + word_count=300, + ) + app = _make_app(app_id="app-1") + sqlite_session.add_all( + [ + account, + dataset, + available_document, + waiting_document, + app, + AppDatasetJoin(app_id=app.id, dataset_id=dataset.id), + ] + ) + sqlite_session.flush() - with patch("models.dataset.db") as mock_db: - assert dataset.get_total_documents(session=session) == 2 - assert dataset.get_total_available_documents(session=session) == 1 - assert dataset.get_app_count(session=session) == 3 - assert dataset.get_document_count(session=session) == 2 - assert dataset.get_word_count(session=session) == 500 - assert dataset.get_author_name(session=session) == "Ada" - assert dataset.get_tags(session=session) == [] - assert dataset.get_doc_form(session=session) == IndexStructureType.PARAGRAPH_INDEX - assert dataset.get_external_knowledge_info(session=session) is None - assert dataset.get_doc_metadata(session=session) == [] - assert dataset.get_is_published(session=session) is False - - assert session.scalar.call_count == 6 - assert session.scalars.call_count == 2 - mock_db.session.scalar.assert_not_called() - mock_db.session.scalars.assert_not_called() + assert dataset.get_total_documents(session=sqlite_session) == 2 + assert dataset.get_total_available_documents(session=sqlite_session) == 1 + assert dataset.get_app_count(session=sqlite_session) == 1 + assert dataset.get_document_count(session=sqlite_session) == 2 + assert dataset.get_word_count(session=sqlite_session) == 500 + assert dataset.get_author_name(session=sqlite_session) == "Ada" + assert dataset.get_tags(session=sqlite_session) == [] + assert dataset.get_doc_form(session=sqlite_session) == IndexStructureType.PARAGRAPH_INDEX + assert dataset.get_external_knowledge_info(session=sqlite_session) is None + assert dataset.get_doc_metadata(session=sqlite_session) == [] + assert dataset.get_is_published(session=sqlite_session) is False def test_dataset_indexing_technique_validation(self): """Test dataset indexing technique values.""" @@ -292,56 +363,82 @@ class TestDatasetModelValidation: assert result["top_k"] == 2 assert result["score_threshold"] == 0.0 - def test_dataset_external_knowledge_info_returns_none_for_cross_tenant_template(self): + @pytest.mark.parametrize( + "sqlite_session", [(Dataset, ExternalKnowledgeBindings, ExternalKnowledgeApis)], indirect=True + ) + def test_dataset_external_knowledge_info_returns_none_for_cross_tenant_template(self, sqlite_session: Session): """Test external datasets fail closed when the bound template is outside the tenant.""" - dataset = Dataset( - tenant_id=str(uuid4()), - name="External Dataset", - data_source_type=DataSourceType.UPLOAD_FILE, - created_by=str(uuid4()), - provider="external", + dataset = _make_dataset(provider="external") + external_api = ExternalKnowledgeApis( + tenant_id="other-tenant", + created_by="account-1", + updated_by=None, + name="Other tenant API", + description="", + settings=json.dumps({"endpoint": "https://example.com"}), ) binding = ExternalKnowledgeBindings( - tenant_id="tenant-id", - external_knowledge_api_id=str(uuid4()), - dataset_id="dataset-id", + tenant_id=dataset.tenant_id, + external_knowledge_api_id=external_api.id, + dataset_id=dataset.id, external_knowledge_id="knowledge-1", created_by="account-id", ) + sqlite_session.add_all([dataset, external_api, binding]) + sqlite_session.flush() - session = Mock() - session.scalar.side_effect = [binding, None] - with patch("models.dataset.db") as mock_db: - assert dataset.get_external_knowledge_info(session=session) is None + assert dataset.get_external_knowledge_info(session=sqlite_session) is None - assert session.scalar.call_count == 2 - mock_db.session.scalar.assert_not_called() - - def test_external_knowledge_api_dataset_bindings_use_caller_session(self): + @pytest.mark.parametrize( + "sqlite_session", [(ExternalKnowledgeApis, ExternalKnowledgeBindings, Dataset)], indirect=True + ) + def test_external_knowledge_api_dataset_bindings_use_caller_session(self, sqlite_session: Session): external_api = ExternalKnowledgeApis( - tenant_id=str(uuid4()), + tenant_id="tenant-1", created_by=str(uuid4()), updated_by=None, name="External API", description="", settings=None, ) - binding = Mock(dataset_id="dataset-1") - dataset = SimpleNamespace(id="dataset-1", name="Dataset") - session = Mock() - session.scalars.side_effect = [ - Mock(all=Mock(return_value=[binding])), - Mock(all=Mock(return_value=[dataset])), - ] + other_api = ExternalKnowledgeApis( + tenant_id="tenant-1", + created_by="account-1", + updated_by=None, + name="Other API", + description="", + settings=None, + ) + dataset = _make_dataset() + decoy_dataset = _make_dataset(dataset_id="dataset-2") + sqlite_session.add_all([external_api, other_api, dataset, decoy_dataset]) + sqlite_session.flush() + sqlite_session.add_all( + [ + ExternalKnowledgeBindings( + tenant_id="tenant-1", + external_knowledge_api_id=external_api.id, + dataset_id=dataset.id, + external_knowledge_id="knowledge-1", + created_by="account-1", + ), + ExternalKnowledgeBindings( + tenant_id="tenant-1", + external_knowledge_api_id=other_api.id, + dataset_id=decoy_dataset.id, + external_knowledge_id="knowledge-2", + created_by="account-1", + ), + ] + ) + sqlite_session.flush() - with patch("models.dataset.db") as mock_db: - result = external_api.get_dataset_bindings(session=session) + result = external_api.get_dataset_bindings(session=sqlite_session) - assert result == [{"id": "dataset-1", "name": "Dataset"}] - assert session.scalars.call_count == 2 - mock_db.session.scalars.assert_not_called() + assert result == [{"id": dataset.id, "name": dataset.name}] - def test_dataset_query_get_queries_uses_caller_session(self): + @pytest.mark.parametrize("sqlite_session", [(DatasetQuery, UploadFile)], indirect=True) + def test_dataset_query_get_queries_uses_caller_session(self, sqlite_session: Session): dataset_query = DatasetQuery( dataset_id=str(uuid4()), content=json.dumps([{"content_type": "image_query", "content": "file-1"}]), @@ -350,21 +447,25 @@ class TestDatasetModelValidation: created_by_role=CreatorUserRole.ACCOUNT, created_by=str(uuid4()), ) - upload_file = SimpleNamespace( - id="file-1", + upload_file = UploadFile( + tenant_id="tenant-1", + storage_type=StorageType.LOCAL, + key="image.png", name="image.png", size=10, extension="png", mime_type="image/png", + created_by_role=CreatorUserRole.ACCOUNT, + created_by="account-1", + created_at=datetime(2024, 1, 1), + used=False, ) - session = Mock() - session.scalar.return_value = upload_file + upload_file.id = "file-1" + sqlite_session.add_all([dataset_query, upload_file]) + sqlite_session.flush() - with ( - patch("models.dataset.db") as mock_db, - patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"), - ): - queries = dataset_query.get_queries(session=session) + with patch("models.dataset.sign_upload_file_preview_url", return_value="signed-url"): + queries = dataset_query.get_queries(session=sqlite_session) assert queries == [ { @@ -380,8 +481,6 @@ class TestDatasetModelValidation: }, } ] - session.scalar.assert_called_once() - mock_db.session.scalar.assert_not_called() def test_dataset_retrieval_model_dict_property(self): """Test retrieval_model_dict property with default values.""" @@ -479,29 +578,35 @@ class TestDocumentModelRelationships: assert "notion_import" in Document.DATA_SOURCES assert "website_crawl" in Document.DATA_SOURCES - def test_session_aware_document_getters_use_caller_session(self): - document = Document( - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - position=1, - data_source_type=DataSourceType.UPLOAD_FILE, - batch="batch_001", - name="test.pdf", - created_from=DocumentCreatedFrom.WEB, - created_by=str(uuid4()), - dataset_process_rule_id=str(uuid4()), + @pytest.mark.parametrize("sqlite_session", [(Document, DatasetProcessRule, DocumentSegment)], indirect=True) + def test_session_aware_document_getters_use_caller_session(self, sqlite_session: Session): + process_rule = DatasetProcessRule( + dataset_id="dataset-1", + mode=ProcessRuleMode.CUSTOM, + rules=None, + created_by="account-1", ) - process_rule = Mock() - session = Mock() - session.get.return_value = process_rule - session.scalar.side_effect = [3, 7] + document = _make_document(process_rule_id=process_rule.id) + segments = [ + DocumentSegment( + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, + position=position, + content=f"Segment {position}", + word_count=2, + tokens=2, + created_by="account-1", + hit_count=hit_count, + ) + for position, hit_count in [(1, 2), (2, 1), (3, 4)] + ] + sqlite_session.add_all([process_rule, document, *segments]) + sqlite_session.flush() - assert document.get_dataset_process_rule(session=session) is process_rule - assert document.get_segment_count(session=session) == 3 - assert document.get_hit_count(session=session) == 7 - - session.get.assert_called_once_with(DatasetProcessRule, document.dataset_process_rule_id) - assert session.scalar.call_count == 2 + assert document.get_dataset_process_rule(session=sqlite_session) is process_rule + assert document.get_segment_count(session=sqlite_session) == 3 + assert document.get_hit_count(session=sqlite_session) == 7 def test_document_display_status_queuing(self): """Test document display_status property for queuing state.""" @@ -706,25 +811,22 @@ class TestDocumentModelRelationships: # Assert assert result == {} - def test_document_get_dataset_uses_caller_session(self): - document = Document( - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - position=1, - data_source_type=DataSourceType.UPLOAD_FILE, - batch="batch_001", - name="test.pdf", - created_from=DocumentCreatedFrom.WEB, - created_by=str(uuid4()), - ) - dataset = Dataset() - session = Mock() - session.get.return_value = dataset + @pytest.mark.parametrize("sqlite_session", [(Document, Dataset)], indirect=True) + def test_document_get_dataset_uses_caller_session(self, sqlite_session: Session): + dataset = _make_dataset() + document = _make_document(dataset_id=dataset.id) + sqlite_session.add_all([dataset, document]) + sqlite_session.flush() - assert document.get_dataset(session=session) is dataset - session.get.assert_called_once_with(Dataset, document.dataset_id) + assert document.get_dataset(session=sqlite_session) is dataset - def test_document_average_segment_length(self): + @pytest.mark.parametrize("sqlite_session", [(Document, DocumentSegment)], indirect=True) + def test_document_average_segment_length( + self, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + monkeypatch: pytest.MonkeyPatch, + ): """Test average_segment_length property calculation.""" # Arrange document = Document( @@ -738,14 +840,17 @@ class TestDocumentModelRelationships: created_by=str(uuid4()), word_count=1000, ) + sqlite_session.add(document) + sqlite_session.flush() + sqlite_session.add_all(_make_segments(document, [0] * 10)) + sqlite_session.commit() + monkeypatch.setattr(dataset_module.db, "session", scoped_session(sqlite_session_factory)) - # Mock segment_count property - with patch.object(Document, "segment_count", new_callable=lambda: property(lambda self: 10)): - # Act - result = document.average_segment_length + # Act + result = document.average_segment_length - # Assert - assert result == 100 + # Assert + assert result == 100 def test_document_average_segment_length_zero(self): """Test average_segment_length property when word_count is zero.""" @@ -772,107 +877,101 @@ class TestDocumentModelRelationships: class TestDocumentSegmentIndexing: """Test suite for DocumentSegment model indexing and operations.""" - def test_get_child_chunks_uses_caller_session(self): + @pytest.mark.parametrize( + "sqlite_session", [(DocumentSegment, Document, DatasetProcessRule, ChildChunk)], indirect=True + ) + def test_get_child_chunks_uses_caller_session(self, sqlite_session: Session): + process_rule = DatasetProcessRule( + dataset_id="dataset-1", + mode=ProcessRuleMode.HIERARCHICAL, + rules=json.dumps({"parent_mode": ParentMode.PARAGRAPH}), + created_by="account-1", + ) + document = _make_document(process_rule_id=process_rule.id) segment = DocumentSegment( - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - document_id=str(uuid4()), + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, position=1, content="Test content", word_count=2, tokens=5, created_by=str(uuid4()), ) - document = Document() - process_rule = Mock(mode="hierarchical", rules_dict={"parent_mode": "paragraph"}) child_chunk = ChildChunk( - tenant_id="tenant-id", - dataset_id="dataset-id", - document_id="document-id", - segment_id="segment-id", + tenant_id=segment.tenant_id, + dataset_id=segment.dataset_id, + document_id=segment.document_id, + segment_id=segment.id, position=1, content="", word_count=0, created_by="account-id", ) - session = Mock() - session.get.return_value = document - session.scalars.return_value.all.return_value = [child_chunk] - - with ( - patch.object(Document, "get_dataset_process_rule", return_value=process_rule) as get_process_rule, - patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode="paragraph")), - ): - result = segment.get_child_chunks(session=session) + sqlite_session.add_all([process_rule, document, segment, child_chunk]) + sqlite_session.flush() + result = segment.get_child_chunks(session=sqlite_session) assert result == [child_chunk] - session.get.assert_called_once_with(Document, segment.document_id) - get_process_rule.assert_called_once_with(session=session) - session.scalars.assert_called_once() - def test_get_child_chunks_includes_full_doc_unless_explicitly_hidden(self): + @pytest.mark.parametrize( + "sqlite_session", [(DocumentSegment, Document, DatasetProcessRule, ChildChunk)], indirect=True + ) + def test_get_child_chunks_includes_full_doc_unless_explicitly_hidden(self, sqlite_session: Session): + process_rule = DatasetProcessRule( + dataset_id="dataset-1", + mode=ProcessRuleMode.HIERARCHICAL, + rules=json.dumps({"parent_mode": ParentMode.FULL_DOC}), + created_by="account-1", + ) + document = _make_document(process_rule_id=process_rule.id) segment = DocumentSegment( - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - document_id=str(uuid4()), + tenant_id=document.tenant_id, + dataset_id=document.dataset_id, + document_id=document.id, position=1, content="Test content", word_count=2, tokens=5, created_by=str(uuid4()), ) - document = Document() - process_rule = Mock( - mode="hierarchical", - rules_dict={"parent_mode": ParentMode.FULL_DOC}, - ) - session = Mock() - session.get.return_value = document child_chunk = ChildChunk( - tenant_id="tenant-id", - dataset_id="dataset-id", - document_id="document-id", - segment_id="segment-id", + tenant_id=segment.tenant_id, + dataset_id=segment.dataset_id, + document_id=segment.document_id, + segment_id=segment.id, position=1, content="", word_count=0, created_by="account-id", ) - session.scalars.return_value.all.return_value = [child_chunk] - - with ( - patch.object(Document, "get_dataset_process_rule", return_value=process_rule), - patch("models.dataset.Rule.model_validate", return_value=Mock(parent_mode=ParentMode.FULL_DOC)), - ): - result = segment.get_child_chunks(session=session) - response_result = segment.get_child_chunks(session=session, include_full_doc=False) + sqlite_session.add_all([process_rule, document, segment, child_chunk]) + sqlite_session.flush() + result = segment.get_child_chunks(session=sqlite_session) + response_result = segment.get_child_chunks(session=sqlite_session, include_full_doc=False) assert result == [child_chunk] assert response_result == [] - session.scalars.assert_called_once() - def test_relationship_getters_use_caller_session(self): + @pytest.mark.parametrize("sqlite_session", [(Dataset, Document, DocumentSegment)], indirect=True) + def test_relationship_getters_use_caller_session(self, sqlite_session: Session): + dataset = _make_dataset() + document = _make_document(dataset_id=dataset.id) segment = DocumentSegment( - tenant_id=str(uuid4()), - dataset_id=str(uuid4()), - document_id=str(uuid4()), + tenant_id=dataset.tenant_id, + dataset_id=dataset.id, + document_id=document.id, position=1, content="Test content", word_count=2, tokens=5, created_by=str(uuid4()), ) - dataset = Dataset() - document = Document() - session = Mock() - session.get.side_effect = [dataset, document] + sqlite_session.add_all([dataset, document, segment]) + sqlite_session.flush() - assert segment.get_dataset(session=session) is dataset - assert segment.get_document(session=session) is document - assert session.get.call_args_list == [ - call(Dataset, segment.dataset_id), - call(Document, segment.document_id), - ] + assert segment.get_dataset(session=sqlite_session) is dataset + assert segment.get_document(session=sqlite_session) is document def test_document_segment_creation_with_required_fields(self): """Test creating a document segment with all required fields.""" @@ -1029,7 +1128,10 @@ class TestDocumentSegmentIndexing: # Assert assert segment.hit_count == 5 - def test_document_segment_attachments_prefers_files_url_for_source_url(self, monkeypatch: pytest.MonkeyPatch): + @pytest.mark.parametrize("sqlite_session", [(DocumentSegment, UploadFile, SegmentAttachmentBinding)], indirect=True) + def test_document_segment_attachments_prefers_files_url_for_source_url( + self, sqlite_session: Session, monkeypatch: pytest.MonkeyPatch + ): """Test attachment source URLs use FILES_URL before falling back to CONSOLE_API_URL.""" # Arrange segment = DocumentSegment( @@ -1057,6 +1159,15 @@ class TestDocumentSegmentIndexing: used=False, ) attachment.id = "upload-1" + binding = SegmentAttachmentBinding( + tenant_id=segment.tenant_id, + dataset_id=segment.dataset_id, + document_id=segment.document_id, + segment_id=segment.id, + attachment_id=attachment.id, + ) + sqlite_session.add_all([segment, attachment, binding]) + sqlite_session.flush() monkeypatch.setattr("models.dataset.time.time", lambda: 1700000000) monkeypatch.setattr("models.dataset.os.urandom", lambda _: b"\x01" * 16) @@ -1064,11 +1175,8 @@ class TestDocumentSegmentIndexing: monkeypatch.setattr("models.dataset.dify_config.FILES_URL", "https://files.example.com") monkeypatch.setattr("models.dataset.dify_config.CONSOLE_API_URL", "https://console.example.com") - session = Mock() - session.execute.return_value.all.return_value = [(Mock(), attachment)] - # Act - attachments = segment.get_attachments(session=session) + attachments = segment.get_attachments(session=sqlite_session) # Assert assert len(attachments) == 1 @@ -1080,7 +1188,6 @@ class TestDocumentSegmentIndexing: assert query["timestamp"] == ["1700000000"] assert query["nonce"] == ["01010101010101010101010101010101"] assert query["sign"][0] - session.execute.assert_called_once() def test_document_segment_error_tracking(self): """Test document segment error tracking.""" @@ -1308,20 +1415,20 @@ class TestDatasetKeywordTable: # Assert assert keyword_table.data_source_type == "file" - def test_get_keyword_table_dict_from_database_uses_caller_session(self): - dataset = Mock(tenant_id="tenant-1") - session = Mock() - session.scalar.return_value = dataset + @pytest.mark.parametrize("sqlite_session", [(Dataset, DatasetKeywordTable)], indirect=True) + def test_get_keyword_table_dict_from_database_uses_caller_session(self, sqlite_session: Session): + dataset = _make_dataset() keyword_table = DatasetKeywordTable( - dataset_id="dataset-1", + dataset_id=dataset.id, keyword_table=json.dumps({"__data__": {"table": {"keyword": ["node-1"]}}}), data_source_type="database", ) + sqlite_session.add_all([dataset, keyword_table]) + sqlite_session.flush() - result = keyword_table.get_keyword_table_dict(session=session) + result = keyword_table.get_keyword_table_dict(session=sqlite_session) assert result == {"__data__": {"table": {"keyword": {"node-1"}}}} - session.scalar.assert_called_once() class TestAppDatasetJoin: @@ -1462,7 +1569,13 @@ class TestModelIntegration: assert document.word_count == 100 assert segment.status == SegmentStatus.COMPLETED - def test_document_to_dict_serialization(self): + @pytest.mark.parametrize("sqlite_session", [(Document, DocumentSegment)], indirect=True) + def test_document_to_dict_serialization( + self, + sqlite_session: Session, + sqlite_session_factory: sessionmaker[Session], + monkeypatch: pytest.MonkeyPatch, + ): """Test document to_dict method for serialization.""" # Arrange tenant_id = str(uuid4()) @@ -1481,20 +1594,20 @@ class TestModelIntegration: word_count=100, indexing_status=IndexingStatus.COMPLETED, ) + sqlite_session.add(document) + sqlite_session.flush() + sqlite_session.add_all(_make_segments(document, [2, 2, 2, 2, 2])) + sqlite_session.commit() + monkeypatch.setattr(dataset_module.db, "session", scoped_session(sqlite_session_factory)) - # Mock segment_count and hit_count - with ( - patch.object(Document, "segment_count", new_callable=lambda: property(lambda self: 5)), - patch.object(Document, "hit_count", new_callable=lambda: property(lambda self: 10)), - ): - # Act - result = document.to_dict() + # Act + result = document.to_dict() - # Assert - assert result["tenant_id"] == tenant_id - assert result["dataset_id"] == dataset_id - assert result["name"] == "test.pdf" - assert result["word_count"] == 100 - assert result["indexing_status"] == IndexingStatus.COMPLETED - assert result["segment_count"] == 5 - assert result["hit_count"] == 10 + # Assert + assert result["tenant_id"] == tenant_id + assert result["dataset_id"] == dataset_id + assert result["name"] == "test.pdf" + assert result["word_count"] == 100 + assert result["indexing_status"] == IndexingStatus.COMPLETED + assert result["segment_count"] == 5 + assert result["hit_count"] == 10 diff --git a/api/tests/unit_tests/models/test_model.py b/api/tests/unit_tests/models/test_model.py index bb3206713d7..a00c2341bb9 100644 --- a/api/tests/unit_tests/models/test_model.py +++ b/api/tests/unit_tests/models/test_model.py @@ -1,12 +1,12 @@ -import importlib -import types -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest +from sqlalchemy.orm import Session from core.workflow.file_reference import build_file_reference from graphon.file import FILE_MODEL_IDENTITY, FileTransferMethod -from models.model import Conversation, Message +from models import model as model_module +from models.model import App, AppMode, Conversation, IconType, Message @pytest.fixture(autouse=True) @@ -14,10 +14,11 @@ def patch_file_helpers(monkeypatch: pytest.MonkeyPatch): """ Patch file_helpers.get_signed_file_url to a deterministic stub. """ - model_module = importlib.import_module("models.model") - dummy = types.SimpleNamespace(get_signed_file_url=lambda fid: f"https://signed.example/{fid}") - # Inject/override file_helpers on models.model - monkeypatch.setattr(model_module, "file_helpers", dummy, raising=False) + monkeypatch.setattr( + model_module.file_helpers, + "get_signed_file_url", + lambda file_id: f"https://signed.example/{file_id}", + ) def _wrap_md(url: str) -> str: @@ -124,17 +125,43 @@ def test_inputs_restore_external_remote_url_file_mappings(owner_cls: type[Conver assert restored_file.remote_url == "https://example.com/report.pdf" -def test_message_inputs_resolve_file_tenant_with_caller_session() -> None: +@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True) +def test_message_inputs_resolve_file_tenant_with_caller_session(sqlite_session: Session) -> None: + app = App( + id="app-1", + tenant_id="tenant-1", + name="File owner", + description="", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="file", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, + ) + decoy = App( + id="other-app", + tenant_id="other-tenant", + name="Decoy", + description="", + mode=AppMode.CHAT, + icon_type=IconType.EMOJI, + icon="file", + icon_background="#FFFFFF", + enable_site=False, + enable_api=False, + max_active_requests=0, + ) + sqlite_session.add_all([decoy, app]) + sqlite_session.flush() message = Message(app_id="app-1") message.inputs = {"file": _build_local_file_mapping("upload-1")} - session = MagicMock() - session.scalar.return_value = "tenant-1" with patch( "models.model.build_file_from_input_mapping", side_effect=lambda **kwargs: kwargs["tenant_resolver"](), ): - inputs = message.inputs_with_session(session=session) + inputs = message.inputs_with_session(session=sqlite_session) assert inputs["file"] == "tenant-1" - session.scalar.assert_called_once() diff --git a/api/tests/unit_tests/models/test_workflow.py b/api/tests/unit_tests/models/test_workflow.py index 9ec7383e1dd..8c202450ff4 100644 --- a/api/tests/unit_tests/models/test_workflow.py +++ b/api/tests/unit_tests/models/test_workflow.py @@ -3,6 +3,9 @@ import json from unittest import mock from uuid import uuid4 +import pytest +from sqlalchemy.orm import Session + from constants import HIDDEN_VALUE from core.helper import encrypter from core.workflow.file_reference import build_file_reference @@ -12,6 +15,7 @@ from graphon.file import File, FileTransferMethod, FileType from graphon.variables import FloatVariable, IntegerVariable, SecretVariable, StringVariable from graphon.variables.segments import IntegerSegment, Segment from models.account import Account +from models.tools import WorkflowToolProvider from models.workflow import ( Workflow, WorkflowDraftVariable, @@ -162,7 +166,14 @@ def test_to_dict(): assert workflow_dict["environment_variables"][1]["value"] == "text" -def test_workflow_account_getters_use_caller_session(): +@pytest.mark.parametrize("sqlite_session", [(Workflow, Account)], indirect=True) +def test_workflow_account_getters_use_caller_session(sqlite_session: Session): + created_account = Account(name="Created Account", email="created@example.com") + created_account.id = "created-account-id" + updated_account = Account(name="Updated Account", email="updated@example.com") + updated_account.id = "updated-account-id" + decoy_account = Account(name="Decoy Account", email="decoy@example.com") + decoy_account.id = "decoy-account-id" workflow = Workflow( tenant_id="tenant_id", app_id="app_id", @@ -175,23 +186,15 @@ def test_workflow_account_getters_use_caller_session(): conversation_variables=[], updated_by="updated-account-id", ) - created_account = Account(name="Test Account", email="test@example.com") - updated_account = Account(name="Test Account", email="test@example.com") - session = mock.Mock() - session.get.side_effect = [created_account, updated_account] + sqlite_session.add_all([decoy_account, updated_account, workflow, created_account]) + sqlite_session.flush() - with mock.patch("models.workflow.db") as mock_db: - assert workflow.get_created_by_account(session=session) is created_account - assert workflow.get_updated_by_account(session=session) is updated_account - - assert session.get.call_args_list == [ - mock.call(Account, "created-account-id"), - mock.call(Account, "updated-account-id"), - ] - mock_db.session.get.assert_not_called() + assert workflow.get_created_by_account(session=sqlite_session) is created_account + assert workflow.get_updated_by_account(session=sqlite_session) is updated_account -def test_workflow_tool_published_getter_uses_caller_session(): +@pytest.mark.parametrize("sqlite_session", [(Workflow, WorkflowToolProvider)], indirect=True) +def test_workflow_tool_published_getter_uses_caller_session(sqlite_session: Session): workflow = Workflow( tenant_id="tenant_id", app_id="app_id", @@ -203,14 +206,30 @@ def test_workflow_tool_published_getter_uses_caller_session(): environment_variables=[], conversation_variables=[], ) - session = mock.Mock() - session.execute.return_value.scalar_one.return_value = True + matching_provider = WorkflowToolProvider( + name="matching-provider", + label="Matching provider", + icon="tool", + app_id=workflow.app_id, + version="1", + user_id="account-id", + tenant_id=workflow.tenant_id, + description="Matching workflow tool", + ) + decoy_provider = WorkflowToolProvider( + name="decoy-provider", + label="Decoy provider", + icon="tool", + app_id="other-app", + version="1", + user_id="account-id", + tenant_id=workflow.tenant_id, + description="Different app", + ) + sqlite_session.add_all([decoy_provider, workflow, matching_provider]) + sqlite_session.flush() - with mock.patch("models.workflow.db") as mock_db: - assert workflow.get_tool_published(session=session) is True - - session.execute.assert_called_once() - mock_db.session.execute.assert_not_called() + assert workflow.get_tool_published(session=sqlite_session) is True def test_normalize_environment_variable_mappings_converts_full_mask_to_hidden_value():