test: use SQLite sessions in services core (#39111)

This commit is contained in:
Asuka Minato 2026-07-22 11:15:01 +09:00 committed by GitHub
parent cb2b36f1aa
commit b38caf3cdb
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 104 additions and 66 deletions

View File

@ -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()

View File

@ -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)
# ===========================================================================