mirror of
https://github.com/langgenius/dify.git
synced 2026-09-21 05:11:22 +08:00
test: use SQLite sessions in core tools (#39108)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user