diff --git a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py index 78e46d272c2..5db648a3014 100644 --- a/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py +++ b/api/tests/unit_tests/services/rag_pipeline/pipeline_template/test_remote_retrieval.py @@ -1,12 +1,14 @@ import pytest from pytest_mock import MockerFixture +from sqlalchemy.orm import Session from services.rag_pipeline.pipeline_template.database.database_retrieval import DatabasePipelineTemplateRetrieval from services.rag_pipeline.pipeline_template.pipeline_template_type import PipelineTemplateType from services.rag_pipeline.pipeline_template.remote.remote_retrieval import RemotePipelineTemplateRetrieval -def test_get_pipeline_templates_fallbacks_to_database_on_error(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_get_pipeline_templates_fallbacks_to_database_on_error(mocker: MockerFixture, sqlite_session: Session) -> None: fetch_mock = mocker.patch.object( RemotePipelineTemplateRetrieval, "fetch_pipeline_templates_from_dify_official", @@ -18,17 +20,20 @@ def test_get_pipeline_templates_fallbacks_to_database_on_error(mocker: MockerFix return_value={"pipeline_templates": [{"id": "db-1"}]}, ) retrieval = RemotePipelineTemplateRetrieval() - session = mocker.Mock() - result = retrieval.get_pipeline_templates("en-US", session=session) + result = retrieval.get_pipeline_templates("en-US", session=sqlite_session) assert retrieval.get_type() == PipelineTemplateType.REMOTE assert result == {"pipeline_templates": [{"id": "db-1"}]} fetch_mock.assert_called_once_with("en-US") - fallback_mock.assert_called_once_with("en-US", session=session) + fallback_mock.assert_called_once_with("en-US", session=sqlite_session) + assert not sqlite_session.in_transaction() -def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: MockerFixture) -> None: +@pytest.mark.parametrize("sqlite_session", [()], indirect=True) +def test_get_pipeline_template_detail_fallbacks_to_database_on_error( + mocker: MockerFixture, sqlite_session: Session +) -> None: fetch_mock = mocker.patch.object( RemotePipelineTemplateRetrieval, "fetch_pipeline_template_detail_from_dify_official", @@ -40,13 +45,13 @@ def test_get_pipeline_template_detail_fallbacks_to_database_on_error(mocker: Moc return_value={"id": "db-1"}, ) retrieval = RemotePipelineTemplateRetrieval() - session = mocker.Mock() - result = retrieval.get_pipeline_template_detail("tpl-1", session=session) + result = retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session) assert result == {"id": "db-1"} fetch_mock.assert_called_once_with("tpl-1") - fallback_mock.assert_called_once_with("tpl-1", session=session) + fallback_mock.assert_called_once_with("tpl-1", session=sqlite_session) + assert not sqlite_session.in_transaction() def test_fetch_pipeline_templates_from_dify_official(mocker: MockerFixture) -> None: