mirror of
https://github.com/langgenius/dify.git
synced 2026-07-22 19:48:36 +08:00
test: use SQLite sessions in services core (#39111)
This commit is contained in:
parent
cb2b36f1aa
commit
b38caf3cdb
@ -1,9 +1,10 @@
|
||||
"""Unit tests for the Agent tool inner invoke service."""
|
||||
"""Unit tests for the Agent tool inner invoke service with SQLite-backed app lookup."""
|
||||
|
||||
from collections.abc import Generator
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.tools.entities.tool_entities import ToolInvokeMessage, ToolProviderType
|
||||
from core.tools.errors import (
|
||||
@ -12,19 +13,44 @@ from core.tools.errors import (
|
||||
ToolProviderCredentialValidationError,
|
||||
ToolProviderNotFoundError,
|
||||
)
|
||||
from models.enums import AppStatus
|
||||
from models.model import App, AppMode
|
||||
from services.agent_tool_inner_service import AgentToolInnerService
|
||||
from services.entities.agent_tool_inner import AgentToolInvokeRequest
|
||||
from services.errors.agent_tool_inner import AgentToolInnerServiceError
|
||||
|
||||
TENANT_ID = "11111111-1111-1111-1111-111111111111"
|
||||
OTHER_TENANT_ID = "22222222-2222-2222-2222-222222222222"
|
||||
USER_ID = "33333333-3333-3333-3333-333333333333"
|
||||
APP_ID = "44444444-4444-4444-4444-444444444444"
|
||||
|
||||
|
||||
def _persist_app(sqlite_session: Session, *, tenant_id: str = TENANT_ID) -> App:
|
||||
app = App(
|
||||
id=APP_ID,
|
||||
tenant_id=tenant_id,
|
||||
name="Test App",
|
||||
description="",
|
||||
mode=AppMode.CHAT,
|
||||
status=AppStatus.NORMAL,
|
||||
enable_site=False,
|
||||
enable_api=False,
|
||||
max_active_requests=None,
|
||||
)
|
||||
sqlite_session.add(app)
|
||||
sqlite_session.commit()
|
||||
sqlite_session.expunge_all()
|
||||
return app
|
||||
|
||||
|
||||
def _request() -> AgentToolInvokeRequest:
|
||||
return AgentToolInvokeRequest.model_validate(
|
||||
{
|
||||
"caller": {
|
||||
"tenant_id": "tenant-1",
|
||||
"user_id": "user-1",
|
||||
"tenant_id": TENANT_ID,
|
||||
"user_id": USER_ID,
|
||||
"user_from": "account",
|
||||
"app_id": "app-1",
|
||||
"app_id": APP_ID,
|
||||
"invoke_from": "service-api",
|
||||
"conversation_id": "conversation-1",
|
||||
"workflow_id": "workflow-1",
|
||||
@ -53,11 +79,10 @@ def _messages() -> Generator[ToolInvokeMessage, None, None]:
|
||||
)
|
||||
|
||||
|
||||
def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
|
||||
def test_invoke_uses_agent_tool_runtime_and_returns_observation(sqlite_session: Session) -> None:
|
||||
fake_tool = MagicMock()
|
||||
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
|
||||
session = MagicMock()
|
||||
session.get.return_value = fake_app
|
||||
_persist_app(sqlite_session)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@ -70,7 +95,7 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None:
|
||||
side_effect=lambda messages, **_kwargs: messages,
|
||||
),
|
||||
):
|
||||
response = AgentToolInnerService().invoke(_request(), session=session)
|
||||
response = AgentToolInnerService().invoke(_request(), session=sqlite_session)
|
||||
|
||||
assert response.observation == "ok"
|
||||
assert response.metadata == {
|
||||
@ -82,56 +107,58 @@ def test_invoke_uses_agent_tool_runtime_and_returns_observation() -> None:
|
||||
assert agent_tool.provider_type is ToolProviderType.PLUGIN
|
||||
assert agent_tool.tool_parameters == {"region": "us"}
|
||||
mock_invoke.assert_called_once()
|
||||
assert mock_invoke.call_args.kwargs["session"] is sqlite_session
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
|
||||
def test_invoke_raises_app_not_found_when_session_has_no_app() -> None:
|
||||
session = MagicMock()
|
||||
session.get.return_value = None
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
|
||||
def test_invoke_raises_app_not_found_when_session_has_no_app(sqlite_session: Session) -> None:
|
||||
with pytest.raises(AgentToolInnerServiceError) as exc_info:
|
||||
AgentToolInnerService().invoke(_request(), session=session)
|
||||
AgentToolInnerService().invoke(_request(), session=sqlite_session)
|
||||
|
||||
assert exc_info.value.error_code == "app_not_found"
|
||||
assert exc_info.value.status_code == 404
|
||||
assert exc_info.value.description == "App not found."
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
|
||||
def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant() -> None:
|
||||
fake_app = MagicMock(id="app-1", tenant_id="tenant-2")
|
||||
session = MagicMock()
|
||||
session.get.return_value = fake_app
|
||||
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
|
||||
def test_invoke_raises_app_tenant_mismatch_when_app_belongs_to_other_tenant(sqlite_session: Session) -> None:
|
||||
_persist_app(sqlite_session, tenant_id=OTHER_TENANT_ID)
|
||||
|
||||
with pytest.raises(AgentToolInnerServiceError) as exc_info:
|
||||
AgentToolInnerService().invoke(_request(), session=session)
|
||||
AgentToolInnerService().invoke(_request(), session=sqlite_session)
|
||||
|
||||
assert exc_info.value.error_code == "app_tenant_mismatch"
|
||||
assert exc_info.value.status_code == 403
|
||||
assert exc_info.value.description == "App does not belong to the caller tenant."
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
|
||||
def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_code() -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
|
||||
def test_invoke_maps_tool_runtime_app_not_found_value_error_to_specific_error_code(
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
fake_tool = MagicMock()
|
||||
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
|
||||
session = MagicMock()
|
||||
session.get.return_value = fake_app
|
||||
_persist_app(sqlite_session)
|
||||
|
||||
with (
|
||||
patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", return_value=fake_tool),
|
||||
patch("services.agent_tool_inner_service.ToolEngine.generic_invoke", side_effect=ValueError("app not found")),
|
||||
):
|
||||
with pytest.raises(AgentToolInnerServiceError) as exc_info:
|
||||
AgentToolInnerService().invoke(_request(), session=session)
|
||||
AgentToolInnerService().invoke(_request(), session=sqlite_session)
|
||||
|
||||
assert exc_info.value.error_code == "app_not_found"
|
||||
assert exc_info.value.status_code == 404
|
||||
assert exc_info.value.description == "App not found."
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
|
||||
def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
|
||||
def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper(sqlite_session: Session) -> None:
|
||||
fake_tool = MagicMock()
|
||||
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
|
||||
session = MagicMock()
|
||||
session.get.return_value = fake_app
|
||||
_persist_app(sqlite_session)
|
||||
|
||||
with (
|
||||
patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", return_value=fake_tool),
|
||||
@ -141,9 +168,10 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N
|
||||
),
|
||||
):
|
||||
with pytest.raises(AgentToolInnerServiceError) as exc_info:
|
||||
AgentToolInnerService().invoke(_request(), session=session)
|
||||
AgentToolInnerService().invoke(_request(), session=sqlite_session)
|
||||
|
||||
assert exc_info.value.error_code == "agent_tool_invoke_failed"
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@ -154,13 +182,17 @@ def test_invoke_maps_tool_invoke_error_without_private_tool_engine_helper() -> N
|
||||
(ToolParameterValidationError("query is required"), "tool_parameters_invalid"),
|
||||
],
|
||||
)
|
||||
def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(error: Exception, expected_code: str) -> None:
|
||||
fake_app = MagicMock(id="app-1", tenant_id="tenant-1")
|
||||
session = MagicMock()
|
||||
session.get.return_value = fake_app
|
||||
@pytest.mark.parametrize("sqlite_session", [(App,)], indirect=True)
|
||||
def test_invoke_maps_runtime_lookup_errors_to_service_error_codes(
|
||||
error: Exception,
|
||||
expected_code: str,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
_persist_app(sqlite_session)
|
||||
|
||||
with patch("services.agent_tool_inner_service.ToolManager.get_agent_tool_runtime", side_effect=error):
|
||||
with pytest.raises(AgentToolInnerServiceError) as exc_info:
|
||||
AgentToolInnerService().invoke(_request(), session=session)
|
||||
AgentToolInnerService().invoke(_request(), session=sqlite_session)
|
||||
|
||||
assert exc_info.value.error_code == expected_code
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
@ -1419,6 +1419,8 @@ class TestWorkflowService:
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("sqlite_session")
|
||||
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
|
||||
class TestWorkflowServiceCredentialValidation:
|
||||
"""
|
||||
Tests for the private credential-validation helpers on WorkflowService.
|
||||
@ -1444,7 +1446,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
# --- _validate_workflow_credentials: tool node (with credential_id) ---
|
||||
|
||||
def test_validate_workflow_credentials_should_check_tool_credential_when_credential_id_present(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
nodes = [
|
||||
@ -1462,11 +1464,11 @@ class TestWorkflowServiceCredentialValidation:
|
||||
# Act + Assert
|
||||
with patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check:
|
||||
# Should not raise; mock allows the call
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
mock_check.assert_called_once()
|
||||
|
||||
def test_validate_workflow_credentials_should_check_default_credential_when_no_credential_id(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
nodes = [
|
||||
@ -1483,14 +1485,13 @@ class TestWorkflowServiceCredentialValidation:
|
||||
|
||||
# Act
|
||||
with patch.object(service, "_check_default_tool_credential") as mock_default:
|
||||
session = MagicMock()
|
||||
service._validate_workflow_credentials(workflow, session=session)
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
# Assert
|
||||
mock_default.assert_called_once_with("tenant-1", "my-provider", session=session)
|
||||
mock_default.assert_called_once_with("tenant-1", "my-provider", session=sqlite_session)
|
||||
|
||||
def test_validate_workflow_credentials_should_skip_tool_node_without_provider(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""Tool nodes without a provider_id should be silently skipped."""
|
||||
# Arrange
|
||||
@ -1499,11 +1500,11 @@ class TestWorkflowServiceCredentialValidation:
|
||||
|
||||
# Act + Assert (no error raised)
|
||||
with patch.object(service, "_check_default_tool_credential") as mock_default:
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
mock_default.assert_not_called()
|
||||
|
||||
def test_validate_workflow_credentials_should_validate_llm_node_with_model_config(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
nodes = [
|
||||
@ -1522,13 +1523,13 @@ class TestWorkflowServiceCredentialValidation:
|
||||
patch.object(service, "_validate_llm_model_config") as mock_llm,
|
||||
patch.object(service, "_validate_load_balancing_credentials"),
|
||||
):
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
# Assert
|
||||
mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4")
|
||||
|
||||
def test_validate_workflow_credentials_should_raise_for_llm_node_missing_model(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""LLM nodes without provider AND name should raise ValueError."""
|
||||
# Arrange
|
||||
@ -1542,10 +1543,10 @@ class TestWorkflowServiceCredentialValidation:
|
||||
|
||||
# Act + Assert
|
||||
with pytest.raises(ValueError, match="Missing provider or model configuration"):
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
def test_validate_workflow_credentials_should_wrap_unexpected_exception_in_value_error(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""Non-ValueError exceptions from validation must be re-raised as ValueError."""
|
||||
# Arrange
|
||||
@ -1563,9 +1564,11 @@ class TestWorkflowServiceCredentialValidation:
|
||||
# Act + Assert
|
||||
with patch.object(service, "_validate_llm_model_config", side_effect=RuntimeError("boom")):
|
||||
with pytest.raises(ValueError, match="boom"):
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
def test_validate_workflow_credentials_should_validate_agent_node_model(self, service: WorkflowService) -> None:
|
||||
def test_validate_workflow_credentials_should_validate_agent_node_model(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
nodes = [
|
||||
{
|
||||
@ -1586,12 +1589,14 @@ class TestWorkflowServiceCredentialValidation:
|
||||
patch.object(service, "_validate_llm_model_config") as mock_llm,
|
||||
patch.object(service, "_validate_load_balancing_credentials"),
|
||||
):
|
||||
service._validate_workflow_credentials(workflow, session=MagicMock())
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
# Assert
|
||||
mock_llm.assert_called_once_with("tenant-1", "openai", "gpt-4")
|
||||
|
||||
def test_validate_workflow_credentials_should_validate_agent_tools(self, service: WorkflowService) -> None:
|
||||
def test_validate_workflow_credentials_should_validate_agent_tools(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""Each agent tool with a provider should be checked for credential compliance."""
|
||||
# Arrange
|
||||
nodes = [
|
||||
@ -1618,12 +1623,11 @@ class TestWorkflowServiceCredentialValidation:
|
||||
patch("core.helper.credential_utils.check_credential_policy_compliance") as mock_check,
|
||||
patch.object(service, "_check_default_tool_credential") as mock_default,
|
||||
):
|
||||
session = MagicMock()
|
||||
service._validate_workflow_credentials(workflow, session=session)
|
||||
service._validate_workflow_credentials(workflow, session=sqlite_session)
|
||||
|
||||
# Assert
|
||||
mock_check.assert_called_once() # provider-a has credential_id
|
||||
mock_default.assert_called_once_with("tenant-1", "provider-b", session=session)
|
||||
mock_default.assert_called_once_with("tenant-1", "provider-b", session=sqlite_session)
|
||||
|
||||
# --- _validate_llm_model_config ---
|
||||
|
||||
@ -1676,14 +1680,12 @@ class TestWorkflowServiceCredentialValidation:
|
||||
|
||||
# --- _check_default_tool_credential ---
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
|
||||
def test_check_default_tool_credential_should_silently_pass_when_no_provider_found(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""Missing BuiltinToolProvider → plugin requires no credentials → no error."""
|
||||
service._check_default_tool_credential("tenant-1", "some-provider", session=sqlite_session)
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(BuiltinToolProvider,)], indirect=True)
|
||||
def test_check_default_tool_credential_should_raise_when_compliance_fails(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
@ -1746,7 +1748,9 @@ class TestWorkflowServiceCredentialValidation:
|
||||
|
||||
# --- _get_load_balancing_configs ---
|
||||
|
||||
def test_get_load_balancing_configs_should_return_empty_list_on_exception(self, service: WorkflowService) -> None:
|
||||
def test_get_load_balancing_configs_should_return_empty_list_on_exception(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""Any exception during LB config retrieval should return an empty list."""
|
||||
# Arrange
|
||||
with patch(
|
||||
@ -1754,12 +1758,14 @@ class TestWorkflowServiceCredentialValidation:
|
||||
side_effect=RuntimeError("fail"),
|
||||
):
|
||||
# Act
|
||||
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock())
|
||||
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=sqlite_session)
|
||||
|
||||
# Assert
|
||||
assert result == []
|
||||
|
||||
def test_get_load_balancing_configs_should_merge_predefined_and_custom(self, service: WorkflowService) -> None:
|
||||
def test_get_load_balancing_configs_should_merge_predefined_and_custom(
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
predefined = [{"credential_id": "cred-a"}, {"credential_id": None}]
|
||||
custom = [{"credential_id": "cred-b"}]
|
||||
@ -1771,7 +1777,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
],
|
||||
):
|
||||
# Act
|
||||
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=MagicMock())
|
||||
result = service._get_load_balancing_configs("tenant-1", "openai", "gpt-4", session=sqlite_session)
|
||||
|
||||
# Assert — only entries with a credential_id should be returned
|
||||
assert len(result) == 2
|
||||
@ -1780,7 +1786,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
# --- _validate_load_balancing_credentials ---
|
||||
|
||||
def test_validate_load_balancing_credentials_should_skip_when_no_model_config(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
"""Missing provider or model in node_data should be a no-op."""
|
||||
# Arrange
|
||||
@ -1788,10 +1794,10 @@ class TestWorkflowServiceCredentialValidation:
|
||||
node_data: dict[str, Any] = {} # no model key
|
||||
|
||||
# Act + Assert (no error expected)
|
||||
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock())
|
||||
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=sqlite_session)
|
||||
|
||||
def test_validate_load_balancing_credentials_should_skip_when_lb_not_enabled(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
workflow = self._make_workflow([])
|
||||
@ -1799,10 +1805,10 @@ class TestWorkflowServiceCredentialValidation:
|
||||
|
||||
# Act + Assert (no error expected)
|
||||
with patch.object(service, "_is_load_balancing_enabled", return_value=False):
|
||||
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock())
|
||||
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=sqlite_session)
|
||||
|
||||
def test_validate_load_balancing_credentials_should_raise_when_compliance_fails(
|
||||
self, service: WorkflowService
|
||||
self, service: WorkflowService, sqlite_session: Session
|
||||
) -> None:
|
||||
# Arrange
|
||||
workflow = self._make_workflow([])
|
||||
@ -1819,7 +1825,7 @@ class TestWorkflowServiceCredentialValidation:
|
||||
),
|
||||
):
|
||||
with pytest.raises(ValueError, match="Invalid load balancing credentials"):
|
||||
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=MagicMock())
|
||||
service._validate_load_balancing_credentials(workflow, node_data, "node-1", session=sqlite_session)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
|
||||
Loading…
Reference in New Issue
Block a user