diff --git a/api/tests/unit_tests/core/tools/test_base_tool.py b/api/tests/unit_tests/core/tools/test_base_tool.py index f164e3fddea..b28b28003c4 100644 --- a/api/tests/unit_tests/core/tools/test_base_tool.py +++ b/api/tests/unit_tests/core/tools/test_base_tool.py @@ -3,9 +3,9 @@ from __future__ import annotations from collections.abc import Generator from dataclasses import dataclass from typing import Any, cast -from unittest.mock import MagicMock import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool import Tool @@ -94,7 +94,8 @@ def _build_tool(runtime: ToolRuntime | None = None) -> DummyTool: return DummyTool(entity=entity, runtime=runtime) -def test_invoke_supports_single_message_and_parameter_casting(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_invoke_supports_single_message_and_parameter_casting(sqlite_session: Session): runtime = ToolRuntime( tenant_id="tenant-1", invoke_from=InvokeFrom.DEBUGGER, @@ -112,7 +113,7 @@ def test_invoke_supports_single_message_and_parameter_casting(): messages = list( tool.invoke( - session=MagicMock(), + session=sqlite_session, user_id="user-1", tool_parameters={"age": "18", "raw": "keep"}, conversation_id="conv-1", @@ -132,7 +133,8 @@ def test_invoke_supports_single_message_and_parameter_casting(): } -def test_invoke_preserves_multiple_select_values(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_invoke_preserves_multiple_select_values(sqlite_session: Session): tool = _build_tool() parameter = ToolParameter.get_simple_instance( name="choice", @@ -144,18 +146,19 @@ def test_invoke_preserves_multiple_select_values(): parameter.multiple = True tool.entity.parameters = [parameter] - list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": ["a", "b"]})) + list(tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={"choice": ["a", "b"]})) assert tool.last_invocation is not None assert tool.last_invocation["tool_parameters"] == {"choice": ["a", "b"]} with pytest.raises(ValueError, match="must be a list"): - tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={"choice": "a"}) + tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={"choice": "a"}) -def test_invoke_supports_list_and_generator_results(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_invoke_supports_list_and_generator_results(sqlite_session: Session): tool = _build_tool() tool.result = [tool.create_text_message("a"), tool.create_text_message("b")] - list_messages = list(tool.invoke(session=MagicMock(), user_id="user-1", tool_parameters={})) + list_messages = list(tool.invoke(session=sqlite_session, user_id="user-1", tool_parameters={})) assert [msg.message.text for msg in list_messages] == ["a", "b"] def _message_generator() -> Generator[ToolInvokeMessage, None, None]: @@ -163,7 +166,7 @@ def test_invoke_supports_list_and_generator_results(): yield tool.create_text_message("g2") tool.result = _message_generator() - generated_messages = list(tool.invoke(session=MagicMock(), user_id="user-2", tool_parameters={})) + generated_messages = list(tool.invoke(session=sqlite_session, user_id="user-2", tool_parameters={})) assert [msg.message.text for msg in generated_messages] == ["g1", "g2"] @@ -372,6 +375,7 @@ def test_message_factory_helpers(): assert variable_message.message.stream is False -def test_base_abstract_invoke_placeholder_returns_none(): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_base_abstract_invoke_placeholder_returns_none(sqlite_session: Session): tool = _build_tool() - assert Tool._invoke(tool, session=MagicMock(), user_id="u", tool_parameters={}) is None + assert Tool._invoke(tool, session=sqlite_session, user_id="u", tool_parameters={}) is None diff --git a/api/tests/unit_tests/core/tools/test_custom_tool.py b/api/tests/unit_tests/core/tools/test_custom_tool.py index 64cd128f18e..933c64bf972 100644 --- a/api/tests/unit_tests/core/tools/test_custom_tool.py +++ b/api/tests/unit_tests/core/tools/test_custom_tool.py @@ -2,10 +2,10 @@ from __future__ import annotations from types import SimpleNamespace from typing import Any -from unittest.mock import MagicMock import httpx import pytest +from sqlalchemy.orm import Session from core.app.entities.app_invoke_entities import InvokeFrom from core.tools.__base.tool_runtime import ToolRuntime @@ -241,7 +241,8 @@ def test_do_http_request_builds_arguments_and_handles_invalid_method(monkeypatch invalid_method_tool.do_http_request("https://api.example.com", "TRACE", headers={}, parameters={}) -def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytest.MonkeyPatch): +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytest.MonkeyPatch, sqlite_session: Session): openapi = { "parameters": [], "requestBody": { @@ -281,11 +282,11 @@ def test_do_http_request_handles_file_upload_and_invoke_paths(monkeypatch: pytes monkeypatch.setattr(tool, "assembling_request", lambda parameters: {}) monkeypatch.setattr(tool, "do_http_request", lambda *args, **kwargs: httpx.Response(200, text='{"a":1}')) monkeypatch.setattr(tool, "validate_and_parse_response", lambda _: ParsedResponse({"a": 1}, True)) - messages = list(tool.invoke(session=MagicMock(), user_id="u1", tool_parameters={})) + messages = list(tool.invoke(session=sqlite_session, user_id="u1", tool_parameters={})) assert [m.type for m in messages] == [ToolInvokeMessage.MessageType.JSON, ToolInvokeMessage.MessageType.TEXT] # _invoke text path monkeypatch.setattr(tool, "validate_and_parse_response", lambda _: ParsedResponse("plain", False)) - messages = list(tool.invoke(session=MagicMock(), user_id="u1", tool_parameters={})) + messages = list(tool.invoke(session=sqlite_session, user_id="u1", tool_parameters={})) assert len(messages) == 1 assert messages[0].message.text == "plain" diff --git a/api/tests/unit_tests/core/tools/test_custom_tool_provider.py b/api/tests/unit_tests/core/tools/test_custom_tool_provider.py index 93ae217e24e..760a6b49657 100644 --- a/api/tests/unit_tests/core/tools/test_custom_tool_provider.py +++ b/api/tests/unit_tests/core/tools/test_custom_tool_provider.py @@ -1,17 +1,31 @@ +"""Tests for custom API tool providers with persisted provider lookup state.""" + from __future__ import annotations +import json +from dataclasses import dataclass from types import SimpleNamespace -from unittest.mock import Mock, patch +from typing import cast +from uuid import uuid4 import pytest +from sqlalchemy import delete +from sqlalchemy.orm import Session +from core.tools.custom_tool import provider as provider_module from core.tools.custom_tool.provider import ApiToolProviderController from core.tools.custom_tool.tool import ApiTool from core.tools.entities.tool_bundle import ApiToolBundle -from core.tools.entities.tool_entities import ApiProviderAuthType, ToolProviderType +from core.tools.entities.tool_entities import ApiProviderAuthType, ApiProviderSchemaType, ToolProviderType +from models.tools import ApiToolProvider -def _db_provider() -> SimpleNamespace: +@dataclass(frozen=True) +class _Database: + session: Session + + +def _db_provider() -> ApiToolProvider: bundle = ApiToolBundle( server_url="https://api.example.com/items", method="GET", @@ -21,18 +35,39 @@ def _db_provider() -> SimpleNamespace: author="author", openapi={"parameters": []}, ) - return SimpleNamespace( - id="provider-id", - tenant_id="tenant-1", - name="provider-a", - description="desc", - icon="icon.svg", - user=SimpleNamespace(name="Alice"), - tools=[bundle], + return cast( + ApiToolProvider, + SimpleNamespace( + id="provider-id", + tenant_id="tenant-1", + name="provider-a", + description="desc", + icon="icon.svg", + user=SimpleNamespace(name="Alice"), + tools=[bundle], + ), ) -def test_api_tool_provider_from_db_and_parse_tool_bundle(): +def _persist_provider(session: Session, *, tenant_id: str, name: str = "provider-a") -> ApiToolProvider: + bundle = _db_provider().tools[0] + provider = ApiToolProvider( + name=name, + icon="icon.svg", + schema="{}", + schema_type_str=ApiProviderSchemaType.OPENAPI, + user_id=str(uuid4()), + tenant_id=tenant_id, + description="desc", + tools_str=json.dumps([bundle.model_dump(mode="json")]), + credentials_str='{"auth_type":"none"}', + ) + session.add(provider) + session.commit() + return provider + + +def test_api_tool_provider_from_db_and_parse_tool_bundle() -> None: controller = ApiToolProviderController.from_db(_db_provider(), ApiProviderAuthType.API_KEY_HEADER) assert controller.provider_type == ToolProviderType.API assert any(c.name == "api_key_value" for c in controller.entity.credentials_schema) @@ -42,7 +77,7 @@ def test_api_tool_provider_from_db_and_parse_tool_bundle(): assert tool.entity.identity.provider == "provider-id" -def test_api_tool_provider_from_db_query_auth_and_none_auth(): +def test_api_tool_provider_from_db_query_auth_and_none_auth() -> None: query_controller = ApiToolProviderController.from_db(_db_provider(), ApiProviderAuthType.API_KEY_QUERY) assert any(c.name == "api_key_query_param" for c in query_controller.entity.credentials_schema) @@ -50,7 +85,9 @@ def test_api_tool_provider_from_db_query_auth_and_none_auth(): assert [c.name for c in none_controller.entity.credentials_schema] == ["auth_type"] -def test_api_tool_provider_load_get_tools_and_get_tool(): +def test_api_tool_provider_load_get_tools_and_get_tool( + monkeypatch: pytest.MonkeyPatch, sqlite_session: Session +) -> None: controller = ApiToolProviderController.from_db(_db_provider(), ApiProviderAuthType.NONE) loaded = controller.load_bundled_tools(_db_provider().tools) assert len(loaded) == 1 @@ -66,10 +103,17 @@ def test_api_tool_provider_load_get_tools_and_get_tool(): # Force DB fetch branch. controller.tools = [] - provider_with_tools = _db_provider() - with patch("core.tools.custom_tool.provider.db") as mock_db: - scalars_result = Mock() - scalars_result.all.return_value = [provider_with_tools] - mock_db.session.scalars.return_value = scalars_result - tools = controller.get_tools("tenant-1") + tenant_id = str(uuid4()) + provider_with_tools = _persist_provider(sqlite_session, tenant_id=tenant_id) + _persist_provider(sqlite_session, tenant_id=str(uuid4())) + controller.tenant_id = tenant_id + monkeypatch.setattr(provider_module, "db", _Database(session=sqlite_session)) + + tools = controller.get_tools(tenant_id) assert len(tools) == 1 + assert tools[0].entity.identity.provider == controller.provider_id + + sqlite_session.execute(delete(ApiToolProvider).where(ApiToolProvider.id == provider_with_tools.id)) + sqlite_session.commit() + controller.tools = [] + assert controller.get_tools(tenant_id) == [] diff --git a/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py b/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py index 44785f939ca..8b75d733452 100644 --- a/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py +++ b/api/tests/unit_tests/core/tools/utils/test_model_invocation_utils.py @@ -4,6 +4,7 @@ Covers success and error branches for ModelInvocationUtils, including InvokeModelError and invoke error mappings for InvokeAuthorizationError, InvokeBadRequestError, InvokeConnectionError, InvokeRateLimitError, and InvokeServerUnavailableError. Assumes mocked model instances and managers. +Invocation logging uses a real SQLite-backed SQLAlchemy session. """ from __future__ import annotations @@ -14,7 +15,10 @@ from typing import Any from unittest.mock import Mock, patch import pytest +from sqlalchemy import select +from sqlalchemy.orm import Session +from core.tools.entities.tool_entities import ToolProviderType from core.tools.utils.model_invocation_utils import InvokeModelError, ModelInvocationUtils from graphon.model_runtime.entities.model_entities import ModelPropertyKey from graphon.model_runtime.errors.invoke import ( @@ -24,6 +28,11 @@ from graphon.model_runtime.errors.invoke import ( InvokeRateLimitError, InvokeServerUnavailableError, ) +from models.tools import ToolModelInvoke + +TENANT_ID = "11111111-1111-1111-1111-111111111111" +USER_ID = "22222222-2222-2222-2222-222222222222" +CALLER_ID = "33333333-3333-3333-3333-333333333333" def _mock_model_instance(*, schema: dict[str, Any] | None = None) -> SimpleNamespace: @@ -80,7 +89,8 @@ def test_calculate_tokens_handles_missing_model(): mock_factory.assert_called_once_with(tenant_id="tenant", user_id=None) -def test_invoke_success_and_error_mappings(): +@pytest.mark.parametrize("sqlite_session", [(ToolModelInvoke,)], indirect=True) +def test_invoke_success_and_error_mappings(sqlite_session: Session): model_instance = _mock_model_instance(schema={ModelPropertyKey.CONTEXT_SIZE: 2048}) model_instance.invoke_llm.return_value = SimpleNamespace( message=SimpleNamespace(content="ok"), @@ -96,28 +106,37 @@ def test_invoke_success_and_error_mappings(): manager = Mock() manager.get_default_model_instance.return_value = model_instance - class _ToolModelInvoke: - def __init__(self, **kwargs): - self.__dict__.update(kwargs) + database = SimpleNamespace(session=sqlite_session) - db_mock = SimpleNamespace(session=Mock()) - - with patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory: - with patch("core.tools.utils.model_invocation_utils.ToolModelInvoke", _ToolModelInvoke): - with patch("core.tools.utils.model_invocation_utils.db", db_mock): - response = ModelInvocationUtils.invoke( - user_id="u1", - tenant_id="tenant", - tool_type="builtin", - tool_name="tool-a", - prompt_messages=[], - caller_user_id="caller-1", - ) + with ( + patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory, + patch("core.tools.utils.model_invocation_utils.db", database), + patch.object(sqlite_session, "add", wraps=sqlite_session.add) as mock_session_add, + patch.object(sqlite_session, "commit", wraps=sqlite_session.commit) as mock_session_commit, + ): + response = ModelInvocationUtils.invoke( + user_id=USER_ID, + tenant_id=TENANT_ID, + tool_type=ToolProviderType.BUILT_IN, + tool_name="tool-a", + prompt_messages=[], + caller_user_id=CALLER_ID, + ) assert response.message.content == "ok" - assert db_mock.session.add.call_count == 1 - assert db_mock.session.commit.call_count == 2 - mock_factory.assert_called_once_with(tenant_id="tenant", user_id="caller-1") + assert mock_session_add.call_count == 1 + assert mock_session_commit.call_count == 2 + assert not sqlite_session.in_transaction() + persisted = sqlite_session.scalar(select(ToolModelInvoke)) + assert persisted is not None + assert persisted.user_id == USER_ID + assert persisted.tenant_id == TENANT_ID + assert persisted.tool_type == ToolProviderType.BUILT_IN + assert persisted.model_response == "ok" + assert persisted.prompt_tokens == 5 + assert persisted.answer_tokens == 7 + assert persisted.total_price == Decimal("0.7000000") + mock_factory.assert_called_once_with(tenant_id=TENANT_ID, user_id=CALLER_ID) @pytest.mark.parametrize( @@ -139,27 +158,31 @@ def test_invoke_success_and_error_mappings(): "generic-error", ], ) -def test_invoke_error_mappings(exc, expected): +@pytest.mark.parametrize("sqlite_session", [(ToolModelInvoke,)], indirect=True) +def test_invoke_error_mappings(exc, expected, sqlite_session: Session): model_instance = _mock_model_instance(schema={ModelPropertyKey.CONTEXT_SIZE: 2048}) model_instance.invoke_llm.side_effect = exc manager = Mock() manager.get_default_model_instance.return_value = model_instance - class _ToolModelInvoke: - def __init__(self, **kwargs): - self.__dict__.update(kwargs) + database = SimpleNamespace(session=sqlite_session) - db_mock = SimpleNamespace(session=Mock()) - - with patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory: - with patch("core.tools.utils.model_invocation_utils.ToolModelInvoke", _ToolModelInvoke): - with patch("core.tools.utils.model_invocation_utils.db", db_mock): - with pytest.raises(InvokeModelError, match=expected): - ModelInvocationUtils.invoke( - user_id="u1", - tenant_id="tenant", - tool_type="builtin", - tool_name="tool-a", - prompt_messages=[], - ) - mock_factory.assert_called_once_with(tenant_id="tenant", user_id="u1") + with ( + patch("core.tools.utils.model_invocation_utils.ModelManager.for_tenant", return_value=manager) as mock_factory, + patch("core.tools.utils.model_invocation_utils.db", database), + ): + with pytest.raises(InvokeModelError, match=expected): + ModelInvocationUtils.invoke( + user_id=USER_ID, + tenant_id=TENANT_ID, + tool_type=ToolProviderType.BUILT_IN, + tool_name="tool-a", + prompt_messages=[], + ) + assert not sqlite_session.in_transaction() + persisted = sqlite_session.scalar(select(ToolModelInvoke)) + assert persisted is not None + assert persisted.model_response == "" + assert persisted.prompt_tokens == 5 + assert persisted.answer_tokens == 0 + mock_factory.assert_called_once_with(tenant_id=TENANT_ID, user_id=USER_ID)