test: use SQLite sessions in core tools (#39108)

This commit is contained in:
Asuka Minato
2026-08-06 03:57:33 +00:00
committed by GitHub
parent 6add64a44d
commit 2e65ab9d1e
4 changed files with 145 additions and 73 deletions
@@ -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
@@ -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"
@@ -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) == []
@@ -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)