diff --git a/api/tests/unit_tests/core/entities/test_entities_mcp_provider.py b/api/tests/unit_tests/core/entities/test_entities_mcp_provider.py index c0943d0f2f2..4ebf265193d 100644 --- a/api/tests/unit_tests/core/entities/test_entities_mcp_provider.py +++ b/api/tests/unit_tests/core/entities/test_entities_mcp_provider.py @@ -1,5 +1,4 @@ from datetime import UTC, datetime -from types import SimpleNamespace from unittest.mock import Mock, patch import pytest @@ -11,6 +10,7 @@ from core.entities.mcp_provider import ( MCPProviderEntity, ) from core.mcp.types import OAuthTokens +from models.tools import MCPToolProvider def _build_mcp_provider_entity() -> MCPProviderEntity: @@ -37,24 +37,25 @@ def _build_mcp_provider_entity() -> MCPProviderEntity: def test_from_db_model_maps_fields() -> None: # Arrange now = datetime(2025, 1, 1, tzinfo=UTC) - db_provider = SimpleNamespace( - id="provider-1", + db_provider = MCPToolProvider( server_identifier="server-1", name="Example MCP", tenant_id="tenant-1", user_id="user-1", server_url="encrypted-server-url", - headers={"Authorization": "enc"}, + server_url_hash="server-url-hash", + encrypted_headers='{"Authorization": "enc"}', timeout=15, sse_read_timeout=120, authed=True, - credentials={"access_token": "enc-token"}, - tool_dict=[{"name": "search"}], + encrypted_credentials='{"access_token": "enc-token"}', + tools='[{"name": "search"}]', icon=None, - created_at=now, - updated_at=now, identity_mode="off", ) + db_provider.id = "provider-1" + db_provider.created_at = now + db_provider.updated_at = now # Act entity = MCPProviderEntity.from_db_model(db_provider) diff --git a/api/tests/unit_tests/core/memory/test_token_buffer_memory.py b/api/tests/unit_tests/core/memory/test_token_buffer_memory.py index db5498bddb4..9c4ddfd3820 100644 --- a/api/tests/unit_tests/core/memory/test_token_buffer_memory.py +++ b/api/tests/unit_tests/core/memory/test_token_buffer_memory.py @@ -11,6 +11,7 @@ import pytest from sqlalchemy import Engine, event from sqlalchemy.orm import Session +import models.model as model_module from core.memory import token_buffer_memory as memory_module from core.memory.token_buffer_memory import TokenBufferMemory from graphon.file import FileTransferMethod, FileType @@ -23,8 +24,14 @@ from graphon.model_runtime.entities import ( ) from models.base import TypeBase from models.enums import ConversationFromSource, CreatorUserRole, MessageFileBelongsTo -from models.model import AppMode, Message, MessageFile -from models.workflow import Workflow, WorkflowType +from models.model import App, AppMode, Conversation, Message, MessageFile +from models.workflow import ( + Workflow, + WorkflowExecutionStatus, + WorkflowRun, + WorkflowRunTriggeredFrom, + WorkflowType, +) # --------------------------------------------------------------------------- # Helpers / shared fixtures @@ -44,7 +51,7 @@ class Database: def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator[Database]: TypeBase.metadata.create_all( sqlite_engine, - tables=[Message.__table__, MessageFile.__table__, Workflow.__table__], + tables=[App.__table__, Conversation.__table__, Message.__table__, MessageFile.__table__, Workflow.__table__], ) statements: list[tuple[str, object]] = [] @@ -55,17 +62,42 @@ def database(sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch) -> Iterator with Session(sqlite_engine, expire_on_commit=False) as session: database = Database(engine=sqlite_engine, session=session, statements=statements) monkeypatch.setattr(memory_module, "db", database) + monkeypatch.setattr(model_module, "db", database) yield database event.remove(sqlite_engine, "before_cursor_execute", record_statement) -def _make_conversation(mode: AppMode = AppMode.CHAT) -> MagicMock: - """Return a minimal Conversation mock.""" - conv = MagicMock() - conv.id = str(uuid4()) - conv.mode = mode - conv.model_config = {} - return conv +def _make_app(*, app_id: str | None = None, mode: AppMode = AppMode.CHAT) -> App: + """Return a real transient app with the ownership fields used by memory.""" + return App( + id=app_id or str(uuid4()), + tenant_id=str(uuid4()), + name="Memory test app", + mode=mode, + enable_site=False, + enable_api=False, + ) + + +def _make_conversation(mode: AppMode = AppMode.CHAT, *, app_id: str | None = None) -> Conversation: + """Return a real transient conversation configured without database-backed model settings.""" + return Conversation( + id=str(uuid4()), + app_id=app_id or str(uuid4()), + mode=mode, + name="Memory test conversation", + override_model_configs="{}", + _inputs={}, + from_source=ConversationFromSource.API, + ) + + +def _persist_conversation(database: Database, mode: AppMode = AppMode.CHAT) -> Conversation: + app = _make_app(mode=mode) + conversation = _make_conversation(mode, app_id=app.id) + database.session.add_all([app, conversation]) + database.session.commit() + return conversation def _make_model_instance() -> MagicMock: @@ -75,15 +107,51 @@ def _make_model_instance() -> MagicMock: return mi -def _make_message(answer: str = "hello", answer_tokens: int = 5) -> MagicMock: - msg = MagicMock() - msg.id = str(uuid4()) - msg.query = "user query" - msg.answer = answer - msg.answer_tokens = answer_tokens - msg.workflow_run_id = str(uuid4()) - msg.created_at = MagicMock() - return msg +def _make_message(answer: str = "hello", answer_tokens: int = 5) -> Message: + return Message( + id=str(uuid4()), + app_id=str(uuid4()), + conversation_id=str(uuid4()), + _inputs={}, + query="user query", + message={}, + message_unit_price=Decimal(0), + answer=answer, + answer_tokens=answer_tokens, + answer_unit_price=Decimal(0), + currency="USD", + from_source=ConversationFromSource.API, + workflow_run_id=str(uuid4()), + created_at=datetime.now(UTC).replace(tzinfo=None), + ) + + +def _make_message_file() -> MessageFile: + return MessageFile( + message_id=str(uuid4()), + type=FileType.IMAGE, + transfer_method=FileTransferMethod.REMOTE_URL, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + belongs_to=MessageFileBelongsTo.USER, + url="https://example.com/image.png", + ) + + +def _make_workflow_run(*, workflow_id: str | None = None) -> WorkflowRun: + return WorkflowRun( + tenant_id=str(uuid4()), + app_id=str(uuid4()), + workflow_id=workflow_id or str(uuid4()), + type=WorkflowType.CHAT, + triggered_from=WorkflowRunTriggeredFrom.APP_RUN, + version="1", + graph="{}", + inputs="{}", + status=WorkflowExecutionStatus.SUCCEEDED, + created_by_role=CreatorUserRole.ACCOUNT, + created_by=str(uuid4()), + ) def _persist_message( @@ -229,7 +297,7 @@ class TestBuildPromptMessageWithFiles: message_files=[], text_content="hello", message=_make_message(), - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @@ -272,9 +340,8 @@ class TestBuildPromptMessageWithFiles: url="http://example.com/img.png", format="png", mime_type="image/png" ) - mock_message_file = MagicMock() - mock_app_record = MagicMock() - mock_app_record.tenant_id = "tenant-1" + message_file = _make_message_file() + app_record = _make_app() with ( patch( @@ -291,10 +358,10 @@ class TestBuildPromptMessageWithFiles: ), ): result = mem._build_prompt_message_with_files( - message_files=[mock_message_file], + message_files=[message_file], text_content="user text", message=_make_message(), - app_record=mock_app_record, + app_record=app_record, is_user_message=True, ) @@ -318,8 +385,7 @@ class TestBuildPromptMessageWithFiles: real_image_content = ImagePromptMessageContent( url="http://example.com/img.png", format="png", mime_type="image/png" ) - mock_app_record = MagicMock() - mock_app_record.tenant_id = "tenant-1" + app_record = _make_app() with ( patch( @@ -336,10 +402,10 @@ class TestBuildPromptMessageWithFiles: ), ): mem._build_prompt_message_with_files( - message_files=[MagicMock()], + message_files=[_make_message_file()], text_content="user text", message=_make_message(), - app_record=mock_app_record, + app_record=app_record, is_user_message=True, ) @@ -359,8 +425,7 @@ class TestBuildPromptMessageWithFiles: real_image_content = ImagePromptMessageContent( url="http://example.com/img.png", format="png", mime_type="image/png" ) - mock_app_record = MagicMock() - mock_app_record.tenant_id = "tenant-1" + app_record = _make_app() with ( patch( @@ -377,10 +442,10 @@ class TestBuildPromptMessageWithFiles: ), ): result = mem._build_prompt_message_with_files( - message_files=[MagicMock()], + message_files=[_make_message_file()], text_content="ai text", message=_make_message(), - app_record=mock_app_record, + app_record=app_record, is_user_message=False, ) @@ -399,8 +464,7 @@ class TestBuildPromptMessageWithFiles: mock_file_extra_config = MagicMock() mock_file_extra_config.image_config = mock_image_config - mock_app_record = MagicMock() - mock_app_record.tenant_id = "tenant-1" + app_record = _make_app() real_image_content = ImagePromptMessageContent( url="http://example.com/img.png", format="png", mime_type="image/png" @@ -421,10 +485,10 @@ class TestBuildPromptMessageWithFiles: ) as mock_to_prompt, ): mem._build_prompt_message_with_files( - message_files=[MagicMock()], + message_files=[_make_message_file()], text_content="user text", message=_make_message(), - app_record=mock_app_record, + app_record=app_record, is_user_message=True, ) # Ensure the LOW detail was passed through @@ -445,7 +509,7 @@ class TestBuildPromptMessageWithFiles: return_value=mock_file_extra_config, ): result = mem._build_prompt_message_with_files( - message_files=[MagicMock()], + message_files=[_make_message_file()], text_content="hello", message=_make_message(), app_record=None, # <-- forces the else branch → file_objs = [] @@ -460,10 +524,9 @@ class TestBuildPromptMessageWithFiles: # ------------------------------------------------------------------ @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def test_workflow_mode_no_app_raises(self, mode): + def test_workflow_mode_no_app_raises(self, mode, database: Database): """Raises ValueError when conversation.app is falsy.""" conv = _make_conversation(mode) - conv.app = None mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) with pytest.raises(ValueError, match="App not found for conversation"): @@ -471,15 +534,14 @@ class TestBuildPromptMessageWithFiles: message_files=[], text_content="text", message=_make_message(), - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def test_workflow_mode_no_workflow_run_id_raises(self, mode): + def test_workflow_mode_no_workflow_run_id_raises(self, mode, database: Database): """Raises ValueError when message.workflow_run_id is falsy.""" - conv = _make_conversation(mode) - conv.app = MagicMock() + conv = _persist_conversation(database, mode) message = _make_message() message.workflow_run_id = None # force missing @@ -491,16 +553,14 @@ class TestBuildPromptMessageWithFiles: message_files=[], text_content="text", message=message, - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) - def test_workflow_mode_workflow_run_not_found_raises(self, mode): + def test_workflow_mode_workflow_run_not_found_raises(self, mode, database: Database): """Raises ValueError when workflow_run_repo returns None.""" - conv = _make_conversation(mode) - mock_app = MagicMock() - conv.app = mock_app + conv = _persist_conversation(database, mode) mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) mem._workflow_run_repo = MagicMock() @@ -511,46 +571,39 @@ class TestBuildPromptMessageWithFiles: message_files=[], text_content="text", message=_make_message(), - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) def test_workflow_mode_workflow_not_found_raises(self, mode, database: Database): """Raises ValueError when Workflow lookup returns None.""" - conv = _make_conversation(mode) - conv.app = MagicMock() - - mock_workflow_run = MagicMock() - mock_workflow_run.workflow_id = str(uuid4()) + conv = _persist_conversation(database, mode) + workflow_run = _make_workflow_run() mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) mem._workflow_run_repo = MagicMock() - mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run + mem._workflow_run_repo.get_workflow_run_by_id.return_value = workflow_run with pytest.raises(ValueError, match="Workflow not found"): mem._build_prompt_message_with_files( message_files=[], text_content="text", message=_make_message(), - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @pytest.mark.parametrize("mode", [AppMode.ADVANCED_CHAT, AppMode.WORKFLOW]) def test_workflow_mode_success_no_files_user(self, mode, database: Database): """Happy path: workflow mode, no message files → plain UserPromptMessage.""" - conv = _make_conversation(mode) - conv.app = MagicMock() - - mock_workflow_run = MagicMock() - mock_workflow_run.workflow_id = str(uuid4()) - - workflow = _persist_workflow(database, workflow_id=mock_workflow_run.workflow_id) + conv = _persist_conversation(database, mode) + workflow_run = _make_workflow_run() + workflow = _persist_workflow(database, workflow_id=workflow_run.workflow_id) mem = TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) mem._workflow_run_repo = MagicMock() - mem._workflow_run_repo.get_workflow_run_by_id.return_value = mock_workflow_run + mem._workflow_run_repo.get_workflow_run_by_id.return_value = workflow_run with patch( "core.memory.token_buffer_memory.FileUploadConfigManager.convert", @@ -560,7 +613,7 @@ class TestBuildPromptMessageWithFiles: message_files=[], text_content="wf text", message=_make_message(), - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @@ -583,7 +636,7 @@ class TestBuildPromptMessageWithFiles: message_files=[], text_content="text", message=_make_message(), - app_record=MagicMock(), + app_record=_make_app(), is_user_message=True, ) @@ -596,23 +649,22 @@ class TestBuildPromptMessageWithFiles: class TestGetHistoryPromptMessages: """Tests for persisted history retrieval, file batching, and pruning.""" - def _make_memory(self, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory: - conv = _make_conversation(mode) - conv.app = MagicMock() + def _make_memory(self, database: Database, mode: AppMode = AppMode.CHAT) -> TokenBufferMemory: + conv = _persist_conversation(database, mode) return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) def test_returns_empty_when_no_messages(self, database: Database) -> None: - assert self._make_memory().get_history_prompt_messages() == [] + assert self._make_memory(database).get_history_prompt_messages() == [] def test_skips_newest_message_without_answer(self, database: Database) -> None: - mem = self._make_memory() + mem = self._make_memory(database) message = _persist_message(database, mem.conversation.id, answer="", answer_tokens=0) assert mem.get_history_prompt_messages() == [] assert database.session.get(Message, message.id) is message def test_message_with_answer_returns_user_and_assistant_prompts(self, database: Database) -> None: - mem = self._make_memory() + mem = self._make_memory(database) _persist_message(database, mem.conversation.id, query="My query", answer="My answer", answer_tokens=10) result = mem.get_history_prompt_messages() @@ -624,7 +676,7 @@ class TestGetHistoryPromptMessages: assert result[1].content == "My answer" def test_history_is_conversation_scoped(self, database: Database) -> None: - mem = self._make_memory() + mem = self._make_memory(database) _persist_message(database, mem.conversation.id, answer="visible") _persist_message(database, "other-conversation", answer="hidden") @@ -642,12 +694,12 @@ class TestGetHistoryPromptMessages: message_limit: int | None, expected_limit: int, ) -> None: - mem = self._make_memory() + mem = self._make_memory(database) before = len(database.statements) mem.get_history_prompt_messages(message_limit=message_limit) - statements = database.statements[before:] + statements = [entry for entry in database.statements[before:] if "FROM messages" in entry[0]] assert len(statements) == 1 sql, parameters = statements[0] assert "LIMIT" in sql @@ -667,7 +719,7 @@ class TestGetHistoryPromptMessages: belongs_to: MessageFileBelongsTo | None, is_user_message: bool, ) -> None: - mem = self._make_memory() + mem = self._make_memory(database) message = _persist_message(database, mem.conversation.id) message_file = _persist_message_file(database, message, belongs_to=belongs_to) built_prompt = ( @@ -685,7 +737,7 @@ class TestGetHistoryPromptMessages: assert built_prompt in result def test_message_files_are_batch_loaded_with_constant_query_count(self, database: Database) -> None: - mem = self._make_memory() + mem = self._make_memory(database) base_time = datetime.now(UTC).replace(tzinfo=None) messages = [ _persist_message( @@ -703,7 +755,8 @@ class TestGetHistoryPromptMessages: result = mem.get_history_prompt_messages() selects = [sql for sql, _ in database.statements[before:] if sql.lstrip().upper().startswith("SELECT")] - assert len(selects) == 3 + assert len(selects) == 4 + assert sum("FROM apps" in sql for sql in selects) == 1 assert len(result) == 10 @pytest.mark.parametrize( @@ -721,7 +774,7 @@ class TestGetHistoryPromptMessages: max_token_limit: int, expected_length: int, ) -> None: - mem = self._make_memory() + mem = self._make_memory(database) mem.model_instance.get_llm_num_tokens.side_effect = token_values _persist_message(database, mem.conversation.id) @@ -740,7 +793,6 @@ class TestGetHistoryPromptText: def _make_memory(self) -> TokenBufferMemory: conv = _make_conversation() - conv.app = MagicMock() return TokenBufferMemory(conversation=conv, model_instance=_make_model_instance()) def test_empty_messages_returns_empty_string(self): diff --git a/api/tests/unit_tests/core/plugin/test_backwards_invocation_model.py b/api/tests/unit_tests/core/plugin/test_backwards_invocation_model.py index c24d3ac0120..9f480d9448f 100644 --- a/api/tests/unit_tests/core/plugin/test_backwards_invocation_model.py +++ b/api/tests/unit_tests/core/plugin/test_backwards_invocation_model.py @@ -1,9 +1,9 @@ -from types import SimpleNamespace from unittest.mock import patch from core.plugin.backwards_invocation.model import PluginModelBackwardsInvocation from core.plugin.entities.request import RequestInvokeSummary from graphon.model_runtime.entities.message_entities import UserPromptMessage +from models.account import Tenant def test_system_model_helpers_forward_user_id(): @@ -36,7 +36,8 @@ def test_system_model_helpers_forward_user_id(): def test_invoke_summary_uses_same_user_scope_for_token_helpers(): - tenant = SimpleNamespace(id="tenant-1") + tenant = Tenant(name="Test Workspace") + tenant.id = "tenant-1" payload = RequestInvokeSummary(text="short", instruction="keep it concise") with ( diff --git a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py index 1642f8eb552..9fa76dd9737 100644 --- a/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py +++ b/api/tests/unit_tests/core/rag/datasource/keyword/test_keyword_factory.py @@ -1,6 +1,5 @@ import sys import types -from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -9,6 +8,7 @@ from sqlalchemy.orm import Session from core.rag.datasource.keyword.keyword_factory import Keyword from core.rag.datasource.keyword.keyword_type import KeyWordType from core.rag.models.document import Document +from models.dataset import Dataset def test_get_keyword_factory_returns_jieba_factory(monkeypatch: pytest.MonkeyPatch): @@ -29,7 +29,13 @@ def test_get_keyword_factory_raises_for_unsupported_type(): def test_keyword_initialization_uses_configured_factory(monkeypatch: pytest.MonkeyPatch): - dataset = SimpleNamespace(id="dataset-1") + dataset = Dataset( + id="dataset-1", + tenant_id="tenant-1", + name="Test Dataset", + description="", + created_by="account-1", + ) fake_processor = MagicMock() monkeypatch.setattr("core.rag.datasource.keyword.keyword_factory.dify_config.KEYWORD_STORE", KeyWordType.JIEBA) diff --git a/api/tests/unit_tests/enterprise/telemetry/test_draft_trace.py b/api/tests/unit_tests/enterprise/telemetry/test_draft_trace.py index c8c8de8595c..747b662de3b 100644 --- a/api/tests/unit_tests/enterprise/telemetry/test_draft_trace.py +++ b/api/tests/unit_tests/enterprise/telemetry/test_draft_trace.py @@ -2,39 +2,60 @@ from __future__ import annotations +import json from datetime import UTC, datetime +from typing import cast from unittest.mock import MagicMock, patch from graphon.enums import WorkflowNodeExecutionMetadataKey +from models.enums import CreatorUserRole +from models.workflow import ( + WorkflowNodeExecutionModel, + WorkflowNodeExecutionStatus, + WorkflowNodeExecutionTriggeredFrom, +) # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- -def _make_execution(**overrides) -> MagicMock: - """Return a minimal WorkflowNodeExecutionModel mock.""" - execution = MagicMock() - execution.tenant_id = overrides.get("tenant_id", "tenant-1") - execution.app_id = overrides.get("app_id", "app-1") - execution.workflow_id = overrides.get("workflow_id", "wf-1") - execution.id = overrides.get("id", "exec-1") - execution.node_id = overrides.get("node_id", "node-1") - execution.node_type = overrides.get("node_type", "llm") - execution.title = overrides.get("title", "My LLM Node") - execution.status = overrides.get("status", "succeeded") - execution.error = overrides.get("error") - execution.elapsed_time = overrides.get("elapsed_time", 1.5) - execution.index = overrides.get("index", 1) - execution.predecessor_node_id = overrides.get("predecessor_node_id") - execution.created_at = overrides.get("created_at", datetime(2024, 1, 1, tzinfo=UTC)) - execution.finished_at = overrides.get("finished_at", datetime(2024, 1, 1, 0, 0, 5, tzinfo=UTC)) - execution.workflow_run_id = overrides.get("workflow_run_id", "run-1") - execution.inputs_dict = overrides.get("inputs_dict", {"prompt": "hello"}) - execution.outputs_dict = overrides.get("outputs_dict", {"answer": "world"}) - execution.process_data_dict = overrides.get("process_data_dict", {}) - execution.execution_metadata_dict = overrides.get("execution_metadata_dict", {}) - return execution +def _make_execution(**overrides: object) -> WorkflowNodeExecutionModel: + """Return a real transient execution with JSON-backed model properties.""" + + inputs = overrides.get("inputs_dict", {"prompt": "hello"}) + outputs = overrides.get("outputs_dict", {"answer": "world"}) + process_data = overrides.get("process_data_dict", {}) + metadata = overrides.get("execution_metadata_dict", {}) + return WorkflowNodeExecutionModel( + id=cast(str, overrides.get("id", "exec-1")), + tenant_id=cast(str, overrides.get("tenant_id", "tenant-1")), + app_id=cast(str, overrides.get("app_id", "app-1")), + workflow_id=cast(str, overrides.get("workflow_id", "wf-1")), + triggered_from=WorkflowNodeExecutionTriggeredFrom.WORKFLOW_RUN, + workflow_run_id=cast(str | None, overrides.get("workflow_run_id", "run-1")), + index=cast(int, overrides.get("index", 1)), + predecessor_node_id=cast(str | None, overrides.get("predecessor_node_id")), + node_execution_id=None, + node_id=cast(str, overrides.get("node_id", "node-1")), + node_type=cast(str, overrides.get("node_type", "llm")), + title=cast(str, overrides.get("title", "My LLM Node")), + agent_workspace_binding_id=None, + inputs=json.dumps(inputs) if inputs is not None else None, + process_data=json.dumps(process_data) if process_data is not None else None, + outputs=json.dumps(outputs) if outputs is not None else None, + status=WorkflowNodeExecutionStatus(cast(str, overrides.get("status", "succeeded"))), + error=cast(str | None, overrides.get("error")), + elapsed_time=cast(float, overrides.get("elapsed_time", 1.5)), + execution_metadata=json.dumps(metadata), + created_at=cast(datetime, overrides.get("created_at", datetime(2024, 1, 1, tzinfo=UTC))), + created_by_role=CreatorUserRole.ACCOUNT, + created_by="user-1", + finished_at=cast( + datetime | None, + overrides.get("finished_at", datetime(2024, 1, 1, 0, 0, 5, tzinfo=UTC)), + ), + ) # --------------------------------------------------------------------------- @@ -289,8 +310,8 @@ class TestEnqueueDraftNodeExecutionTrace: # --------------------------------------------------------------------------- -def _make_llm_execution() -> MagicMock: - """Return a WorkflowNodeExecutionModel mock that mimics a real LLM node. +def _make_llm_execution() -> WorkflowNodeExecutionModel: + """Return a real WorkflowNodeExecutionModel that mimics an LLM node. The field values match what graphon/nodes/llm/node.py produces: - process_data_dict contains model_provider, model_name, and usage diff --git a/api/tests/unit_tests/libs/test_token_manager.py b/api/tests/unit_tests/libs/test_token_manager.py index bbe8a7e30bc..e8f975ae974 100644 --- a/api/tests/unit_tests/libs/test_token_manager.py +++ b/api/tests/unit_tests/libs/test_token_manager.py @@ -15,6 +15,7 @@ from pydantic import ValidationError import libs.helper as helper_module from libs.helper import TokenManager +from models.account import Account def _build_fake_redis(storage: dict[str, str]): @@ -70,7 +71,8 @@ def test_token_manager_roundtrip_uses_explicit_email_with_account(monkeypatch: p storage: dict[str, str] = {} monkeypatch.setattr(helper_module, "redis_client", _build_fake_redis(storage)) - account = SimpleNamespace(id="acc-1", email="old@example.com") + account = Account(name="Test User", email="old@example.com") + account.id = "acc-1" token = TokenManager.generate_token( account=account,