From c5b0772ab056f91749ee63865f06f1ee8e62362e Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Wed, 29 Jul 2026 14:52:33 +0900 Subject: [PATCH] test: use SQLite sessions in core tools (#39107) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- .../core/tools/test_builtin_tools_extra.py | 84 +++++++------ .../unit_tests/core/tools/test_tool_engine.py | 115 ++++++++++++------ 2 files changed, 119 insertions(+), 80 deletions(-) diff --git a/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py b/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py index 4dac9b7260d..d6bb7503793 100644 --- a/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py +++ b/api/tests/unit_tests/core/tools/test_builtin_tools_extra.py @@ -4,10 +4,10 @@ import calendar import math from datetime import date from types import SimpleNamespace -from unittest.mock import MagicMock from zoneinfo import ZoneInfo import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool_runtime import ToolRuntime @@ -51,24 +51,26 @@ def _raise_runtime_error(*_args: object, **_kwargs: object) -> None: raise RuntimeError("boom") -def test_current_time_tool(): +def test_current_time_tool(sqlite_session: Session): current_tool = _build_builtin_tool(CurrentTimeTool) - utc_text = list(current_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"timezone": "UTC"}))[ + utc_text = list(current_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"timezone": "UTC"}))[ 0 ].message.text assert utc_text invalid_tz = list( - current_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"timezone": "Invalid/TZ"}) + current_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"timezone": "Invalid/TZ"}) )[0].message.text assert "Invalid timezone" in invalid_tz -def test_localtime_to_timestamp_tool(): +def test_localtime_to_timestamp_tool(sqlite_session: Session): localtime_tool = _build_builtin_tool(LocaltimeToTimestampTool) ts_message = list( localtime_tool.invoke( - session=MagicMock(), user_id="u", tool_parameters={"localtime": "2024-01-01 10:00:00", "timezone": "UTC"} + session=sqlite_session, + user_id="u", + tool_parameters={"localtime": "2024-01-01 10:00:00", "timezone": "UTC"}, ) )[0].message.text ts_value = float(ts_message.strip()) @@ -92,11 +94,11 @@ def test_localtime_to_timestamp_tool(): LocaltimeToTimestampTool.localtime_to_timestamp("bad", "%Y-%m-%d %H:%M:%S", "UTC") -def test_timestamp_to_localtime_tool(): +def test_timestamp_to_localtime_tool(sqlite_session: Session): to_local_tool = _build_builtin_tool(TimestampToLocaltimeTool) local_text = list( to_local_tool.invoke( - session=MagicMock(), user_id="u", tool_parameters={"timestamp": 1704067200, "timezone": "UTC"} + session=sqlite_session, user_id="u", tool_parameters={"timestamp": 1704067200, "timezone": "UTC"} ) )[0].message.text assert "2024" in local_text @@ -104,11 +106,11 @@ def test_timestamp_to_localtime_tool(): TimestampToLocaltimeTool.timestamp_to_localtime("bad", "UTC") # type: ignore[arg-type] -def test_timezone_conversion_tool(): +def test_timezone_conversion_tool(sqlite_session: Session): timezone_tool = _build_builtin_tool(TimezoneConversionTool) converted = list( timezone_tool.invoke( - session=MagicMock(), + session=sqlite_session, user_id="u", tool_parameters={ "current_time": "2024-01-01 08:00:00", @@ -122,10 +124,10 @@ def test_timezone_conversion_tool(): TimezoneConversionTool.timezone_convert("bad", "UTC", "Asia/Tokyo") -def test_weekday_tool(): +def test_weekday_tool(sqlite_session: Session): weekday_tool = _build_builtin_tool(WeekdayTool) valid = list( - weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "month": 1, "day": 1}) + weekday_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"year": 2024, "month": 1, "day": 1}) )[0].message.text expected_date = date(2024, 1, 1) expected_message = ( @@ -135,14 +137,14 @@ def test_weekday_tool(): ) assert valid == expected_message invalid = list( - weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "month": 2, "day": 31}) + weekday_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"year": 2024, "month": 2, "day": 31}) )[0].message.text assert "Invalid date" in invalid with pytest.raises(ValueError, match="Month is required"): - list(weekday_tool.invoke(session=MagicMock(), user_id="u", tool_parameters={"year": 2024, "day": 1})) + list(weekday_tool.invoke(session=sqlite_session, user_id="u", tool_parameters={"year": 2024, "day": 1})) -def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): +def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): simple_code = _build_builtin_tool(SimpleCode) monkeypatch.setattr( @@ -151,7 +153,7 @@ def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): ) result = list( simple_code.invoke( - session=MagicMock(), + session=sqlite_session, user_id="u", tool_parameters={"language": "python3", "code": "print(1)"}, ) @@ -159,18 +161,18 @@ def test_simple_code_valid_execution(monkeypatch: pytest.MonkeyPatch): assert result == "ok" -def test_simple_code_invalid_language(): +def test_simple_code_invalid_language(sqlite_session: Session): simple_code = _build_builtin_tool(SimpleCode) with pytest.raises(ValueError, match="Only python3 and javascript"): list( simple_code.invoke( - session=MagicMock(), user_id="u", tool_parameters={"language": "go", "code": "fmt.Println(1)"} + session=sqlite_session, user_id="u", tool_parameters={"language": "go", "code": "fmt.Println(1)"} ) ) -def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch): +def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): simple_code = _build_builtin_tool(SimpleCode) monkeypatch.setattr( @@ -180,33 +182,35 @@ def test_simple_code_execution_error(monkeypatch: pytest.MonkeyPatch): with pytest.raises(ToolInvokeError, match="boom"): list( simple_code.invoke( - session=MagicMock(), user_id="u", tool_parameters={"language": "python3", "code": "print(1)"} + session=sqlite_session, + user_id="u", + tool_parameters={"language": "python3", "code": "print(1)"}, ) ) -def test_webscraper_empty_url(): +def test_webscraper_empty_url(sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) - empty = list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": ""}))[0].message.text + empty = list(webscraper.invoke(session=sqlite_session, user_id="u", tool_parameters={"url": ""}))[0].message.text assert empty == "Please input url" -def test_webscraper_fetch(monkeypatch: pytest.MonkeyPatch): +def test_webscraper_fetch(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr("core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", lambda *a, **k: "page") - full = list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com"}))[ + full = list(webscraper.invoke(session=sqlite_session, user_id="u", tool_parameters={"url": "https://example.com"}))[ 0 ].message.text assert full == "page" -def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch): +def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr("core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", lambda *a, **k: "page") monkeypatch.setattr(webscraper, "summary", lambda user_id, content: "summary") summarized = list( webscraper.invoke( - session=MagicMock(), + session=sqlite_session, user_id="u", tool_parameters={"url": "https://example.com", "generate_summary": True}, ) @@ -214,26 +218,26 @@ def test_webscraper_summary(monkeypatch: pytest.MonkeyPatch): assert summarized == "summary" -def test_webscraper_fetch_error(monkeypatch: pytest.MonkeyPatch): +def test_webscraper_fetch_error(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): webscraper = _build_builtin_tool(WebscraperTool) monkeypatch.setattr( "core.tools.builtin_tool.providers.webscraper.tools.webscraper.get_url", _raise_runtime_error, ) with pytest.raises(ToolInvokeError, match="boom"): - list(webscraper.invoke(session=MagicMock(), user_id="u", tool_parameters={"url": "https://example.com"})) + list(webscraper.invoke(session=sqlite_session, user_id="u", tool_parameters={"url": "https://example.com"})) -def test_asr_invalid_file(): +def test_asr_invalid_file(sqlite_session: Session): asr = _build_builtin_tool(ASRTool) file_obj = SimpleNamespace(type=FileType.DOCUMENT) - invalid_file = list(asr.invoke(session=MagicMock(), user_id="u", tool_parameters={"audio_file": file_obj}))[ + invalid_file = list(asr.invoke(session=sqlite_session, user_id="u", tool_parameters={"audio_file": file_obj}))[ 0 ].message.text assert "not a valid audio file" in invalid_file -def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch): +def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): asr = _build_builtin_tool(ASRTool) model_instance = type("M", (), {"invoke_speech2text": lambda self, file: "transcript"})() model_manager = type("Mgr", (), {"get_model_instance": lambda *a, **k: model_instance})() @@ -245,9 +249,9 @@ def test_asr_valid_file_invocation(monkeypatch: pytest.MonkeyPatch): lambda **kwargs: captured_manager_kwargs.update(kwargs) or model_manager, ) audio_file = SimpleNamespace(type=FileType.AUDIO) - ok = list(asr.invoke(session=MagicMock(), user_id="u", tool_parameters={"audio_file": audio_file, "model": "p#m"}))[ - 0 - ].message.text + ok = list( + asr.invoke(session=sqlite_session, user_id="u", tool_parameters={"audio_file": audio_file, "model": "p#m"}) + )[0].message.text assert ok == "transcript" assert captured_manager_kwargs == {"tenant_id": "tenant-1", "user_id": "u"} @@ -263,7 +267,7 @@ def test_asr_available_models_and_runtime_parameters(monkeypatch: pytest.MonkeyP assert asr.get_runtime_parameters()[0].name == "model" -def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch): +def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): tts = _build_builtin_tool(TTSTool) captured_manager_kwargs = {} voices_model_instance = type( @@ -281,7 +285,7 @@ def test_tts_invoke_returns_messages(monkeypatch: pytest.MonkeyPatch): or type("M", (), {"get_model_instance": lambda *a, **k: voices_model_instance})() ), ) - messages = list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + messages = list(tts.invoke(session=sqlite_session, user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) assert [m.type for m in messages] == [ToolInvokeMessage.MessageType.TEXT, ToolInvokeMessage.MessageType.BLOB] assert captured_manager_kwargs == {"tenant_id": "tenant-1", "user_id": "u"} @@ -293,18 +297,18 @@ def test_tts_get_available_models_requires_runtime(): tts.get_available_models() -def test_tts_tool_raises_when_runtime_missing(): +def test_tts_tool_raises_when_runtime_missing(sqlite_session: Session): tts = _build_builtin_tool(TTSTool) tts.runtime = None with pytest.raises(ValueError, match="Runtime is required"): - list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + list(tts.invoke(session=sqlite_session, user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) @pytest.mark.parametrize( "voices", [[{"value": None}], []], ) -def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices): +def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices, sqlite_session: Session): tts = _build_builtin_tool(TTSTool) tts.runtime = ToolRuntime(tenant_id="tenant-1", invoke_from=InvokeFrom.DEBUGGER) model_without_voice = type( @@ -320,7 +324,7 @@ def test_tts_tool_raises_when_voice_unavailable(monkeypatch, voices): lambda **_: type("Manager", (), {"get_model_instance": lambda *args, **kwargs: model_without_voice})(), ) with pytest.raises(ValueError, match="no voice available"): - list(tts.invoke(session=MagicMock(), user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) + list(tts.invoke(session=sqlite_session, user_id="u", tool_parameters={"model": "p#m", "text": "hello"})) def test_tts_tool_get_available_models_and_runtime_parameters(monkeypatch: pytest.MonkeyPatch): diff --git a/api/tests/unit_tests/core/tools/test_tool_engine.py b/api/tests/unit_tests/core/tools/test_tool_engine.py index f38ab2a2fab..f688c68cfa6 100644 --- a/api/tests/unit_tests/core/tools/test_tool_engine.py +++ b/api/tests/unit_tests/core/tools/test_tool_engine.py @@ -1,11 +1,14 @@ from __future__ import annotations from collections.abc import Generator -from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock, patch +from uuid import uuid4 import pytest +from sqlalchemy import select +from sqlalchemy.engine import Engine +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool import Tool @@ -26,6 +29,45 @@ from core.tools.errors import ( ToolParameterValidationError, ) from core.tools.tool_engine import ToolEngine +from models.model import AppMode, Message, MessageFile + + +class _DatabaseBinding: + engine: Engine + + def __init__(self, engine: Engine) -> None: + self.engine = engine + + +def _message() -> Message: + message = Message( + app_id=str(uuid4()), + model_provider="provider", + model_id="model", + override_model_configs=None, + conversation_id=str(uuid4()), + inputs={}, + query="query", + message="", + message_tokens=0, + message_unit_price=0, + message_price_unit=0, + answer="", + answer_tokens=0, + answer_unit_price=0, + answer_price_unit=0, + parent_message_id=None, + provider_response_latency=0, + total_price=0, + currency="USD", + invoke_from="debugger", + from_source="console", + from_end_user_id=None, + from_account_id=str(uuid4()), + app_mode=AppMode.CHAT, + ) + message.id = str(uuid4()) + return message class _DummyTool(Tool): @@ -120,52 +162,41 @@ def test_convert_tool_response_to_str_and_extract_binary_messages(): ) -def test_create_message_files_and_invoke_generator(): +@pytest.mark.parametrize("sqlite_session", [(MessageFile,)], indirect=True) +def test_create_message_files_and_invoke_generator(sqlite_engine: Engine, sqlite_session: Session): binaries = [ ToolInvokeMessageBinary(mimetype="image/png", url="https://example.com/abc.png"), ToolInvokeMessageBinary(mimetype="audio/wav", url="https://example.com/def.wav"), ] - created = [] - - def _message_file_factory(**kwargs): - obj = SimpleNamespace(id=f"mf-{len(created) + 1}", **kwargs) - created.append(obj) - return obj - - file_session = MagicMock() - session_factory = MagicMock() - session_factory.begin.return_value.__enter__.return_value = file_session - with ( - patch("core.tools.tool_engine.MessageFile", side_effect=_message_file_factory), - patch("core.tools.tool_engine.db") as mock_db, - patch("core.tools.tool_engine.sessionmaker", return_value=session_factory) as mock_sessionmaker, - ): + agent_message = _message() + with patch("core.tools.tool_engine.db", _DatabaseBinding(sqlite_engine)): ids = ToolEngine._create_message_files( tool_messages=binaries, - agent_message=SimpleNamespace(id="msg-1"), + agent_message=agent_message, invoke_from=InvokeFrom.DEBUGGER, - user_id="user-1", + user_id=str(uuid4()), ) - assert ids == ["mf-1", "mf-2"] - mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) - assert file_session.add.call_count == 2 - mock_db.session.close.assert_not_called() + message_files = list(sqlite_session.scalars(select(MessageFile).order_by(MessageFile.created_at)).all()) + assert ids == [message_file.id for message_file in message_files] + assert len(message_files) == 2 + assert {message_file.message_id for message_file in message_files} == {agent_message.id} tool = _build_tool() - invoked = list(ToolEngine._invoke(MagicMock(), tool, {"a": 1}, user_id="u")) + invoked = list(ToolEngine._invoke(sqlite_session, tool, {"a": 1}, user_id="u")) assert invoked[0].type == ToolInvokeMessage.MessageType.TEXT assert isinstance(invoked[-1], ToolInvokeMeta) assert invoked[-1].error is None -def test_generic_invoke_success_and_error_paths(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_generic_invoke_success_and_error_paths(sqlite_session: Session): tool = _build_tool() callback = Mock() callback.on_tool_execution.side_effect = lambda **kwargs: kwargs["tool_outputs"] response = list( ToolEngine.generic_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"x": 1}, user_id="u1", @@ -186,7 +217,7 @@ def test_generic_invoke_success_and_error_paths(): with pytest.raises(RuntimeError, match="boom"): list( ToolEngine.generic_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"x": 1}, user_id="u1", @@ -197,10 +228,11 @@ def test_generic_invoke_success_and_error_paths(): error_callback.on_tool_error.assert_called_once() -def test_agent_invoke_success(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_success(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() meta = ToolInvokeMeta.empty() with patch.object(ToolEngine, "_invoke", return_value=iter([tool.create_text_message("ok"), meta])): @@ -211,7 +243,7 @@ def test_agent_invoke_success(): with patch.object(ToolEngine, "_extract_tool_response_binary_and_text", return_value=iter([])): with patch.object(ToolEngine, "_create_message_files", return_value=[]): result_text, message_files, result_meta = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters="hello", user_id="u1", @@ -228,14 +260,15 @@ def test_agent_invoke_success(): callback.on_tool_end.assert_called_once() -def test_agent_invoke_param_validation_error(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_param_validation_error(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() with patch.object(ToolEngine, "_invoke", side_effect=ToolParameterValidationError("bad-param")): error_text, files, error_meta = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"a": 1}, user_id="u1", @@ -250,15 +283,16 @@ def test_agent_invoke_param_validation_error(): assert error_meta.error -def test_agent_invoke_engine_meta_error(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_engine_meta_error(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() engine_error = ToolEngineInvokeError(ToolInvokeMeta.error_instance("meta failure")) with patch.object(ToolEngine, "_invoke", side_effect=engine_error): error_text, files, error_meta = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"a": 1}, user_id="u1", @@ -295,14 +329,15 @@ def test_convert_tool_response_excludes_variable_messages(): assert "variable_name" not in result -def test_agent_invoke_tool_invoke_error(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_agent_invoke_tool_invoke_error(sqlite_session: Session): tool = _build_tool(with_llm_parameter=True) callback = Mock() - message = SimpleNamespace(id="m1", conversation_id="c1") + message = _message() with patch.object(ToolEngine, "_invoke", side_effect=ToolInvokeError("invoke boom")): error_text, files, _ = ToolEngine.agent_invoke( - session=MagicMock(), + session=sqlite_session, tool=tool, tool_parameters={"a": 1}, user_id="u1",