test: use SQLite sessions in core callback_handler (#39081)

This commit is contained in:
Asuka Minato 2026-07-31 21:54:20 +09:00 committed by GitHub
parent 753583c06c
commit 917207414b
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,10 +1,38 @@
from collections.abc import Iterator
from dataclasses import dataclass
from uuid import uuid4
import pytest
from pytest_mock import MockerFixture
from sqlalchemy import select
from sqlalchemy.engine import Engine
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session, SessionTransaction
import core.callback_handler.index_tool_callback_handler as callback_module
from core.app.entities.app_invoke_entities import InvokeFrom
from core.callback_handler.index_tool_callback_handler import (
DatasetIndexToolCallbackHandler,
)
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.rag.index_processor.constant.index_type import IndexStructureType
from core.rag.models.document import Document
from models.dataset import ChildChunk, DatasetQuery, DocumentSegment
from models.dataset import Document as DatasetDocument
from models.enums import CreatorUserRole, DatasetQuerySource, DataSourceType, DocumentCreatedFrom
class _DatabaseBinding:
engine: Engine
def __init__(self, engine: Engine) -> None:
self.engine = engine
@dataclass(frozen=True)
class _CallerSessionBoundary:
"""Caller session state that callback-owned transactions must not disturb."""
session: Session
transaction: SessionTransaction
pending_query: DatasetQuery
@pytest.fixture
@ -13,16 +41,75 @@ def mock_queue_manager(mocker: MockerFixture):
@pytest.fixture
def handler(mock_queue_manager, mocker: MockerFixture):
mocker.patch(
"core.callback_handler.index_tool_callback_handler.db",
)
def handler(mock_queue_manager, sqlite_engine: Engine, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(callback_module, "db", _DatabaseBinding(sqlite_engine))
return DatasetIndexToolCallbackHandler(
queue_manager=mock_queue_manager,
app_id="app-1",
message_id="msg-1",
user_id="user-1",
invoke_from=mocker.Mock(),
app_id=str(uuid4()),
message_id=str(uuid4()),
user_id=str(uuid4()),
invoke_from=InvokeFrom.DEBUGGER,
)
@pytest.fixture
def caller_session_boundary(sqlite_engine: Engine) -> Iterator[_CallerSessionBoundary]:
"""Keep an unflushed caller transaction open to prove callback isolation."""
with Session(sqlite_engine, expire_on_commit=False) as session:
pending_query = DatasetQuery(
dataset_id=str(uuid4()),
content="caller-owned pending query",
source=DatasetQuerySource.APP,
source_app_id=str(uuid4()),
created_by_role=CreatorUserRole.ACCOUNT,
created_by=str(uuid4()),
)
session.add(pending_query)
transaction = session.get_transaction()
assert transaction is not None
yield _CallerSessionBoundary(
session=session,
transaction=transaction,
pending_query=pending_query,
)
assert session.get_transaction() is transaction
assert list(session.new) == [pending_query]
assert not session.dirty
assert not session.deleted
with Session(sqlite_engine) as observer_session:
assert observer_session.get(DatasetQuery, pending_query.id) is None
def _dataset_document(*, doc_form: IndexStructureType) -> DatasetDocument:
return DatasetDocument(
tenant_id=str(uuid4()),
dataset_id=str(uuid4()),
position=1,
data_source_type=DataSourceType.UPLOAD_FILE,
batch="batch",
name="Document",
created_from=DocumentCreatedFrom.API,
created_by=str(uuid4()),
doc_form=doc_form,
)
def _segment(document: DatasetDocument, *, index_node_id: str) -> DocumentSegment:
return DocumentSegment(
tenant_id=document.tenant_id,
dataset_id=document.dataset_id,
document_id=document.id,
position=1,
content="content",
word_count=1,
tokens=1,
created_by=document.created_by,
index_node_id=index_node_id,
hit_count=0,
)
@ -35,53 +122,40 @@ class TestOnQuery:
(InvokeFrom.WEB_APP, "end_user"),
],
)
def test_on_query_success_roles(self, mocker: MockerFixture, mock_queue_manager, invoke_from, expected_role):
# Arrange — the caller passes a session, but our fix uses an independent one
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")
def test_on_query_success_roles(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
caller_session_boundary: _CallerSessionBoundary,
mock_queue_manager,
invoke_from: InvokeFrom,
expected_role: str,
) -> None:
monkeypatch.setattr(callback_module, "db", _DatabaseBinding(sqlite_engine))
handler = DatasetIndexToolCallbackHandler(
queue_manager=mock_queue_manager,
app_id="app-1",
message_id="msg-1",
user_id="user-1",
invoke_from=mocker.Mock(),
app_id=str(uuid4()),
message_id=str(uuid4()),
user_id=str(uuid4()),
invoke_from=invoke_from,
)
handler._invoke_from = invoke_from
handler.on_query("test query", str(uuid4()), caller_session_boundary.session)
# Act — pass caller_session as required by signature
handler.on_query("test query", "dataset-1", caller_session)
# Assert — independent session used, not the caller's session
independent_session.add.assert_called_once()
dataset_query = independent_session.add.call_args.args[0]
dataset_query = sqlite_session.scalar(select(DatasetQuery))
assert dataset_query is not None
assert dataset_query.created_by_role == expected_role
caller_session.add.assert_not_called()
caller_session.commit.assert_not_called()
def test_on_query_none_values(self, mocker: MockerFixture, mock_queue_manager):
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")
def test_on_query_none_values_roll_back_independent_transaction(
self,
monkeypatch: pytest.MonkeyPatch,
sqlite_engine: Engine,
sqlite_session: Session,
caller_session_boundary: _CallerSessionBoundary,
mock_queue_manager,
) -> None:
monkeypatch.setattr(callback_module, "db", _DatabaseBinding(sqlite_engine))
handler = DatasetIndexToolCallbackHandler(
queue_manager=mock_queue_manager,
app_id=None,
@ -90,136 +164,109 @@ class TestOnQuery:
invoke_from=None,
)
handler.on_query(None, None, caller_session)
with pytest.raises(IntegrityError):
handler.on_query(None, None, caller_session_boundary.session) # type: ignore[arg-type]
independent_session.add.assert_called_once()
caller_session.add.assert_not_called()
assert sqlite_session.scalar(select(DatasetQuery)) is None
class TestOnToolEnd:
def test_on_tool_end_no_metadata(self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture):
caller_session = mocker.Mock()
def test_on_tool_end_no_metadata(
self,
handler: DatasetIndexToolCallbackHandler,
caller_session_boundary: _CallerSessionBoundary,
) -> None:
document = Document.model_construct(page_content="content", metadata=None, provider="dify")
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.metadata = None
handler.on_tool_end([document], caller_session)
independent_session.commit.assert_called_once()
independent_session.execute.assert_not_called()
caller_session.commit.assert_not_called()
handler.on_tool_end([document], caller_session_boundary.session)
def test_on_tool_end_dataset_document_not_found(
self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture
):
caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
self,
handler: DatasetIndexToolCallbackHandler,
sqlite_session: Session,
caller_session_boundary: _CallerSessionBoundary,
) -> None:
document = Document(
page_content="content",
metadata={"document_id": str(uuid4()), "doc_id": "node-1"},
)
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.metadata = {"document_id": "doc-1", "doc_id": "node-1"}
handler.on_tool_end([document], caller_session_boundary.session)
handler.on_tool_end([document], caller_session)
independent_session.scalar.assert_called_once()
caller_session.scalar.assert_not_called()
assert sqlite_session.scalar(select(DatasetDocument)) is None
def test_on_tool_end_parent_child_index_with_child(
self, handler: DatasetIndexToolCallbackHandler, mocker: MockerFixture
):
caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
self,
handler: DatasetIndexToolCallbackHandler,
sqlite_session: Session,
caller_session_boundary: _CallerSessionBoundary,
) -> None:
dataset_document = _dataset_document(doc_form=IndexStructureType.PARENT_CHILD_INDEX)
sqlite_session.add(dataset_document)
sqlite_session.flush()
segment = _segment(dataset_document, index_node_id="parent-node")
sqlite_session.add(segment)
sqlite_session.flush()
child = ChildChunk(
tenant_id=dataset_document.tenant_id,
dataset_id=dataset_document.dataset_id,
document_id=dataset_document.id,
segment_id=segment.id,
position=1,
content="child",
word_count=1,
created_by=dataset_document.created_by,
index_node_id="child-node",
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
mock_dataset_doc = mocker.Mock()
from core.callback_handler.index_tool_callback_handler import IndexStructureType
mock_dataset_doc.doc_form = IndexStructureType.PARENT_CHILD_INDEX
mock_dataset_doc.dataset_id = "dataset-1"
mock_dataset_doc.id = "doc-1"
mock_child_chunk = mocker.Mock()
mock_child_chunk.segment_id = "segment-1"
independent_session.scalar.side_effect = [mock_dataset_doc, mock_child_chunk]
document = mocker.Mock()
document.metadata = {"document_id": "doc-1", "doc_id": "node-1"}
handler.on_tool_end([document], caller_session)
independent_session.execute.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):
caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
sqlite_session.add(child)
sqlite_session.commit()
document = Document(
page_content="content",
metadata={"document_id": dataset_document.id, "doc_id": child.index_node_id},
)
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.doc_form = "OTHER"
handler.on_tool_end([document], caller_session_boundary.session)
independent_session.scalar.return_value = mock_dataset_doc
sqlite_session.expire_all()
assert sqlite_session.get(DocumentSegment, segment.id).hit_count == 1 # type: ignore[union-attr]
document = mocker.Mock()
document.metadata = {
"document_id": "doc-1",
"doc_id": "node-1",
"dataset_id": "dataset-1",
}
handler.on_tool_end([document], caller_session)
independent_session.execute.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):
caller_session = mocker.Mock()
independent_session = mocker.MagicMock()
mocker.patch(
"core.callback_handler.index_tool_callback_handler.Session",
return_value=independent_session,
def test_on_tool_end_non_parent_child_index(
self,
handler: DatasetIndexToolCallbackHandler,
sqlite_session: Session,
caller_session_boundary: _CallerSessionBoundary,
) -> None:
dataset_document = _dataset_document(doc_form=IndexStructureType.PARAGRAPH_INDEX)
sqlite_session.add(dataset_document)
sqlite_session.flush()
segment = _segment(dataset_document, index_node_id="node-1")
sqlite_session.add(segment)
sqlite_session.commit()
document = Document(
page_content="content",
metadata={
"document_id": dataset_document.id,
"doc_id": segment.index_node_id,
"dataset_id": dataset_document.dataset_id,
},
)
independent_session.__enter__ = mocker.MagicMock(return_value=independent_session)
independent_session.__exit__ = mocker.MagicMock(return_value=False)
handler.on_tool_end([], caller_session)
handler.on_tool_end([document], caller_session_boundary.session)
sqlite_session.expire_all()
assert sqlite_session.get(DocumentSegment, segment.id).hit_count == 1 # type: ignore[union-attr]
def test_on_tool_end_empty_documents(
self,
handler: DatasetIndexToolCallbackHandler,
caller_session_boundary: _CallerSessionBoundary,
) -> None:
handler.on_tool_end([], caller_session_boundary.session)
class TestReturnRetrieverResourceInfo:
def test_publish_called(self, handler: DatasetIndexToolCallbackHandler, mock_queue_manager, mocker: MockerFixture):
mock_event = mocker.patch("core.callback_handler.index_tool_callback_handler.QueueRetrieverResourcesEvent")
resources = [mocker.Mock()]
handler.return_retriever_resource_info(resources)