fix(api): isolate side-effect database writes (#37895)

Co-authored-by: FFXN <31929997+FFXN@users.noreply.github.com>
This commit is contained in:
linhongkuan 2026-06-29 14:20:34 +08:00 committed by GitHub
parent fc16fcba36
commit d8f3be4bcd
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 114 additions and 63 deletions

View File

@ -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."
)

View File

@ -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,

View File

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

View File

@ -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:

View File

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

View File

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