diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py index 753236f6dfc..4bdc862b651 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_datasource_auth.py @@ -27,12 +27,19 @@ from controllers.console.datasets.rag_pipeline.datasource_auth import ( ) from core.plugin.impl.oauth import OAuthHandler from graphon.model_runtime.errors.validate import CredentialsValidateFailedError +from models.account import Account from services.datasource_provider_service import DatasourceProviderService from services.plugin.oauth_service import OAuthProxyService _PROVIDER_ID = "langgenius/notion_datasource/notion" +def _account() -> Account: + account = Account(name="Datasource Auth Tester", email="datasource-auth@example.com") + account.id = "user-1" + return account + + def _i18n(text: str) -> dict[str, str]: return {"en_US": text, "zh_Hans": text, "pt_BR": text, "ja_JP": text} @@ -106,7 +113,7 @@ class TestDatasourcePluginOAuthAuthorizationUrl: api = DatasourcePluginOAuthAuthorizationUrl() method = inspect.unwrap(api.get) - user = MagicMock(id="user-1") + user = _account() oauth_client = {"client_id": "abc", "client_secret": "shh", "scopes": ["read", "write"]} auth_url_payload = { "authorization_url": "https://auth.example.com/oauth?client_id=abc&state=xyz", @@ -155,7 +162,7 @@ class TestDatasourcePluginOAuthAuthorizationUrl: def test_get_no_oauth_config(self, app: Flask): api = DatasourcePluginOAuthAuthorizationUrl() method = inspect.unwrap(api.get) - user = MagicMock(id="user-1") + user = _account() with ( app.test_request_context("/"), @@ -172,7 +179,7 @@ class TestDatasourcePluginOAuthAuthorizationUrl: api = DatasourcePluginOAuthAuthorizationUrl() method = inspect.unwrap(api.get) - user = MagicMock(id="user-1") + user = _account() with ( app.test_request_context("/"), @@ -443,7 +450,7 @@ class TestDatasourceAuth: def test_get_success(self, app: Flask): api = DatasourceAuth() method = inspect.unwrap(api.get) - user = MagicMock(id="user-1") + user = _account() with ( app.test_request_context("/"), @@ -474,7 +481,7 @@ class TestDatasourceAuth: def test_get_empty_list(self, app: Flask): api = DatasourceAuth() method = inspect.unwrap(api.get) - user = MagicMock(id="user-1") + user = _account() with ( app.test_request_context("/"), diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py index 3f8e1a2756c..3e2c9cefd04 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_datasets.py @@ -17,9 +17,16 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_datasets import ( CreateRagPipelineDatasetApi, RagPipelineDatasetImportPayload, ) +from models.account import Account, TenantAccountRole from services.entities.dsl_entities import ImportStatus +def _account(*, editor: bool) -> Account: + account = Account(name="RAG Pipeline Tester", email="rag-pipeline@example.com") + account.role = TenantAccountRole.EDITOR if editor else TenantAccountRole.NORMAL + return account + + class TestCreateRagPipelineDatasetApi: def _valid_payload(self) -> dict[str, str]: return {"yaml_content": "name: test"} @@ -29,7 +36,7 @@ class TestCreateRagPipelineDatasetApi: method = unwrap(api.post) payload = self._valid_payload() - user = MagicMock(is_dataset_editor=True) + user = _account(editor=True) import_info = { "id": "import-1", "status": ImportStatus.COMPLETED, @@ -69,7 +76,7 @@ class TestCreateRagPipelineDatasetApi: method = unwrap(api.post) payload = self._valid_payload() - user = MagicMock(is_dataset_editor=False) + user = _account(editor=False) with ( app.test_request_context("/", json=payload), @@ -83,7 +90,7 @@ class TestCreateRagPipelineDatasetApi: method = unwrap(api.post) payload = self._valid_payload() - user = MagicMock(is_dataset_editor=True) + user = _account(editor=True) mock_service = MagicMock() mock_service.create_rag_pipeline_dataset.side_effect = services.errors.dataset.DatasetNameDuplicateError() @@ -104,7 +111,7 @@ class TestCreateRagPipelineDatasetApi: method = unwrap(api.post) payload: dict[str, str] = {} - user = MagicMock(is_dataset_editor=True) + user = _account(editor=True) with ( app.test_request_context("/", json=payload), @@ -119,7 +126,7 @@ class TestCreateEmptyRagPipelineDatasetApi: api = CreateEmptyRagPipelineDatasetApi() method = unwrap(api.post) - user = MagicMock(is_dataset_editor=False) + user = _account(editor=False) with app.test_request_context("/"): with pytest.raises(Forbidden): diff --git a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py index ad6d974bdb8..c2e80c24bd5 100644 --- a/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py +++ b/api/tests/unit_tests/controllers/console/datasets/rag_pipeline/test_rag_pipeline_import.py @@ -19,12 +19,19 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_import import ( RagPipelineImportPayload, ) from core.plugin.entities.plugin import PluginDependency, PluginDependencyType +from models.account import Account from models.dataset import Pipeline from models.engine import db from services.entities.dsl_entities import CheckDependenciesResult, ImportStatus from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelineImportInfo +def _account() -> Account: + account = Account(name="RAG Import Tester", email="rag-import@example.com") + account.id = "account-1" + return account + + @pytest.fixture def app() -> Iterator[Flask]: app = Flask(__name__) @@ -48,7 +55,7 @@ class TestRagPipelineImportApi: method = unwrap(api.post) payload = self._payload() - user = MagicMock() + user = _account() result = RagPipelineImportInfo( id="import-1", status=ImportStatus.COMPLETED, @@ -87,7 +94,7 @@ class TestRagPipelineImportApi: method = unwrap(api.post) payload = self._payload() - user = MagicMock() + user = _account() result = RagPipelineImportInfo( id="import-1", status=ImportStatus.FAILED, @@ -120,7 +127,7 @@ class TestRagPipelineImportApi: method = unwrap(api.post) payload = self._payload() - user = MagicMock() + user = _account() result = RagPipelineImportInfo( id="import-1", status=ImportStatus.PENDING, @@ -154,7 +161,7 @@ class TestRagPipelineImportConfirmApi: api = RagPipelineImportConfirmApi() method = unwrap(api.post) - user = MagicMock() + user = _account() result = RagPipelineImportInfo( id="import-1", status=ImportStatus.COMPLETED, @@ -184,7 +191,7 @@ class TestRagPipelineImportConfirmApi: api = RagPipelineImportConfirmApi() method = unwrap(api.post) - user = MagicMock() + user = _account() result = RagPipelineImportInfo( id="import-1", status=ImportStatus.FAILED, diff --git a/api/tests/unit_tests/controllers/console/workspace/test_agent_providers.py b/api/tests/unit_tests/controllers/console/workspace/test_agent_providers.py index f7310b5fec8..f62a0818333 100644 --- a/api/tests/unit_tests/controllers/console/workspace/test_agent_providers.py +++ b/api/tests/unit_tests/controllers/console/workspace/test_agent_providers.py @@ -1,5 +1,5 @@ from inspect import unwrap -from unittest.mock import MagicMock, patch +from unittest.mock import patch from flask import Flask @@ -7,6 +7,13 @@ from controllers.console.workspace.agent_providers import ( AgentProviderApi, AgentProviderListApi, ) +from models.account import Account + + +def _account() -> Account: + account = Account(name="Agent Provider Tester", email="agent-provider@example.com") + account.id = "user1" + return account class TestAgentProviderListApi: @@ -14,7 +21,7 @@ class TestAgentProviderListApi: api = AgentProviderListApi() method = unwrap(api.get) - user = MagicMock(id="user1") + user = _account() tenant_id = "tenant1" providers = [{"name": "openai"}, {"name": "anthropic"}] @@ -33,7 +40,7 @@ class TestAgentProviderListApi: api = AgentProviderListApi() method = unwrap(api.get) - user = MagicMock(id="user1") + user = _account() tenant_id = "tenant1" with ( @@ -53,7 +60,7 @@ class TestAgentProviderApi: api = AgentProviderApi() method = unwrap(api.get) - user = MagicMock(id="user1") + user = _account() tenant_id = "tenant1" provider_name = "openai" provider_data = {"name": "openai", "models": ["gpt-4"]} @@ -73,7 +80,7 @@ class TestAgentProviderApi: api = AgentProviderApi() method = unwrap(api.get) - user = MagicMock(id="user1") + user = _account() tenant_id = "tenant1" provider_name = "unknown" diff --git a/api/tests/unit_tests/controllers/service_api/conftest.py b/api/tests/unit_tests/controllers/service_api/conftest.py index 126ac8e8763..38bae518d56 100644 --- a/api/tests/unit_tests/controllers/service_api/conftest.py +++ b/api/tests/unit_tests/controllers/service_api/conftest.py @@ -9,7 +9,6 @@ Service API controller tests. import uuid from collections.abc import Iterator from dataclasses import dataclass -from unittest.mock import Mock import pytest from flask import Flask @@ -19,7 +18,7 @@ from sqlalchemy.orm import Session from core.rag.index_processor.constant.index_type import IndexStructureType from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus from models.base import TypeBase -from models.model import App, AppMode, EndUser +from models.model import ApiToken, App, AppMode, EndUser, EndUserType @dataclass(frozen=True) @@ -82,11 +81,15 @@ def mock_app_id(): @pytest.fixture def mock_end_user(mock_tenant_id): - """Create a mock EndUser model with required attributes.""" + """Create a real EndUser model with required attributes.""" user = EndUser( id=str(uuid.uuid4()), external_user_id=f"external_{uuid.uuid4().hex[:8]}", tenant_id=mock_tenant_id, + app_id=None, + type=EndUserType.SERVICE_API, + name="Service API User", + session_id=str(uuid.uuid4()), ) return user @@ -103,52 +106,44 @@ def mock_app_model(mock_app_id, mock_tenant_id): status="normal", enable_api=True, ) - app.author_name = "Test Author" - app.tags = [] - - # Mock workflow for workflow apps - app.workflow = None - app.app_model_config = None - return app @pytest.fixture def mock_tenant(mock_tenant_id): """Create a Tenant model.""" - tenant = Mock() + tenant = Tenant(name="Service API Tenant", status=TenantStatus.NORMAL) tenant.id = mock_tenant_id - tenant.status = TenantStatus.NORMAL return tenant @pytest.fixture def mock_account(): """Create an Account model.""" - account = Mock() + account = Account(name="Service API Account", email=f"service-{uuid.uuid4()}@example.com") account.id = str(uuid.uuid4()) return account @pytest.fixture def mock_api_token(mock_app_id, mock_tenant_id): - """Create a mock API token for authentication tests.""" - token = Mock() - token.app_id = mock_app_id - token.tenant_id = mock_tenant_id - token.token = f"test_token_{uuid.uuid4().hex[:8]}" - token.type = "app" - return token + """Create a real API token for authentication tests.""" + return ApiToken( + app_id=mock_app_id, + tenant_id=mock_tenant_id, + token=f"test_token_{uuid.uuid4().hex[:8]}", + type="app", + ) @pytest.fixture def mock_dataset_api_token(mock_tenant_id): - """Create a mock API token for dataset endpoints.""" - token = Mock() - token.tenant_id = mock_tenant_id - token.token = f"dataset_token_{uuid.uuid4().hex[:8]}" - token.type = "dataset" - return token + """Create a real API token for dataset endpoints.""" + return ApiToken( + tenant_id=mock_tenant_id, + token=f"dataset_token_{uuid.uuid4().hex[:8]}", + type="dataset", + ) @pytest.fixture diff --git a/api/tests/unit_tests/controllers/web/test_message_list.py b/api/tests/unit_tests/controllers/web/test_message_list.py index 4a1cf603f69..3f45840a73a 100644 --- a/api/tests/unit_tests/controllers/web/test_message_list.py +++ b/api/tests/unit_tests/controllers/web/test_message_list.py @@ -4,8 +4,10 @@ from __future__ import annotations import builtins import inspect +import json import uuid from datetime import datetime +from pathlib import Path from types import ModuleType, SimpleNamespace from unittest.mock import ANY, patch from uuid import uuid4 @@ -15,7 +17,18 @@ from flask import Flask from flask.views import MethodView from controllers.common.controller_schemas import MessageListQuery +from core.app.entities.app_invoke_entities import InvokeFrom from core.entities.execution_extra_content import HumanInputContent +from models.enums import ConversationFromSource, EndUserType, FeedbackFromSource, FeedbackRating +from models.model import ( + App, + AppMode, + Conversation, + EndUser, + Message, + MessageAgentThought, + MessageFeedback, +) # Ensure flask_restx.api finds MethodView during import. if not hasattr(builtins, "MethodView"): @@ -36,8 +49,9 @@ def _load_controller_module(): from flask_restx import Namespace stub = ModuleType(parent_module_name) - stub.__file__ = "controllers/web/__init__.py" - stub.__path__ = ["controllers/web"] + web_controller_dir = Path(__file__).resolve().parents[4] / "controllers" / "web" + stub.__file__ = str(web_controller_dir / "__init__.py") + stub.__path__ = [str(web_controller_dir)] stub.__package__ = "controllers" stub.__spec__ = importlib.util.spec_from_loader(parent_module_name, loader=None, is_package=True) stub.web_ns = Namespace("web", description="Web API", path="/") @@ -67,7 +81,12 @@ def app() -> Flask: return app -def test_message_list_mapping(app: Flask) -> None: +@pytest.mark.parametrize( + "sqlite_session", + [(App, Conversation, EndUser, Message, MessageAgentThought, MessageFeedback)], + indirect=True, +) +def test_message_list_mapping(app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_session) -> None: conversation_id = str(uuid4()) message_id = str(uuid4()) @@ -75,87 +94,135 @@ def test_message_list_mapping(app: Flask) -> None: resource_created_at = datetime(2024, 1, 1, 13, 0, 0) thought_created_at = datetime(2024, 1, 1, 14, 0, 0) - retriever_resource_obj = SimpleNamespace( - id="res-obj", - message_id=message_id, - position=2, - dataset_id="ds-1", - dataset_name="dataset", - document_id="doc-1", - document_name="document", - data_source_type="file", - segment_id="seg-1", - score=0.9, - hit_count=1, - word_count=10, - segment_position=0, - index_node_hash="hash", - content="content", - created_at=resource_created_at, - ) + retriever_resource = { + "id": "res-obj", + "message_id": message_id, + "position": 2, + "dataset_id": "ds-1", + "dataset_name": "dataset", + "document_id": "doc-1", + "document_name": "document", + "data_source_type": "file", + "segment_id": "seg-1", + "score": 0.9, + "hit_count": 1, + "word_count": 10, + "segment_position": 0, + "index_node_hash": "hash", + "content": "content", + "created_at": int(resource_created_at.timestamp()), + } - agent_thought = SimpleNamespace( - id="thought-1", - chain_id=None, + agent_thought = MessageAgentThought( message_chain_id="chain-1", message_id=message_id, position=1, + created_by_role="end_user", + created_by="end-user-1", thought="thinking", tool="tool", - tool_labels={"label": "value"}, + tool_labels_str=json.dumps({"label": "value"}), tool_input="{}", - created_at=thought_created_at, observation="observed", - files=["file-a"], + message_files=json.dumps(["file-a"]), ) + agent_thought.id = "thought-1" + agent_thought.created_at = thought_created_at - message_file_obj = SimpleNamespace( - id="file-obj", - filename="b.txt", - type="file", - url=None, - mime_type=None, - size=None, - transfer_method="local", - belongs_to=None, - upload_file_id=None, + message_files = [ + {"id": "file-dict", "filename": "a.txt", "type": "file", "transfer_method": "local"}, + {"id": "file-obj", "filename": "b.txt", "type": "file", "transfer_method": "local"}, + ] + + app_model = App( + id="app-1", + tenant_id="tenant-1", + name="Chat App", + mode=AppMode.CHAT, + enable_site=False, + enable_api=False, ) - - message = SimpleNamespace( + end_user = EndUser( + id="end-user-1", + tenant_id="tenant-1", + app_id=app_model.id, + type=EndUserType.BROWSER, + name="Web User", + session_id="session-1", + ) + conversation = Conversation( + id=conversation_id, + app_id=app_model.id, + mode=AppMode.CHAT, + name="Conversation", + _inputs={}, + status="normal", + from_source=ConversationFromSource.API, + from_end_user_id=end_user.id, + ) + message = Message( id=message_id, + app_id=app_model.id, conversation_id=conversation_id, parent_message_id=None, - inputs={"foo": "bar"}, + _inputs={"foo": "bar"}, query="hello", - re_sign_file_url_answer="answer", - user_feedback=SimpleNamespace(rating="like"), - retriever_resources=[ - {"id": "res-dict", "message_id": message_id, "position": 1}, - retriever_resource_obj, - ], + message={}, + answer="answer", + message_unit_price=0, + message_price_unit=0, + answer_unit_price=0, + answer_price_unit=0, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from=InvokeFrom.SERVICE_API, + from_source=ConversationFromSource.API, + from_end_user_id=end_user.id, + app_mode=AppMode.CHAT, + message_metadata=json.dumps( + { + "meta": "value", + "retriever_resources": [ + {"id": "res-dict", "message_id": message_id, "position": 1}, + retriever_resource, + ], + } + ), created_at=created_at, - agent_thoughts=[agent_thought], - message_files=[ - {"id": "file-dict", "filename": "a.txt", "type": "file", "transfer_method": "local"}, - message_file_obj, - ], status="normal", error=None, - message_metadata_dict={"meta": "value"}, - extra_contents=[ + ) + message.set_extra_contents( + [ HumanInputContent( workflow_run_id=str(uuid.uuid4()), submitted=True, - ) - ], + ).model_dump(mode="json") + ] ) + feedback = MessageFeedback( + app_id=app_model.id, + conversation_id=conversation_id, + message_id=message_id, + rating=FeedbackRating.LIKE, + from_source=FeedbackFromSource.USER, + from_end_user_id=end_user.id, + ) + sqlite_session.add_all([app_model, end_user, conversation, message, feedback, agent_thought]) + sqlite_session.commit() pagination = SimpleNamespace(limit=20, has_more=False, data=[message]) - app_model = SimpleNamespace(mode="chat") - end_user = SimpleNamespace() + + def message_files_with_session(_self, *, session): + del session + return message_files + + monkeypatch.setattr(Message, "message_files_with_session", message_files_with_session) with ( patch.object(message_module.MessageService, "pagination_by_first_id", return_value=pagination) as mock_page, + patch.object(message_module.db, "session", return_value=sqlite_session), app.test_request_context(f"/messages?conversation_id={conversation_id}&limit=20"), ): query = MessageListQuery.model_validate({"conversation_id": conversation_id, "limit": 20}) @@ -172,7 +239,7 @@ def test_message_list_mapping(app: Flask) -> None: assert item["inputs"] == {"foo": "bar"} assert item["answer"] == "answer" assert item["feedback"]["rating"] == "like" - assert item["metadata"] == {"meta": "value"} + assert item["metadata"]["meta"] == "value" assert item["created_at"] == int(created_at.timestamp()) assert item["retriever_resources"][0]["id"] == "res-dict" @@ -181,8 +248,8 @@ def test_message_list_mapping(app: Flask) -> None: assert item["agent_thoughts"][0]["chain_id"] == "chain-1" assert item["agent_thoughts"][0]["created_at"] == int(thought_created_at.timestamp()) - assert item["extra_contents"][0]["workflow_run_id"] == message.extra_contents[0].workflow_run_id - assert item["extra_contents"][0]["submitted"] == message.extra_contents[0].submitted + assert item["extra_contents"][0]["workflow_run_id"] == message.extra_contents[0]["workflow_run_id"] + assert item["extra_contents"][0]["submitted"] == message.extra_contents[0]["submitted"] assert item["message_files"][0]["id"] == "file-dict" assert item["message_files"][1]["id"] == "file-obj"