From b38caf3cdb6ff383e236e196fb8b3face2e0c9de Mon Sep 17 00:00:00 2001 From: Asuka Minato Date: Wed, 22 Jul 2026 11:15:01 +0900 Subject: [PATCH] test: use SQLite sessions in services core (#39111) --- .../services/test_agent_tool_inner_service.py | 100 ++++++++++++------ .../services/test_workflow_service.py | 70 ++++++------ 2 files changed, 104 insertions(+), 66 deletions(-) diff --git a/api/tests/unit_tests/services/test_agent_tool_inner_service.py b/api/tests/unit_tests/services/test_agent_tool_inner_service.py index 61049d29e9e..8f222b33093 100644 --- a/api/tests/unit_tests/services/test_agent_tool_inner_service.py +++ b/api/tests/unit_tests/services/test_agent_tool_inner_service.py @@ -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() diff --git a/api/tests/unit_tests/services/test_workflow_service.py b/api/tests/unit_tests/services/test_workflow_service.py index 37450cb253a..b2e0e4129c9 100644 --- a/api/tests/unit_tests/services/test_workflow_service.py +++ b/api/tests/unit_tests/services/test_workflow_service.py @@ -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) # ===========================================================================