From 5f8ef3c3219828023a025152c2bdabd1636cc850 Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Mon, 13 Jul 2026 15:05:07 +0900 Subject: [PATCH] test: use sqlite3 session in test_trace_session_metadata (#38750) --- .../core/ops/test_trace_session_metadata.py | 102 +++++++++++------- 1 file changed, 61 insertions(+), 41 deletions(-) diff --git a/api/tests/unit_tests/core/ops/test_trace_session_metadata.py b/api/tests/unit_tests/core/ops/test_trace_session_metadata.py index 9a5cb0b65fe..a0d7cfbce14 100644 --- a/api/tests/unit_tests/core/ops/test_trace_session_metadata.py +++ b/api/tests/unit_tests/core/ops/test_trace_session_metadata.py @@ -3,38 +3,29 @@ from datetime import datetime, timedelta from types import SimpleNamespace from unittest.mock import MagicMock +import pytest +from sqlalchemy import Engine +from sqlalchemy.orm import Session + from core.ops.entities.trace_entity import TraceTaskName from core.ops.ops_trace_manager import TraceTask +from models.model import App, AppMode, Conversation, Message, MessageFile +from models.workflow import WorkflowAppLog, WorkflowNodeExecutionModel + +TABLES = (App, Conversation, Message, MessageFile, WorkflowAppLog, WorkflowNodeExecutionModel) -class _DummySession: - scalar_values: list[object | None] = [] - - def __init__(self, engine): - self._values = list(self.scalar_values) - self._index = 0 - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - return False - - def execute(self, *args, **kwargs): - return self - - def scalar(self, *args, **kwargs): - if self._index >= len(self._values): - return None - value = self._values[self._index] - self._index += 1 - return value - - def scalars(self, *args, **kwargs): - return self - - def all(self): - return [] +@pytest.fixture(autouse=True) +def _bind_trace_database( + monkeypatch: pytest.MonkeyPatch, + sqlite_engine: Engine, + sqlite_session: Session, +) -> None: + """Use real SQLite sessions for ORM lookups without changing trace-domain data.""" + monkeypatch.setattr( + "core.ops.ops_trace_manager.db", + SimpleNamespace(engine=sqlite_engine, session=sqlite_session), + ) def _make_workflow_run(): @@ -92,14 +83,12 @@ def _make_message_data(): return _MessageData(data) -def test_workflow_trace_metadata_includes_trace_session_id(monkeypatch): +@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) +def test_workflow_trace_metadata_includes_trace_session_id(monkeypatch, sqlite_session: Session): repo = MagicMock() repo.get_workflow_run_by_id_without_tenant.return_value = _make_workflow_run() monkeypatch.setattr(TraceTask, "_get_workflow_run_repo", classmethod(lambda cls: repo)) - monkeypatch.setattr("core.ops.ops_trace_manager.Session", _DummySession) - monkeypatch.setattr("core.ops.ops_trace_manager.db", SimpleNamespace(engine=MagicMock())) monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False) - _DummySession.scalar_values = [None, None] task = TraceTask( TraceTaskName.WORKFLOW_TRACE, @@ -115,18 +104,49 @@ def test_workflow_trace_metadata_includes_trace_session_id(monkeypatch): assert trace_info.metadata["trace_session_id"] == "session-1" -def test_message_trace_metadata_includes_trace_session_id(monkeypatch): - db_session = MagicMock() - db_session.scalars.return_value.all.return_value = ["chat"] - db_session.scalar.return_value = None - monkeypatch.setattr( - "core.ops.ops_trace_manager.db", - SimpleNamespace(engine=MagicMock(), session=db_session), +@pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True) +def test_message_trace_metadata_includes_trace_session_id(monkeypatch, sqlite_session: Session): + app = App( + id="app-1", + tenant_id="tenant-1", + name="Trace App", + description="", + mode=AppMode.CHAT, + icon_type=None, + icon=None, + icon_background=None, + app_model_config_id=None, + workflow_id=None, + enable_site=True, + enable_api=True, + max_active_requests=None, + created_by=None, ) - monkeypatch.setattr("core.ops.ops_trace_manager.Session", _DummySession) + conversation = Conversation( + id="conv-1", + app_id=app.id, + app_model_config_id=None, + model_provider=None, + override_model_configs=None, + model_id=None, + mode=AppMode.CHAT, + name="Trace Conversation", + summary=None, + inputs={}, + introduction="", + system_instruction="", + invoke_from=None, + from_source="api", + from_end_user_id="end-user-1", + from_account_id=None, + read_at=None, + read_account_id=None, + ) + sqlite_session.add_all([app, conversation]) + sqlite_session.commit() + monkeypatch.setattr("core.ops.ops_trace_manager.get_message_data", lambda message_id: _make_message_data()) monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False) - _DummySession.scalar_values = ["tenant-1"] task = TraceTask( TraceTaskName.MESSAGE_TRACE,