mirror of
https://github.com/langgenius/dify.git
synced 2026-08-01 18:30:48 +08:00
test: use SQLite sessions in core callback_handler (#39081)
This commit is contained in:
parent
753583c06c
commit
917207414b
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user