fix(api): isolate side-effect session writes in multimodal and RAG handlers (#38210)

Co-authored-by: FFXN <31929997+FFXN@users.noreply.github.com>
This commit is contained in:
Pranav Agarwal 2026-07-06 10:47:12 +05:30 committed by GitHub
parent 93eb6d32b5
commit d9c99daf29
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 286 additions and 231 deletions

View File

@ -5,6 +5,8 @@ from collections.abc import Generator, Mapping, Sequence
from mimetypes import guess_extension from mimetypes import guess_extension
from typing import TYPE_CHECKING, Any, Union from typing import TYPE_CHECKING, Any, Union
from sqlalchemy.orm import sessionmaker
from core.app.app_config.entities import ExternalDataVariableEntity, PromptTemplateEntity from core.app.app_config.entities import ExternalDataVariableEntity, PromptTemplateEntity
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.apps.exc import GenerateTaskStoppedError from core.app.apps.exc import GenerateTaskStoppedError
@ -423,7 +425,9 @@ class AppRunner:
_logger.exception("Failed to save image file") _logger.exception("Failed to save image file")
return return
# Create MessageFile record # Create MessageFile record.
# Use an independent session so this side-effect write does not
# commit or close the caller's request-scoped session.
message_file = MessageFile( message_file = MessageFile(
message_id=message_id, message_id=message_id,
type=FileType.IMAGE, type=FileType.IMAGE,
@ -437,9 +441,8 @@ class AppRunner:
created_by=user_id, created_by=user_id,
) )
db.session.add(message_file) with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
db.session.commit() session.add(message_file)
db.session.refresh(message_file)
# Publish QueueMessageFileEvent # Publish QueueMessageFileEvent
queue_manager.publish( queue_manager.publish(

View File

@ -2,7 +2,7 @@ import logging
from collections.abc import Sequence from collections.abc import Sequence
from sqlalchemy import select, update from sqlalchemy import select, update
from sqlalchemy.orm import scoped_session from sqlalchemy.orm import Session, scoped_session, sessionmaker
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
@ -10,6 +10,7 @@ from core.app.entities.queue_entities import QueueRetrieverResourcesEvent
from core.rag.entities import RetrievalSourceMetadata from core.rag.entities import RetrievalSourceMetadata
from core.rag.index_processor.constant.index_type import IndexStructureType from core.rag.index_processor.constant.index_type import IndexStructureType
from core.rag.models.document import Document from core.rag.models.document import Document
from extensions.ext_database import db
from models.dataset import ChildChunk, DatasetQuery, DocumentSegment from models.dataset import ChildChunk, DatasetQuery, DocumentSegment
from models.dataset import Document as DatasetDocument from models.dataset import Document as DatasetDocument
from models.enums import CreatorUserRole, DatasetQuerySource from models.enums import CreatorUserRole, DatasetQuerySource
@ -46,47 +47,52 @@ class DatasetIndexToolCallbackHandler:
created_by=self._user_id, created_by=self._user_id,
) )
session.add(dataset_query) # Use an independent session so this audit-log side effect does
session.commit() # not commit or close the caller's request-scoped session.
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as independent_session:
independent_session.add(dataset_query)
def on_tool_end(self, documents: list[Document], session: scoped_session): def on_tool_end(self, documents: list[Document], session: scoped_session):
"""Handle tool end.""" """Handle tool end."""
for document in documents: # Use an independent session so hit-count updates do not
if document.metadata is not None: # interfere with the caller's request-scoped session.
document_id = document.metadata["document_id"] with Session(db.engine, expire_on_commit=False) as independent_session:
dataset_document_stmt = select(DatasetDocument).where(DatasetDocument.id == document_id) for document in documents:
dataset_document = session.scalar(dataset_document_stmt) if document.metadata is not None:
if not dataset_document: document_id = document.metadata["document_id"]
_logger.warning( dataset_document_stmt = select(DatasetDocument).where(DatasetDocument.id == document_id)
"Expected DatasetDocument record to exist, but none was found, document_id=%s", dataset_document = independent_session.scalar(dataset_document_stmt)
document_id, if not dataset_document:
) _logger.warning(
continue "Expected DatasetDocument record to exist, but none was found, document_id=%s",
if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX: document_id,
child_chunk_stmt = select(ChildChunk).where(
ChildChunk.index_node_id == document.metadata["doc_id"],
ChildChunk.dataset_id == dataset_document.dataset_id,
ChildChunk.document_id == dataset_document.id,
)
child_chunk = session.scalar(child_chunk_stmt)
if child_chunk:
session.execute(
update(DocumentSegment)
.where(DocumentSegment.id == child_chunk.segment_id)
.values(hit_count=DocumentSegment.hit_count + 1)
) )
else: continue
conditions = [DocumentSegment.index_node_id == document.metadata["doc_id"]] if dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX:
child_chunk_stmt = select(ChildChunk).where(
ChildChunk.index_node_id == document.metadata["doc_id"],
ChildChunk.dataset_id == dataset_document.dataset_id,
ChildChunk.document_id == dataset_document.id,
)
child_chunk = independent_session.scalar(child_chunk_stmt)
if child_chunk:
independent_session.execute(
update(DocumentSegment)
.where(DocumentSegment.id == child_chunk.segment_id)
.values(hit_count=DocumentSegment.hit_count + 1)
)
else:
conditions = [DocumentSegment.index_node_id == document.metadata["doc_id"]]
if "dataset_id" in document.metadata: if "dataset_id" in document.metadata:
conditions.append(DocumentSegment.dataset_id == document.metadata["dataset_id"]) conditions.append(DocumentSegment.dataset_id == document.metadata["dataset_id"])
# add hit count to document segment # add hit count to document segment
session.execute( independent_session.execute(
update(DocumentSegment).where(*conditions).values(hit_count=DocumentSegment.hit_count + 1) update(DocumentSegment).where(*conditions).values(hit_count=DocumentSegment.hit_count + 1)
) )
session.commit() independent_session.commit()
# TODO(-LAN-): Improve type check # TODO(-LAN-): Improve type check
def return_retriever_resource_info(self, resource: Sequence[RetrievalSourceMetadata]): def return_retriever_resource_info(self, resource: Sequence[RetrievalSourceMetadata]):

View File

@ -5,7 +5,6 @@ from uuid import uuid4
import pytest import pytest
from core.app.apps.base_app_queue_manager import PublishFrom
from core.app.apps.base_app_runner import AppRunner from core.app.apps.base_app_runner import AppRunner
from core.app.entities.app_invoke_entities import InvokeFrom from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.queue_entities import QueueMessageFileEvent from core.app.entities.queue_entities import QueueMessageFileEvent
@ -81,59 +80,55 @@ class TestBaseAppRunnerMultimodal:
# Setup mock message file # Setup mock message file
mock_msg_file_class.return_value = mock_message_file mock_msg_file_class.return_value = mock_message_file
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: file_session = MagicMock()
mock_session.add = MagicMock() mock_session_factory = MagicMock()
mock_session.commit = MagicMock() mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session)
mock_session.refresh = MagicMock() mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False)
# Act with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory) as mock_sm:
# Create a mock runner with the method bound with patch("core.app.apps.base_app_runner.db") as mock_db:
runner = MagicMock() # Act
runner = MagicMock()
method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(
runner, *args, **kwargs
)
method = AppRunner._handle_multimodal_image_content runner._handle_multimodal_image_content(
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) content=content,
message_id=mock_message_id,
user_id=mock_user_id,
tenant_id=mock_tenant_id,
queue_manager=mock_queue_manager,
)
runner._handle_multimodal_image_content( # Assert
content=content, mock_mgr.create_file_by_url.assert_called_once_with(
message_id=mock_message_id, user_id=mock_user_id,
user_id=mock_user_id, tenant_id=mock_tenant_id,
tenant_id=mock_tenant_id, file_url=image_url,
queue_manager=mock_queue_manager, conversation_id=None,
) )
# Assert mock_msg_file_class.assert_called_once()
# Verify tool file was created from URL call_kwargs = mock_msg_file_class.call_args[1]
mock_mgr.create_file_by_url.assert_called_once_with( assert call_kwargs["message_id"] == mock_message_id
user_id=mock_user_id, assert call_kwargs["type"] == FileType.IMAGE
tenant_id=mock_tenant_id, assert call_kwargs["transfer_method"] == FileTransferMethod.TOOL_FILE
file_url=image_url, assert call_kwargs["belongs_to"] == "assistant"
conversation_id=None, assert call_kwargs["created_by"] == mock_user_id
)
# Verify message file was created with correct parameters # Verify independent session was used (not db.session)
mock_msg_file_class.assert_called_once() mock_sm.assert_called_once_with(bind=mock_db.engine, expire_on_commit=False)
call_kwargs = mock_msg_file_class.call_args[1] file_session.add.assert_called_once_with(mock_message_file)
assert call_kwargs["message_id"] == mock_message_id mock_db.session.commit.assert_not_called()
assert call_kwargs["type"] == FileType.IMAGE mock_db.session.close.assert_not_called()
assert call_kwargs["transfer_method"] == FileTransferMethod.TOOL_FILE
assert call_kwargs["belongs_to"] == "assistant"
assert call_kwargs["created_by"] == mock_user_id
# Verify database operations # Verify event was published
mock_session.add.assert_called_once_with(mock_message_file) mock_queue_manager.publish.assert_called_once()
mock_session.commit.assert_called_once() publish_call = mock_queue_manager.publish.call_args
mock_session.refresh.assert_called_once_with(mock_message_file) assert isinstance(publish_call[0][0], QueueMessageFileEvent)
assert publish_call[0][0].message_file_id == mock_message_file.id
# Verify event was published
mock_queue_manager.publish.assert_called_once()
publish_call = mock_queue_manager.publish.call_args
assert isinstance(publish_call[0][0], QueueMessageFileEvent)
assert publish_call[0][0].message_file_id == mock_message_file.id
# publish_from might be passed as positional or keyword argument
assert (
publish_call[0][1] == PublishFrom.APPLICATION_MANAGER
or publish_call.kwargs.get("publish_from") == PublishFrom.APPLICATION_MANAGER
)
def test_handle_multimodal_image_content_with_base64( def test_handle_multimodal_image_content_with_base64(
self, self,
@ -165,50 +160,44 @@ class TestBaseAppRunnerMultimodal:
mock_mgr_class.return_value = mock_mgr mock_mgr_class.return_value = mock_mgr
with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class:
# Setup mock message file
mock_msg_file_class.return_value = mock_message_file mock_msg_file_class.return_value = mock_message_file
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: file_session = MagicMock()
mock_session.add = MagicMock() mock_session_factory = MagicMock()
mock_session.commit = MagicMock() mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session)
mock_session.refresh = MagicMock() mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False)
# Act with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory):
# Create a mock runner with the method bound with patch("core.app.apps.base_app_runner.db") as mock_db:
runner = MagicMock() runner = MagicMock()
method = AppRunner._handle_multimodal_image_content method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) runner._handle_multimodal_image_content = lambda *args, **kwargs: method(
runner, *args, **kwargs
)
runner._handle_multimodal_image_content( runner._handle_multimodal_image_content(
content=content, content=content,
message_id=mock_message_id, message_id=mock_message_id,
user_id=mock_user_id, user_id=mock_user_id,
tenant_id=mock_tenant_id, tenant_id=mock_tenant_id,
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
) )
# Assert mock_mgr.create_file_by_raw.assert_called_once()
# Verify tool file was created from base64 call_kwargs = mock_mgr.create_file_by_raw.call_args[1]
mock_mgr.create_file_by_raw.assert_called_once() assert call_kwargs["user_id"] == mock_user_id
call_kwargs = mock_mgr.create_file_by_raw.call_args[1] assert call_kwargs["tenant_id"] == mock_tenant_id
assert call_kwargs["user_id"] == mock_user_id assert call_kwargs["conversation_id"] is None
assert call_kwargs["tenant_id"] == mock_tenant_id assert "file_binary" in call_kwargs
assert call_kwargs["conversation_id"] is None assert call_kwargs["mimetype"] == "image/png"
assert "file_binary" in call_kwargs assert call_kwargs["filename"].startswith("generated_image")
assert call_kwargs["mimetype"] == "image/png" assert call_kwargs["filename"].endswith(".png")
assert call_kwargs["filename"].startswith("generated_image")
assert call_kwargs["filename"].endswith(".png")
# Verify message file was created mock_msg_file_class.assert_called_once()
mock_msg_file_class.assert_called_once() file_session.add.assert_called_once()
mock_db.session.commit.assert_not_called()
# Verify database operations mock_queue_manager.publish.assert_called_once()
mock_session.add.assert_called_once()
mock_session.commit.assert_called_once()
mock_session.refresh.assert_called_once()
# Verify event was published
mock_queue_manager.publish.assert_called_once()
def test_handle_multimodal_image_content_with_base64_data_uri( def test_handle_multimodal_image_content_with_base64_data_uri(
self, self,
@ -238,33 +227,32 @@ class TestBaseAppRunnerMultimodal:
mock_mgr_class.return_value = mock_mgr mock_mgr_class.return_value = mock_mgr
with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class:
# Setup mock message file
mock_msg_file_class.return_value = mock_message_file mock_msg_file_class.return_value = mock_message_file
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: file_session = MagicMock()
mock_session.add = MagicMock() mock_session_factory = MagicMock()
mock_session.commit = MagicMock() mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session)
mock_session.refresh = MagicMock() mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False)
# Act with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory):
# Create a mock runner with the method bound with patch("core.app.apps.base_app_runner.db"):
runner = MagicMock() runner = MagicMock()
method = AppRunner._handle_multimodal_image_content method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) runner._handle_multimodal_image_content = lambda *args, **kwargs: method(
runner, *args, **kwargs
)
runner._handle_multimodal_image_content( runner._handle_multimodal_image_content(
content=content, content=content,
message_id=mock_message_id, message_id=mock_message_id,
user_id=mock_user_id, user_id=mock_user_id,
tenant_id=mock_tenant_id, tenant_id=mock_tenant_id,
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
) )
# Assert - verify that base64 data was extracted correctly (without prefix) mock_mgr.create_file_by_raw.assert_called_once()
mock_mgr.create_file_by_raw.assert_called_once() call_kwargs = mock_mgr.create_file_by_raw.call_args[1]
call_kwargs = mock_mgr.create_file_by_raw.call_args[1] assert "file_binary" in call_kwargs
# The base64 data should be decoded, so we check the binary was passed
assert "file_binary" in call_kwargs
def test_handle_multimodal_image_content_without_url_or_base64( def test_handle_multimodal_image_content_without_url_or_base64(
self, self,
@ -284,9 +272,7 @@ class TestBaseAppRunnerMultimodal:
with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class:
with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class:
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: with patch("core.app.apps.base_app_runner.db"):
# Act
# Create a mock runner with the method bound
runner = MagicMock() runner = MagicMock()
method = AppRunner._handle_multimodal_image_content method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs)
@ -299,10 +285,8 @@ class TestBaseAppRunnerMultimodal:
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
) )
# Assert - should not create any files or publish events
mock_mgr_class.assert_not_called() mock_mgr_class.assert_not_called()
mock_msg_file_class.assert_not_called() mock_msg_file_class.assert_not_called()
mock_session.add.assert_not_called()
mock_queue_manager.publish.assert_not_called() mock_queue_manager.publish.assert_not_called()
def test_handle_multimodal_image_content_with_error( def test_handle_multimodal_image_content_with_error(
@ -322,20 +306,16 @@ class TestBaseAppRunnerMultimodal:
) )
with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class:
# Setup mock to raise exception
mock_mgr = MagicMock() mock_mgr = MagicMock()
mock_mgr.create_file_by_url.side_effect = Exception("Network error") mock_mgr.create_file_by_url.side_effect = Exception("Network error")
mock_mgr_class.return_value = mock_mgr mock_mgr_class.return_value = mock_mgr
with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class:
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: with patch("core.app.apps.base_app_runner.db"):
# Act
# Create a mock runner with the method bound
runner = MagicMock() runner = MagicMock()
method = AppRunner._handle_multimodal_image_content method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs)
# Should not raise exception, just log it
runner._handle_multimodal_image_content( runner._handle_multimodal_image_content(
content=content, content=content,
message_id=mock_message_id, message_id=mock_message_id,
@ -344,9 +324,7 @@ class TestBaseAppRunnerMultimodal:
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
) )
# Assert - should not create message file or publish event on error
mock_msg_file_class.assert_not_called() mock_msg_file_class.assert_not_called()
mock_session.add.assert_not_called()
mock_queue_manager.publish.assert_not_called() mock_queue_manager.publish.assert_not_called()
def test_handle_multimodal_image_content_debugger_mode( def test_handle_multimodal_image_content_debugger_mode(
@ -369,37 +347,36 @@ class TestBaseAppRunnerMultimodal:
mock_queue_manager.invoke_from = InvokeFrom.DEBUGGER mock_queue_manager.invoke_from = InvokeFrom.DEBUGGER
with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class:
# Setup mock tool file manager
mock_mgr = MagicMock() mock_mgr = MagicMock()
mock_mgr.create_file_by_url.return_value = mock_tool_file mock_mgr.create_file_by_url.return_value = mock_tool_file
mock_mgr_class.return_value = mock_mgr mock_mgr_class.return_value = mock_mgr
with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class:
# Setup mock message file
mock_msg_file_class.return_value = mock_message_file mock_msg_file_class.return_value = mock_message_file
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: file_session = MagicMock()
mock_session.add = MagicMock() mock_session_factory = MagicMock()
mock_session.commit = MagicMock() mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session)
mock_session.refresh = MagicMock() mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False)
# Act with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory):
# Create a mock runner with the method bound with patch("core.app.apps.base_app_runner.db"):
runner = MagicMock() runner = MagicMock()
method = AppRunner._handle_multimodal_image_content method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) runner._handle_multimodal_image_content = lambda *args, **kwargs: method(
runner, *args, **kwargs
)
runner._handle_multimodal_image_content( runner._handle_multimodal_image_content(
content=content, content=content,
message_id=mock_message_id, message_id=mock_message_id,
user_id=mock_user_id, user_id=mock_user_id,
tenant_id=mock_tenant_id, tenant_id=mock_tenant_id,
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
) )
# Assert - verify created_by_role is ACCOUNT for debugger mode call_kwargs = mock_msg_file_class.call_args[1]
call_kwargs = mock_msg_file_class.call_args[1] assert call_kwargs["created_by_role"] == CreatorUserRole.ACCOUNT
assert call_kwargs["created_by_role"] == CreatorUserRole.ACCOUNT
def test_handle_multimodal_image_content_service_api_mode( def test_handle_multimodal_image_content_service_api_mode(
self, self,
@ -421,34 +398,33 @@ class TestBaseAppRunnerMultimodal:
mock_queue_manager.invoke_from = InvokeFrom.SERVICE_API mock_queue_manager.invoke_from = InvokeFrom.SERVICE_API
with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class: with patch("core.app.apps.base_app_runner.ToolFileManager", autospec=True) as mock_mgr_class:
# Setup mock tool file manager
mock_mgr = MagicMock() mock_mgr = MagicMock()
mock_mgr.create_file_by_url.return_value = mock_tool_file mock_mgr.create_file_by_url.return_value = mock_tool_file
mock_mgr_class.return_value = mock_mgr mock_mgr_class.return_value = mock_mgr
with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class: with patch("core.app.apps.base_app_runner.MessageFile", autospec=True) as mock_msg_file_class:
# Setup mock message file
mock_msg_file_class.return_value = mock_message_file mock_msg_file_class.return_value = mock_message_file
with patch("core.app.apps.base_app_runner.db.session", autospec=True) as mock_session: file_session = MagicMock()
mock_session.add = MagicMock() mock_session_factory = MagicMock()
mock_session.commit = MagicMock() mock_session_factory.begin.return_value.__enter__ = MagicMock(return_value=file_session)
mock_session.refresh = MagicMock() mock_session_factory.begin.return_value.__exit__ = MagicMock(return_value=False)
# Act with patch("core.app.apps.base_app_runner.sessionmaker", return_value=mock_session_factory):
# Create a mock runner with the method bound with patch("core.app.apps.base_app_runner.db"):
runner = MagicMock() runner = MagicMock()
method = AppRunner._handle_multimodal_image_content method = AppRunner._handle_multimodal_image_content
runner._handle_multimodal_image_content = lambda *args, **kwargs: method(runner, *args, **kwargs) runner._handle_multimodal_image_content = lambda *args, **kwargs: method(
runner, *args, **kwargs
)
runner._handle_multimodal_image_content( runner._handle_multimodal_image_content(
content=content, content=content,
message_id=mock_message_id, message_id=mock_message_id,
user_id=mock_user_id, user_id=mock_user_id,
tenant_id=mock_tenant_id, tenant_id=mock_tenant_id,
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
) )
# Assert - verify created_by_role is END_USER for service API call_kwargs = mock_msg_file_class.call_args[1]
call_kwargs = mock_msg_file_class.call_args[1] assert call_kwargs["created_by_role"] == CreatorUserRole.END_USER
assert call_kwargs["created_by_role"] == CreatorUserRole.END_USER

View File

@ -14,6 +14,9 @@ def mock_queue_manager(mocker: MockerFixture):
@pytest.fixture @pytest.fixture
def handler(mock_queue_manager, mocker: MockerFixture): def handler(mock_queue_manager, mocker: MockerFixture):
mocker.patch(
"core.callback_handler.index_tool_callback_handler.db",
)
return DatasetIndexToolCallbackHandler( return DatasetIndexToolCallbackHandler(
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
app_id="app-1", app_id="app-1",
@ -33,8 +36,18 @@ class TestOnQuery:
], ],
) )
def test_on_query_success_roles(self, mocker: MockerFixture, mock_queue_manager, invoke_from, expected_role): def test_on_query_success_roles(self, mocker: MockerFixture, mock_queue_manager, invoke_from, expected_role):
# Arrange # Arrange — the caller passes a session, but our fix uses an independent one
mock_session = mocker.Mock() caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mock_session_factory = mocker.MagicMock()
mock_session_factory.begin.return_value.__enter__ = mocker.MagicMock(return_value=independent_session)
mock_session_factory.begin.return_value.__exit__ = mocker.MagicMock(return_value=False)
mocker.patch(
"core.callback_handler.index_tool_callback_handler.sessionmaker",
return_value=mock_session_factory,
)
mocker.patch("core.callback_handler.index_tool_callback_handler.db")
handler = DatasetIndexToolCallbackHandler( handler = DatasetIndexToolCallbackHandler(
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
@ -46,17 +59,28 @@ class TestOnQuery:
handler._invoke_from = invoke_from handler._invoke_from = invoke_from
# Act # Act — pass caller_session as required by signature
handler.on_query("test query", "dataset-1", mock_session) handler.on_query("test query", "dataset-1", caller_session)
# Assert # Assert — independent session used, not the caller's session
mock_session.add.assert_called_once() independent_session.add.assert_called_once()
dataset_query = mock_session.add.call_args.args[0] dataset_query = independent_session.add.call_args.args[0]
assert dataset_query.created_by_role == expected_role assert dataset_query.created_by_role == expected_role
mock_session.commit.assert_called_once() caller_session.add.assert_not_called()
caller_session.commit.assert_not_called()
def test_on_query_none_values(self, mocker: MockerFixture, mock_queue_manager): def test_on_query_none_values(self, mocker: MockerFixture, mock_queue_manager):
mock_session = mocker.Mock() caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mock_session_factory = mocker.MagicMock()
mock_session_factory.begin.return_value.__enter__ = mocker.MagicMock(return_value=independent_session)
mock_session_factory.begin.return_value.__exit__ = mocker.MagicMock(return_value=False)
mocker.patch(
"core.callback_handler.index_tool_callback_handler.sessionmaker",
return_value=mock_session_factory,
)
mocker.patch("core.callback_handler.index_tool_callback_handler.db")
handler = DatasetIndexToolCallbackHandler( handler = DatasetIndexToolCallbackHandler(
queue_manager=mock_queue_manager, queue_manager=mock_queue_manager,
@ -66,40 +90,67 @@ class TestOnQuery:
invoke_from=None, invoke_from=None,
) )
handler.on_query(None, None, mock_session) handler.on_query(None, None, caller_session)
mock_session.add.assert_called_once() independent_session.add.assert_called_once()
mock_session.commit.assert_called_once() caller_session.add.assert_not_called()
class TestOnToolEnd: class TestOnToolEnd:
def test_on_tool_end_no_metadata(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): def test_on_tool_end_no_metadata(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture):
mock_session = mocker.Mock() caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
document = mocker.Mock() document = mocker.Mock()
document.metadata = None document.metadata = None
handler.on_tool_end([document], mock_session) handler.on_tool_end([document], caller_session)
mock_session.commit.assert_not_called() independent_session.commit.assert_called_once()
independent_session.execute.assert_not_called()
caller_session.commit.assert_not_called()
def test_on_tool_end_dataset_document_not_found( def test_on_tool_end_dataset_document_not_found(
self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture
): ):
mock_session = mocker.Mock() caller_session = mocker.Mock()
mock_session.scalar.return_value = None
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
independent_session.scalar.return_value = None
document = mocker.Mock() document = mocker.Mock()
document.metadata = {"document_id": "doc-1", "doc_id": "node-1"} document.metadata = {"document_id": "doc-1", "doc_id": "node-1"}
handler.on_tool_end([document], mock_session) handler.on_tool_end([document], caller_session)
mock_session.scalar.assert_called_once() independent_session.scalar.assert_called_once()
caller_session.scalar.assert_not_called()
def test_on_tool_end_parent_child_index_with_child( def test_on_tool_end_parent_child_index_with_child(
self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture
): ):
mock_session = mocker.Mock() caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
mock_dataset_doc = mocker.Mock() mock_dataset_doc = mocker.Mock()
from core.callback_handler.index_tool_callback_handler import IndexStructureType from core.callback_handler.index_tool_callback_handler import IndexStructureType
@ -111,23 +162,32 @@ class TestOnToolEnd:
mock_child_chunk = mocker.Mock() mock_child_chunk = mocker.Mock()
mock_child_chunk.segment_id = "segment-1" mock_child_chunk.segment_id = "segment-1"
mock_session.scalar.side_effect = [mock_dataset_doc, mock_child_chunk] independent_session.scalar.side_effect = [mock_dataset_doc, mock_child_chunk]
document = mocker.Mock() document = mocker.Mock()
document.metadata = {"document_id": "doc-1", "doc_id": "node-1"} document.metadata = {"document_id": "doc-1", "doc_id": "node-1"}
handler.on_tool_end([document], mock_session) handler.on_tool_end([document], caller_session)
mock_session.execute.assert_called_once() independent_session.execute.assert_called_once()
mock_session.commit.assert_called_once() independent_session.commit.assert_called_once()
caller_session.execute.assert_not_called()
def test_on_tool_end_non_parent_child_index(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): def test_on_tool_end_non_parent_child_index(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture):
mock_session = mocker.Mock() caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
mock_dataset_doc = mocker.Mock() mock_dataset_doc = mocker.Mock()
mock_dataset_doc.doc_form = "OTHER" mock_dataset_doc.doc_form = "OTHER"
mock_session.scalar.return_value = mock_dataset_doc independent_session.scalar.return_value = mock_dataset_doc
document = mocker.Mock() document = mocker.Mock()
document.metadata = { document.metadata = {
@ -136,14 +196,24 @@ class TestOnToolEnd:
"dataset_id": "dataset-1", "dataset_id": "dataset-1",
} }
handler.on_tool_end([document], mock_session) handler.on_tool_end([document], caller_session)
mock_session.execute.assert_called_once() independent_session.execute.assert_called_once()
mock_session.commit.assert_called_once() independent_session.commit.assert_called_once()
caller_session.execute.assert_not_called()
def test_on_tool_end_empty_documents(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture): def test_on_tool_end_empty_documents(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture):
mock_session = mocker.Mock() caller_session = mocker.Mock()
handler.on_tool_end([], mock_session)
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
handler.on_tool_end([], caller_session)
class TestReturnRetrieverResourceInfo: class TestReturnRetrieverResourceInfo: