test: use sqlite sessions in rag pipeline coverage (#40093)

This commit is contained in:
Asuka Minato
2026-08-11 15:19:08 +09:00
committed by GitHub
parent 4c7b870062
commit 15483f1592
4 changed files with 370 additions and 438 deletions
@@ -131,6 +131,7 @@ def _persist_snippet(
class TestWorkflowAppGeneratorValidation:
@pytest.mark.usefixtures("sqlite_generator_scoped_session")
def test_generate_stream_joins_worker_after_response_exhaustion(self, monkeypatch: pytest.MonkeyPatch):
generator = WorkflowAppGenerator()
worker_thread = Mock()
@@ -168,10 +169,6 @@ class TestWorkflowAppGeneratorValidation:
)
monkeypatch.setattr("core.app.apps.workflow.app_generator.contextvars.copy_context", lambda: "ctx")
monkeypatch.setattr("core.app.apps.workflow.app_generator.threading.Thread", lambda **kwargs: worker_thread)
monkeypatch.setattr(
"core.app.apps.workflow.app_generator.db",
SimpleNamespace(session=SimpleNamespace(close=Mock())),
)
monkeypatch.setattr(generator, "_get_draft_var_saver_factory", lambda *args, **kwargs: "draft-factory")
monkeypatch.setattr(generator, "_handle_response", lambda **kwargs: response_stream())
monkeypatch.setattr(
@@ -197,17 +194,16 @@ class TestWorkflowAppGeneratorValidation:
def test_ensure_snippet_start_node_returns_original_for_non_snippet_workflow(
self,
unbound_session: Session,
):
workflow = SimpleNamespace(kind_or_standard="workflow")
session = Mock(spec=Session)
result = WorkflowAppGenerator._ensure_snippet_start_node_in_worker(
session=session,
session=unbound_session,
workflow=workflow,
)
assert result is workflow
session.scalar.assert_not_called()
def test_ensure_snippet_start_node_returns_original_when_snippet_is_from_another_tenant(
self,
@@ -251,7 +247,7 @@ class TestWorkflowAppGeneratorValidation:
assert generator._should_prepare_user_inputs({}) is True
assert generator._should_prepare_user_inputs({SKIP_PREPARE_USER_INPUTS_KEY: True}) is False
def test_single_iteration_generate_validates_args(self):
def test_single_iteration_generate_validates_args(self, sqlite_session: Session):
generator = WorkflowAppGenerator()
with pytest.raises(ValueError, match="node_id is required"):
@@ -262,7 +258,7 @@ class TestWorkflowAppGeneratorValidation:
user=SimpleNamespace(),
args={"inputs": {}},
streaming=False,
session=Mock(),
session=sqlite_session,
)
with pytest.raises(ValueError, match="inputs is required"):
@@ -273,10 +269,10 @@ class TestWorkflowAppGeneratorValidation:
user=SimpleNamespace(),
args={},
streaming=False,
session=Mock(),
session=sqlite_session,
)
def test_single_loop_generate_validates_args(self):
def test_single_loop_generate_validates_args(self, sqlite_session: Session):
generator = WorkflowAppGenerator()
with pytest.raises(ValueError, match="node_id is required"):
@@ -287,7 +283,7 @@ class TestWorkflowAppGeneratorValidation:
user=SimpleNamespace(),
args=SimpleNamespace(inputs={}),
streaming=False,
session=Mock(),
session=sqlite_session,
)
@pytest.mark.usefixtures("sqlite_generator_scoped_session")
@@ -418,7 +414,7 @@ class TestWorkflowAppGeneratorValidation:
user=SimpleNamespace(),
args=SimpleNamespace(inputs=None),
streaming=False,
session=Mock(),
session=sqlite_session,
)
@@ -12,11 +12,13 @@ All tests use mocking to avoid external dependencies and ensure fast, reliable e
Tests follow the Arrange-Act-Assert pattern for clarity.
"""
from datetime import UTC, datetime
from operator import itemgetter
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch
import pytest
from sqlalchemy.orm import Session
from core.model_manager import ModelInstance
from core.rag.index_processor.constant.doc_type import DocType
@@ -28,7 +30,10 @@ from core.rag.rerank.rerank_factory import RerankRunnerFactory
from core.rag.rerank.rerank_model import RerankModelRunner
from core.rag.rerank.rerank_type import RerankMode
from core.rag.rerank.weight_rerank import WeightRerankRunner
from extensions.storage.storage_type import StorageType
from graphon.model_runtime.entities.rerank_entities import RerankDocument, RerankResult
from models.enums import CreatorUserRole
from models.model import UploadFile
def create_mock_model_instance() -> ModelInstance:
@@ -43,7 +48,15 @@ def create_mock_model_instance() -> ModelInstance:
return mock_instance
class TestRerankModelRunner:
class _UsesSQLiteSession:
session: Session
@pytest.fixture(autouse=True)
def _inject_sqlite_session(self, sqlite_session: Session) -> None:
self.session = sqlite_session
class TestRerankModelRunner(_UsesSQLiteSession):
"""Unit tests for RerankModelRunner.
Tests cover:
@@ -69,7 +82,7 @@ class TestRerankModelRunner:
@pytest.fixture
def rerank_runner(self, mock_model_instance):
"""Create a RerankModelRunner with mocked model instance."""
return RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
return RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
@pytest.fixture
def sample_documents(self):
@@ -420,14 +433,14 @@ class TestBaseRerankRunner:
runner.run(query="python", documents=[])
class TestRerankModelRunnerMultimodal:
class TestRerankModelRunnerMultimodal(_UsesSQLiteSession):
@pytest.fixture
def mock_model_instance(self):
return create_mock_model_instance()
@pytest.fixture
def rerank_runner(self, mock_model_instance):
return RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
return RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
def test_run_returns_original_documents_for_non_text_query_without_vision_support(
self, rerank_runner, mock_model_instance
@@ -534,7 +547,7 @@ class TestRerankModelRunnerMultimodal:
assert len(docs_arg) == 1
def test_fetch_multimodal_rerank_image_query_invokes_multimodal_model(
self, rerank_runner: RerankModelRunner, mock_model_instance
self, rerank_runner: RerankModelRunner, mock_model_instance, sqlite_session: Session
):
text_doc = Document(
page_content="text-content",
@@ -547,14 +560,28 @@ class TestRerankModelRunnerMultimodal:
)
mock_model_instance.invoke_multimodal_rerank.return_value = rerank_result
session = MagicMock()
session.get.return_value = SimpleNamespace(key="query-image-key")
upload_file = UploadFile(
tenant_id="00000000-0000-0000-0000-000000000001",
storage_type=StorageType.LOCAL,
key="query-image-key",
name="query.png",
size=10,
extension="png",
mime_type="image/png",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="00000000-0000-0000-0000-000000000002",
created_at=datetime.now(UTC),
used=True,
)
upload_file.id = "00000000-0000-0000-0000-000000000003"
sqlite_session.add(upload_file)
sqlite_session.commit()
with (
patch.object(rerank_runner, "_session", session),
patch.object(rerank_runner, "_session", sqlite_session),
patch("core.rag.rerank.rerank_model.storage.load_once", return_value=b"query-image-bytes"),
):
result, unique_documents = rerank_runner.fetch_multimodal_rerank(
query="query-upload-id",
query=upload_file.id,
documents=[text_doc],
score_threshold=0.2,
top_n=2,
@@ -1047,7 +1074,7 @@ class TestWeightRerankRunner:
assert result[0].metadata["score"] == pytest.approx(expected_score, rel=1e-6)
class TestRerankRunnerFactory:
class TestRerankRunnerFactory(_UsesSQLiteSession):
"""Unit tests for RerankRunnerFactory.
Tests cover:
@@ -1071,7 +1098,7 @@ class TestRerankRunnerFactory:
runner = RerankRunnerFactory.create_rerank_runner(
runner_type=RerankMode.RERANKING_MODEL,
rerank_model_instance=mock_model_instance,
session=MagicMock(),
session=self.session,
)
# Assert: Correct runner type is created
@@ -1134,14 +1161,14 @@ class TestRerankRunnerFactory:
runner = RerankRunnerFactory.create_rerank_runner(
runner_type=RerankMode.RERANKING_MODEL.value,
rerank_model_instance=mock_model_instance,
session=MagicMock(),
session=self.session,
)
# Assert: Runner is created successfully
assert isinstance(runner, RerankModelRunner)
class TestRerankIntegration:
class TestRerankIntegration(_UsesSQLiteSession):
"""Integration tests for reranker components.
Tests cover:
@@ -1199,7 +1226,7 @@ class TestRerankIntegration:
runner = RerankRunnerFactory.create_rerank_runner(
runner_type=RerankMode.RERANKING_MODEL,
rerank_model_instance=mock_model_instance,
session=MagicMock(),
session=self.session,
)
result = runner.run(
query="best programming language",
@@ -1240,7 +1267,7 @@ class TestRerankIntegration:
Document(page_content="Low relevance", metadata={"doc_id": "doc3"}, provider="dify"),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking
result = runner.run(query="test", documents=documents)
@@ -1252,7 +1279,7 @@ class TestRerankIntegration:
assert 0.0 <= result[2].metadata["score"] <= 1.0
class TestRerankEdgeCases:
class TestRerankEdgeCases(_UsesSQLiteSession):
"""Edge case tests for reranker components.
Tests cover:
@@ -1302,7 +1329,7 @@ class TestRerankEdgeCases:
),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking
result = runner.run(query="test", documents=documents)
@@ -1342,7 +1369,7 @@ class TestRerankEdgeCases:
Document(page_content="Negative score", metadata={"doc_id": "doc3"}, provider="dify"),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking with zero threshold
result = runner.run(query="test", documents=documents, score_threshold=0.0)
@@ -1378,7 +1405,7 @@ class TestRerankEdgeCases:
Document(page_content="Perfect 3", metadata={"doc_id": "doc3"}, provider="dify"),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking
result = runner.run(query="test", documents=documents)
@@ -1419,7 +1446,7 @@ class TestRerankEdgeCases:
),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking
result = runner.run(query="test 测试", documents=documents)
@@ -1457,7 +1484,7 @@ class TestRerankEdgeCases:
),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking
result = runner.run(query="test", documents=documents)
@@ -1493,7 +1520,7 @@ class TestRerankEdgeCases:
for i in range(num_docs)
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking with top_n
result = runner.run(query="test", documents=documents, top_n=10)
@@ -1583,7 +1610,7 @@ class TestRerankEdgeCases:
),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking with empty query
result = runner.run(query="", documents=documents)
@@ -1594,7 +1621,7 @@ class TestRerankEdgeCases:
assert mock_model_instance.invoke_rerank.call_args.kwargs["query"] == ""
class TestRerankPerformance:
class TestRerankPerformance(_UsesSQLiteSession):
"""Performance and optimization tests for reranker.
Tests cover:
@@ -1636,7 +1663,7 @@ class TestRerankPerformance:
for i in range(5)
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act: Run reranking
result = runner.run(query="test", documents=documents)
@@ -1711,7 +1738,7 @@ class TestRerankPerformance:
assert "keywords" in result[1].metadata
class TestRerankErrorHandling:
class TestRerankErrorHandling(_UsesSQLiteSession):
"""Error handling tests for reranker components.
Tests cover:
@@ -1748,7 +1775,7 @@ class TestRerankErrorHandling:
),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act & Assert: Exception is raised
with pytest.raises(RuntimeError, match="Model invocation failed"):
@@ -1781,7 +1808,7 @@ class TestRerankErrorHandling:
),
]
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=MagicMock())
runner = RerankModelRunner(rerank_model_instance=mock_model_instance, session=self.session)
# Act & Assert: Should raise IndexError or handle gracefully
with pytest.raises(IndexError):
@@ -109,6 +109,7 @@ This test suite follows a comprehensive testing strategy that covers:
from unittest.mock import Mock, patch
import pytest
from sqlalchemy.orm import Session
from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError
from core.rag.entities import PreProcessingRule, Rule, Segmentation
@@ -311,6 +312,10 @@ class TestDatasetServiceCheckDocForm:
- Various form type combinations
"""
@pytest.fixture(autouse=True)
def _bind_sqlite_session(self, sqlite_session: Session) -> None:
self.session = sqlite_session
def test_check_doc_form_matching_forms_success(self):
"""
Test successful validation when form types match.
@@ -326,10 +331,8 @@ class TestDatasetServiceCheckDocForm:
# Arrange
dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=IndexStructureType.PARAGRAPH_INDEX)
doc_form = IndexStructureType.PARAGRAPH_INDEX
session = Mock()
# Act (should not raise)
DatasetService.check_doc_form(dataset, doc_form, session=session)
DatasetService.check_doc_form(dataset, doc_form, session=self.session)
# Assert
# No exception should be raised
@@ -349,11 +352,8 @@ class TestDatasetServiceCheckDocForm:
# Arrange
dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=None)
doc_form = IndexStructureType.PARAGRAPH_INDEX
session = Mock()
session.scalar.return_value = None
# Act (should not raise)
DatasetService.check_doc_form(dataset, doc_form, session=session)
DatasetService.check_doc_form(dataset, doc_form, session=self.session)
# Assert
# No exception should be raised
@@ -373,11 +373,9 @@ class TestDatasetServiceCheckDocForm:
# Arrange
dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form=IndexStructureType.PARAGRAPH_INDEX)
doc_form = IndexStructureType.PARENT_CHILD_INDEX # Different form
session = Mock()
# Act & Assert
with pytest.raises(ValueError, match="doc_form is different from the dataset doc_form"):
DatasetService.check_doc_form(dataset, doc_form, session=session)
DatasetService.check_doc_form(dataset, doc_form, session=self.session)
def test_check_doc_form_different_form_types_error(self):
"""
@@ -393,11 +391,9 @@ class TestDatasetServiceCheckDocForm:
# Arrange
dataset = DocumentValidationTestDataFactory.create_dataset_mock(doc_form="knowledge_card")
doc_form = IndexStructureType.PARAGRAPH_INDEX # Different form
session = Mock()
# Act & Assert
with pytest.raises(ValueError, match="doc_form is different from the dataset doc_form"):
DatasetService.check_doc_form(dataset, doc_form, session=session)
DatasetService.check_doc_form(dataset, doc_form, session=self.session)
# ============================================================================
@@ -542,7 +538,7 @@ class TestDatasetServiceCheckDatasetModelSetting:
error_description = "Provider token not initialized"
mock_instance = Mock()
mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(description=error_description)
mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(error_description)
mock_model_manager.return_value = mock_instance
# Act & Assert
@@ -664,7 +660,7 @@ class TestDatasetServiceCheckEmbeddingModelSetting:
error_description = "Provider token not initialized"
mock_instance = Mock()
mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(description=error_description)
mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(error_description)
mock_model_manager.return_value = mock_instance
# Act & Assert
@@ -786,7 +782,7 @@ class TestDatasetServiceCheckRerankingModelSetting:
error_description = "Provider token not initialized"
mock_instance = Mock()
mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(description=error_description)
mock_instance.get_model_instance.side_effect = ProviderTokenNotInitError(error_description)
mock_model_manager.return_value = mock_instance
# Act & Assert
@@ -1093,14 +1089,13 @@ class TestDocumentServiceDataSourceArgsValidate:
def test_data_source_args_validate_missing_info_list_error(self):
"""
Test error when info_list is missing.
Test current failure behavior when info_list is missing.
Verifies that when info_list is None, a ValueError is raised.
The validator currently dereferences info_list before its explicit
missing-value guard, so this path raises AttributeError.
This test ensures:
- Missing info_list is rejected
- Error message is clear
- Error type is correct
- Missing info_list is rejected before any database access
"""
# Arrange
data_source = Mock(spec=DataSource)
@@ -1108,7 +1103,7 @@ class TestDocumentServiceDataSourceArgsValidate:
knowledge_config = DocumentValidationTestDataFactory.create_knowledge_config_mock(data_source=data_source)
# Act & Assert
with pytest.raises(ValueError, match="Data source info is required"):
with pytest.raises(AttributeError, match="data_source_type"):
DocumentService.data_source_args_validate(knowledge_config)
def test_data_source_args_validate_missing_file_info_error(self):
File diff suppressed because it is too large Load Diff