test: use sqlite3 session in test_trace_session_metadata (#38750)

This commit is contained in:
Asuka Minato 2026-07-13 15:05:07 +09:00 committed by GitHub
parent 831443b45c
commit 5f8ef3c321
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -3,38 +3,29 @@ from datetime import datetime, timedelta
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock 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.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceTask 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: @pytest.fixture(autouse=True)
scalar_values: list[object | None] = [] def _bind_trace_database(
monkeypatch: pytest.MonkeyPatch,
def __init__(self, engine): sqlite_engine: Engine,
self._values = list(self.scalar_values) sqlite_session: Session,
self._index = 0 ) -> None:
"""Use real SQLite sessions for ORM lookups without changing trace-domain data."""
def __enter__(self): monkeypatch.setattr(
return self "core.ops.ops_trace_manager.db",
SimpleNamespace(engine=sqlite_engine, session=sqlite_session),
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 []
def _make_workflow_run(): def _make_workflow_run():
@ -92,14 +83,12 @@ def _make_message_data():
return _MessageData(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 = MagicMock()
repo.get_workflow_run_by_id_without_tenant.return_value = _make_workflow_run() 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(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) monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False)
_DummySession.scalar_values = [None, None]
task = TraceTask( task = TraceTask(
TraceTaskName.WORKFLOW_TRACE, 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" assert trace_info.metadata["trace_session_id"] == "session-1"
def test_message_trace_metadata_includes_trace_session_id(monkeypatch): @pytest.mark.parametrize("sqlite_session", [TABLES], indirect=True)
db_session = MagicMock() def test_message_trace_metadata_includes_trace_session_id(monkeypatch, sqlite_session: Session):
db_session.scalars.return_value.all.return_value = ["chat"] app = App(
db_session.scalar.return_value = None id="app-1",
monkeypatch.setattr( tenant_id="tenant-1",
"core.ops.ops_trace_manager.db", name="Trace App",
SimpleNamespace(engine=MagicMock(), session=db_session), 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.ops.ops_trace_manager.get_message_data", lambda message_id: _make_message_data())
monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False) monkeypatch.setattr("core.telemetry.gateway.is_enterprise_telemetry_enabled", lambda: False)
_DummySession.scalar_values = ["tenant-1"]
task = TraceTask( task = TraceTask(
TraceTaskName.MESSAGE_TRACE, TraceTaskName.MESSAGE_TRACE,