From d8f3be4bcdb6e80e42b2c8a45a4bc07743a22d91 Mon Sep 17 00:00:00 2001 From: linhongkuan Date: Mon, 29 Jun 2026 14:20:34 +0800 Subject: [PATCH] fix(api): isolate side-effect database writes (#37895) Co-authored-by: FFXN <31929997+FFXN@users.noreply.github.com> --- api/controllers/service_api/wraps.py | 5 +- api/core/rag/retrieval/dataset_retrieval.py | 18 ++++- api/core/tools/tool_engine.py | 75 ++++++++++--------- .../controllers/service_api/test_wraps.py | 14 +++- .../rag/retrieval/test_dataset_retrieval.py | 36 +++++++-- .../unit_tests/core/tools/test_tool_engine.py | 29 ++++--- 6 files changed, 114 insertions(+), 63 deletions(-) diff --git a/api/controllers/service_api/wraps.py b/api/controllers/service_api/wraps.py index 32e95b481f8..8cb339a1491 100644 --- a/api/controllers/service_api/wraps.py +++ b/api/controllers/service_api/wraps.py @@ -12,6 +12,7 @@ from flask_restx import Resource from flask_restx.utils import merge from pydantic import BaseModel from sqlalchemy import select +from sqlalchemy.orm import sessionmaker from werkzeug.exceptions import Forbidden, NotFound, Unauthorized from configs import dify_config @@ -269,8 +270,8 @@ def cloud_edition_billing_rate_limit_check[**P, R]( subscription_plan=knowledge_rate_limit.subscription_plan, operation="knowledge", ) - db.session.add(rate_limit_log) - db.session.commit() + with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: + session.add(rate_limit_log) raise Forbidden( "Sorry, you have reached the knowledge base request rate limit of your subscription." ) diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index 474c9f90c78..c8f1210ee8c 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -1030,6 +1030,10 @@ class DatasetRetrieval: ): """ Persist dataset query audit rows for retrieval requests. + + Query audit logging is a side effect of retrieval. Keep it in an + independent transaction so failures or commits here do not affect the + request/workflow transaction that called the retriever. """ if not query and not attachment_ids: return @@ -1041,6 +1045,9 @@ class DatasetRetrieval: app_id, ) return + created_by_role = self._resolve_creator_user_role(user_from) + if created_by_role is None: + return dataset_queries = [] for dataset_id in dataset_ids: contents = [] @@ -1055,13 +1062,16 @@ class DatasetRetrieval: content=json.dumps(contents), source=DatasetQuerySource.APP, source_app_id=app_id, - created_by_role=CreatorUserRole(user_from), + created_by_role=created_by_role, created_by=created_by, ) dataset_queries.append(dataset_query) - if dataset_queries: - db.session.add_all(dataset_queries) - db.session.commit() + + if not dataset_queries: + return + + with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: + session.add_all(dataset_queries) def _retriever( self, diff --git a/api/core/tools/tool_engine.py b/api/core/tools/tool_engine.py index 16fe15d5aed..6a2eb207641 100644 --- a/api/core/tools/tool_engine.py +++ b/api/core/tools/tool_engine.py @@ -7,6 +7,7 @@ from datetime import UTC, datetime from mimetypes import guess_type from typing import Any, Union, cast +from sqlalchemy.orm import sessionmaker from yarl import URL from core.app.entities.app_invoke_entities import InvokeFrom @@ -338,47 +339,49 @@ class ToolEngine: user_id: str, ) -> list[str]: """ - Create message file + Create message files produced by a tool call. + + Tool file persistence is a side effect of agent execution. Use an + independent transaction so this helper never commits or closes the + caller's request-scoped session. :return: message file ids """ result = [] - for message in tool_messages: - if "image" in message.mimetype: - file_type = FileType.IMAGE - elif "video" in message.mimetype: - file_type = FileType.VIDEO - elif "audio" in message.mimetype: - file_type = FileType.AUDIO - elif "text" in message.mimetype or "pdf" in message.mimetype: - file_type = FileType.DOCUMENT - else: - file_type = FileType.CUSTOM + with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session: + for message in tool_messages: + # extract tool file id from url + tool_file_id = message.url.split("/")[-1].split(".")[0] + message_file = MessageFile( + message_id=agent_message.id, + type=ToolEngine._resolve_tool_file_type(message), + transfer_method=FileTransferMethod.TOOL_FILE, + belongs_to=MessageFileBelongsTo.ASSISTANT, + url=message.url, + upload_file_id=tool_file_id, + created_by_role=( + CreatorUserRole.ACCOUNT + if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER} + else CreatorUserRole.END_USER + ), + created_by=user_id, + ) - # extract tool file id from url - tool_file_id = message.url.split("/")[-1].split(".")[0] - message_file = MessageFile( - message_id=agent_message.id, - type=file_type, - transfer_method=FileTransferMethod.TOOL_FILE, - belongs_to=MessageFileBelongsTo.ASSISTANT, - url=message.url, - upload_file_id=tool_file_id, - created_by_role=( - CreatorUserRole.ACCOUNT - if invoke_from in {InvokeFrom.EXPLORE, InvokeFrom.DEBUGGER} - else CreatorUserRole.END_USER - ), - created_by=user_id, - ) - - db.session.add(message_file) - db.session.commit() - db.session.refresh(message_file) - - result.append(message_file.id) - - db.session.close() + session.add(message_file) + result.append(message_file.id) return result + + @staticmethod + def _resolve_tool_file_type(message: ToolInvokeMessageBinary) -> FileType: + if "image" in message.mimetype: + return FileType.IMAGE + elif "video" in message.mimetype: + return FileType.VIDEO + elif "audio" in message.mimetype: + return FileType.AUDIO + elif "text" in message.mimetype or "pdf" in message.mimetype: + return FileType.DOCUMENT + else: + return FileType.CUSTOM diff --git a/api/tests/unit_tests/controllers/service_api/test_wraps.py b/api/tests/unit_tests/controllers/service_api/test_wraps.py index 0b5c5d95b69..5857e5c639a 100644 --- a/api/tests/unit_tests/controllers/service_api/test_wraps.py +++ b/api/tests/unit_tests/controllers/service_api/test_wraps.py @@ -3,7 +3,7 @@ Unit tests for Service API wraps (authentication decorators) """ import uuid -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest from flask import Flask @@ -469,7 +469,10 @@ class TestCloudEditionBillingRateLimitCheck: @patch("controllers.service_api.wraps.validate_and_get_api_token") @patch("controllers.service_api.wraps.FeatureService.get_knowledge_rate_limit") @patch("controllers.service_api.wraps.db") - def test_rejects_over_rate_limit(self, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask): + @patch("controllers.service_api.wraps.sessionmaker") + def test_rejects_over_rate_limit( + self, mock_sessionmaker, mock_db, mock_get_rate_limit, mock_validate_token, app: Flask + ): """Test that Forbidden is raised when over rate limit.""" # Arrange mock_validate_token.return_value = Mock(tenant_id="tenant123") @@ -479,6 +482,10 @@ class TestCloudEditionBillingRateLimitCheck: mock_rate_limit.limit = 10 mock_rate_limit.subscription_plan = "pro" mock_get_rate_limit.return_value = mock_rate_limit + rate_limit_log_session = MagicMock() + session_factory = MagicMock() + session_factory.begin.return_value.__enter__.return_value = rate_limit_log_session + mock_sessionmaker.return_value = session_factory with patch("controllers.service_api.wraps.redis_client") as mock_redis: mock_redis.zcard.return_value = 15 # Over limit @@ -492,6 +499,9 @@ class TestCloudEditionBillingRateLimitCheck: with pytest.raises(Forbidden) as exc_info: knowledge_request() assert "rate limit" in str(exc_info.value) + mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) + rate_limit_log_session.add.assert_called_once() + mock_db.session.commit.assert_not_called() class TestValidateDatasetToken: diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py index 46cf0f7ac49..f9a58069ab3 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py @@ -3888,7 +3888,17 @@ class TestDatasetRetrievalAdditionalHelpers: trace_manager.add_trace_task.assert_not_called() def test_on_query(self, retrieval: DatasetRetrieval) -> None: - with patch("core.rag.retrieval.dataset_retrieval.db.session") as mock_session: + db_mock = Mock() + audit_session = MagicMock() + session_factory = MagicMock() + session_factory.begin.return_value.__enter__.return_value = audit_session + + with ( + patch("core.rag.retrieval.dataset_retrieval.db", db_mock), + patch( + "core.rag.retrieval.dataset_retrieval.sessionmaker", return_value=session_factory + ) as sessionmaker_mock, + ): retrieval._on_query( query=None, attachment_ids=None, @@ -3897,7 +3907,7 @@ class TestDatasetRetrievalAdditionalHelpers: user_from="account", user_id="u1", ) - mock_session.add_all.assert_not_called() + audit_session.add_all.assert_not_called() retrieval._on_query( query="python", @@ -3907,11 +3917,22 @@ class TestDatasetRetrievalAdditionalHelpers: user_from="account", user_id="u1", ) - mock_session.add_all.assert_called() - mock_session.commit.assert_called() + sessionmaker_mock.assert_called_once_with(bind=db_mock.engine, expire_on_commit=False) + audit_session.add_all.assert_called_once() + added_queries = audit_session.add_all.call_args.args[0] + assert len(added_queries) == 2 + db_mock.session.commit.assert_not_called() def test_on_query_normalizes_workflow_end_user_role(self, retrieval: DatasetRetrieval) -> None: - with patch("core.rag.retrieval.dataset_retrieval.db.session") as mock_session: + db_mock = Mock() + audit_session = MagicMock() + session_factory = MagicMock() + session_factory.begin.return_value.__enter__.return_value = audit_session + + with ( + patch("core.rag.retrieval.dataset_retrieval.db", db_mock), + patch("core.rag.retrieval.dataset_retrieval.sessionmaker", return_value=session_factory), + ): retrieval._on_query( query="python", attachment_ids=None, @@ -3921,12 +3942,11 @@ class TestDatasetRetrievalAdditionalHelpers: user_id="u1", ) - mock_session.add_all.assert_called_once() - added_queries = mock_session.add_all.call_args.args[0] + audit_session.add_all.assert_called_once() + added_queries = audit_session.add_all.call_args.args[0] assert len(added_queries) == 1 assert added_queries[0].created_by_role == CreatorUserRole.END_USER - mock_session.commit.assert_called_once() def test_handle_invoke_result(self, retrieval: DatasetRetrieval) -> None: usage = LLMUsage.empty_usage() diff --git a/api/tests/unit_tests/core/tools/test_tool_engine.py b/api/tests/unit_tests/core/tools/test_tool_engine.py index cd16557ef64..8118f40e8f7 100644 --- a/api/tests/unit_tests/core/tools/test_tool_engine.py +++ b/api/tests/unit_tests/core/tools/test_tool_engine.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections.abc import Generator from types import SimpleNamespace from typing import Any -from unittest.mock import Mock, patch +from unittest.mock import MagicMock, Mock, patch import pytest @@ -131,18 +131,25 @@ def test_create_message_files_and_invoke_generator(): created.append(obj) return obj - with patch("core.tools.tool_engine.MessageFile", side_effect=_message_file_factory): - with patch("core.tools.tool_engine.db") as mock_db: - ids = ToolEngine._create_message_files( - tool_messages=binaries, - agent_message=SimpleNamespace(id="msg-1"), - invoke_from=InvokeFrom.DEBUGGER, - user_id="user-1", - ) + file_session = MagicMock() + session_factory = MagicMock() + session_factory.begin.return_value.__enter__.return_value = file_session + with ( + patch("core.tools.tool_engine.MessageFile", side_effect=_message_file_factory), + patch("core.tools.tool_engine.db") as mock_db, + patch("core.tools.tool_engine.sessionmaker", return_value=session_factory) as mock_sessionmaker, + ): + ids = ToolEngine._create_message_files( + tool_messages=binaries, + agent_message=SimpleNamespace(id="msg-1"), + invoke_from=InvokeFrom.DEBUGGER, + user_id="user-1", + ) assert ids == ["mf-1", "mf-2"] - assert mock_db.session.add.call_count == 2 - mock_db.session.close.assert_called_once() + mock_sessionmaker.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False) + assert file_session.add.call_count == 2 + mock_db.session.close.assert_not_called() tool = _build_tool() invoked = list(ToolEngine._invoke(tool, {"a": 1}, user_id="u"))