mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix: pipeline template render (#36168)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot]
Copilot Autofix powered by AI
parent
ebcc1200a3
commit
7210f856c9
@@ -259,6 +259,60 @@ workflow:
|
||||
if result.status == ImportStatus.FAILED:
|
||||
print(f"DEBUG: {result.error}")
|
||||
assert result.status == ImportStatus.COMPLETED
|
||||
session.commit.assert_not_called()
|
||||
session.flush.assert_called()
|
||||
|
||||
|
||||
def test_import_rag_pipeline_flushes_new_collection_binding_without_commit(mocker) -> None:
|
||||
yaml_content = """
|
||||
version: 0.1.0
|
||||
kind: rag_pipeline
|
||||
rag_pipeline:
|
||||
name: Test Pipeline
|
||||
workflow:
|
||||
graph:
|
||||
nodes:
|
||||
- data:
|
||||
type: knowledge-index
|
||||
"""
|
||||
pipeline = Mock(id="p1", description="desc", is_published=False)
|
||||
pipeline.name = "Test Pipeline"
|
||||
mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline)
|
||||
|
||||
config_mock = Mock()
|
||||
config_mock.indexing_technique = "high_quality"
|
||||
config_mock.embedding_model = "m"
|
||||
config_mock.embedding_model_provider = "p"
|
||||
config_mock.chunk_structure = "text_model"
|
||||
config_mock.retrieval_model.model_dump.return_value = {}
|
||||
config_mock.summary_index_setting = None
|
||||
mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate",
|
||||
return_value=config_mock,
|
||||
)
|
||||
|
||||
dataset_mock = Mock(id="d1")
|
||||
binding_mock = Mock(id="b1")
|
||||
mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset_mock)
|
||||
binding_cls = mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_dsl_service.DatasetCollectionBinding",
|
||||
return_value=binding_mock,
|
||||
)
|
||||
mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock())
|
||||
|
||||
session = cast(MagicMock, Mock())
|
||||
session.scalar.return_value = None
|
||||
session.scalars.return_value.all.return_value = []
|
||||
service = RagPipelineDslService(session=cast(Session, session))
|
||||
account = Mock(current_tenant_id="t1")
|
||||
|
||||
result = service.import_rag_pipeline(account=account, import_mode="yaml-content", yaml_content=yaml_content)
|
||||
|
||||
assert result.status == ImportStatus.COMPLETED
|
||||
binding_cls.assert_called_once()
|
||||
assert dataset_mock.collection_binding_id == "b1"
|
||||
session.commit.assert_not_called()
|
||||
assert session.flush.call_count >= 2
|
||||
|
||||
|
||||
def test_import_rag_pipeline_pending_version(mocker) -> None:
|
||||
@@ -338,6 +392,67 @@ workflow:
|
||||
assert result.dataset_id == "d1"
|
||||
|
||||
|
||||
def test_confirm_import_flushes_new_collection_binding_without_commit(mocker) -> None:
|
||||
from services.rag_pipeline.rag_pipeline_dsl_service import RagPipelinePendingData
|
||||
|
||||
yaml_content = """
|
||||
version: 0.1.0
|
||||
kind: rag_pipeline
|
||||
rag_pipeline:
|
||||
name: Test Pipeline
|
||||
workflow:
|
||||
graph:
|
||||
nodes:
|
||||
- data:
|
||||
type: knowledge-index
|
||||
"""
|
||||
pending = RagPipelinePendingData(import_mode="yaml-content", yaml_content=yaml_content, pipeline_id="p1")
|
||||
mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_dsl_service.redis_client.get",
|
||||
return_value=pending.model_dump_json(),
|
||||
)
|
||||
mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.redis_client.delete")
|
||||
|
||||
pipeline = Mock(id="p1", description="desc")
|
||||
pipeline.name = "Test Pipeline"
|
||||
pipeline.retrieve_dataset.return_value = None
|
||||
mocker.patch.object(RagPipelineDslService, "_create_or_update_pipeline", return_value=pipeline)
|
||||
|
||||
config_mock = Mock()
|
||||
config_mock.indexing_technique = "high_quality"
|
||||
config_mock.embedding_model = "m"
|
||||
config_mock.embedding_model_provider = "p"
|
||||
config_mock.chunk_structure = "text_model"
|
||||
config_mock.retrieval_model.model_dump.return_value = {}
|
||||
config_mock.summary_index_setting = None
|
||||
mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_dsl_service.KnowledgeConfiguration.model_validate",
|
||||
return_value=config_mock,
|
||||
)
|
||||
|
||||
dataset_mock = Mock(id="d1")
|
||||
binding_mock = Mock(id="b1")
|
||||
mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.Dataset", return_value=dataset_mock)
|
||||
binding_cls = mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_dsl_service.DatasetCollectionBinding",
|
||||
return_value=binding_mock,
|
||||
)
|
||||
mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.select", return_value=MagicMock())
|
||||
|
||||
session = cast(MagicMock, Mock())
|
||||
session.scalar.side_effect = [pipeline, None]
|
||||
service = RagPipelineDslService(session=cast(Session, session))
|
||||
account = Mock(id="u1", current_tenant_id="t1")
|
||||
|
||||
result = service.confirm_import(account=account, import_id="imp-1")
|
||||
|
||||
assert result.status == ImportStatus.COMPLETED
|
||||
binding_cls.assert_called_once()
|
||||
assert dataset_mock.collection_binding_id == "b1"
|
||||
session.commit.assert_not_called()
|
||||
assert session.flush.call_count >= 2
|
||||
|
||||
|
||||
# --- _extract_dependencies_from_workflow_graph all types ---
|
||||
|
||||
|
||||
@@ -421,6 +536,8 @@ def test_create_or_update_pipeline_create_new(mocker) -> None:
|
||||
|
||||
assert result == pipeline_instance
|
||||
session.add.assert_called()
|
||||
session.commit.assert_not_called()
|
||||
session.flush.assert_called()
|
||||
|
||||
|
||||
# --- export_rag_pipeline_dsl comprehensive ---
|
||||
|
||||
Reference in New Issue
Block a user