test: use SQLite sessions in rag pipeline transform (#39050)

This commit is contained in:
Asuka Minato 2026-07-30 11:43:48 +09:00 committed by GitHub
parent 87aa982185
commit b855d71adb
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 164 additions and 160 deletions

View File

@ -29,6 +29,11 @@ logger = logging.getLogger(__name__)
class RagPipelineTransformService:
def transform_dataset(self, dataset_id: str, session: Session):
"""Transform a vendor dataset within the caller-owned transaction.
Dataset and document state is read through ``session`` so uncommitted caller changes remain visible. The
transformation commits only after the pipeline and migrated document metadata have been persisted.
"""
dataset = session.get(Dataset, dataset_id)
if not dataset:
raise ValueError("Dataset not found")
@ -46,7 +51,7 @@ class RagPipelineTransformService:
if not datasource_type and not indexing_technique:
return self._transform_to_empty_pipeline(dataset, session=session)
doc_form = dataset.doc_form
doc_form = dataset.get_doc_form(session=session)
if not doc_form:
return self._transform_to_empty_pipeline(dataset, session=session)
retrieval_model = RetrievalSetting.model_validate(dataset.retrieval_model) if dataset.retrieval_model else None

View File

@ -5,12 +5,47 @@ from typing import cast
import pytest
from pytest_mock import MockerFixture
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.dataset import Dataset
from extensions.storage.storage_type import StorageType
from models.dataset import Dataset, Document, DocumentPipelineExecutionLog, Pipeline
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom
from models.model import UploadFile
from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration
from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService
def _dataset(**overrides: object) -> Dataset:
values = {
"id": "dataset-1",
"tenant_id": "tenant-1",
"name": "Dataset",
"description": "desc",
"created_by": "user-1",
"provider": "vendor",
}
values.update(overrides)
return Dataset(**values)
def _document(**overrides: object) -> Document:
values = {
"id": "document-1",
"tenant_id": "tenant-1",
"dataset_id": "dataset-1",
"position": 1,
"data_source_type": DataSourceType.UPLOAD_FILE,
"data_source_info": None,
"batch": "batch-1",
"name": "Document",
"created_from": DocumentCreatedFrom.WEB,
"created_by": "user-1",
}
values.update(overrides)
return Document(**values)
@pytest.mark.parametrize(
("doc_form", "datasource_type", "indexing_technique"),
[
@ -92,129 +127,93 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi
install_mock.assert_called_once_with("tenant-1", ["missing-plugin:1.0.0"])
def test_transform_to_empty_pipeline_updates_dataset_and_commits(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset, Pipeline)], indirect=True)
def test_transform_to_empty_pipeline_updates_dataset_and_commits(
mocker: MockerFixture, sqlite_session: Session
) -> None:
service = RagPipelineTransformService()
mocker.patch(
"services.rag_pipeline.rag_pipeline_transform_service.current_user",
SimpleNamespace(id="user-1"),
)
class FakePipeline:
def __init__(self, **kwargs):
self.id = "pipeline-1"
self.tenant_id = kwargs["tenant_id"]
self.name = kwargs["name"]
self.description = kwargs["description"]
self.created_by = kwargs["created_by"]
dataset = _dataset()
sqlite_session.add(dataset)
sqlite_session.commit()
mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.Pipeline", FakePipeline)
session_mock = mocker.Mock()
add_mock = session_mock.add
flush_mock = session_mock.flush
commit_mock = session_mock.commit
result = service._transform_to_empty_pipeline(dataset, session=sqlite_session)
dataset = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
name="Dataset",
description="desc",
pipeline_id=None,
runtime_mode="general",
updated_by=None,
updated_at=None,
)
result = service._transform_to_empty_pipeline(cast(Dataset, dataset), session=session_mock)
assert result == {"pipeline_id": "pipeline-1", "dataset_id": "dataset-1", "status": "success"}
assert dataset.pipeline_id == "pipeline-1"
pipeline = sqlite_session.get(Pipeline, result["pipeline_id"])
assert pipeline is not None
assert pipeline.name == "Dataset"
assert result == {"pipeline_id": pipeline.id, "dataset_id": "dataset-1", "status": "success"}
assert dataset.pipeline_id == pipeline.id
assert dataset.runtime_mode == "rag_pipeline"
assert dataset.updated_by == "user-1"
add_mock.assert_called()
flush_mock.assert_called_once()
commit_mock.assert_called_once()
# --- transform_dataset ---
def test_transform_dataset_returns_early_when_pipeline_exists(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_returns_early_when_pipeline_exists(sqlite_session: Session) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
id="d1",
pipeline_id="p1",
runtime_mode="rag_pipeline",
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
dataset = _dataset(id="d1", pipeline_id="p1", runtime_mode="rag_pipeline")
sqlite_session.add(dataset)
sqlite_session.commit()
result = service.transform_dataset("d1", session_mock)
result = service.transform_dataset("d1", sqlite_session)
assert result == {"pipeline_id": "p1", "dataset_id": "d1", "status": "success"}
def test_transform_dataset_raises_for_dataset_not_found(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_raises_for_dataset_not_found(sqlite_session: Session) -> None:
service = RagPipelineTransformService()
session_mock = mocker.Mock()
session_mock.get.return_value = None
with pytest.raises(ValueError, match="Dataset not found"):
service.transform_dataset("d1", session_mock)
service.transform_dataset("d1", sqlite_session)
def test_transform_dataset_raises_for_external_dataset(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_raises_for_external_dataset(sqlite_session: Session) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
id="d1",
pipeline_id=None,
runtime_mode=None,
provider="external",
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
sqlite_session.add(_dataset(id="d1", provider="external"))
sqlite_session.commit()
with pytest.raises(ValueError, match="External dataset is not supported"):
service.transform_dataset("d1", session_mock)
service.transform_dataset("d1", sqlite_session)
def test_transform_dataset_calls_empty_pipeline_when_no_datasource(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_calls_empty_pipeline_when_no_datasource(
mocker: MockerFixture, sqlite_session: Session
) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
id="d1",
pipeline_id=None,
runtime_mode=None,
provider="vendor",
data_source_type=None,
indexing_technique=None,
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
sqlite_session.add(_dataset(id="d1", data_source_type=None, indexing_technique=None))
sqlite_session.commit()
empty_result = {"pipeline_id": "p-empty", "dataset_id": "d1", "status": "success"}
mocker.patch.object(service, "_transform_to_empty_pipeline", return_value=empty_result)
result = service.transform_dataset("d1", session_mock)
result = service.transform_dataset("d1", sqlite_session)
assert result == empty_result
def test_transform_dataset_calls_empty_pipeline_when_no_doc_form(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset, Document)], indirect=True)
def test_transform_dataset_calls_empty_pipeline_when_no_doc_form(
mocker: MockerFixture, sqlite_session: Session
) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
id="d1",
pipeline_id=None,
runtime_mode=None,
provider="vendor",
data_source_type="upload_file",
indexing_technique="high_quality",
doc_form=None,
sqlite_session.add(
_dataset(id="d1", data_source_type="upload_file", indexing_technique="high_quality", chunk_structure=None)
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
sqlite_session.commit()
empty_result = {"pipeline_id": "p-empty", "dataset_id": "d1", "status": "success"}
mocker.patch.object(service, "_transform_to_empty_pipeline", return_value=empty_result)
result = service.transform_dataset("d1", session_mock)
result = service.transform_dataset("d1", sqlite_session)
assert result == empty_result
@ -274,78 +273,65 @@ def test_deal_knowledge_index_high_quality_sets_embedding(mocker: MockerFixture)
# --- _deal_document_data ---
def test_deal_document_data_notion(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog)], indirect=True)
def test_deal_document_data_notion(sqlite_session: Session) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(id="d1", pipeline_id="p1")
doc = SimpleNamespace(
dataset = _dataset(id="d1", pipeline_id="p1")
doc = _document(
id="doc1",
dataset_id="d1",
data_source_type="notion_import",
data_source_info_dict={
"notion_workspace_id": "ws1",
"notion_page_id": "page1",
"notion_page_icon": "icon1",
"type": "page",
"last_edited_time": 12345,
},
data_source_info=(
'{"notion_workspace_id":"ws1","notion_page_id":"page1","notion_page_icon":"icon1",'
'"type":"page","last_edited_time":12345}'
),
name="Notion Doc",
created_by="u1",
created_at=datetime.now(UTC).replace(tzinfo=None),
data_source_info=None,
)
sqlite_session.add(doc)
sqlite_session.commit()
scalars_mock = mocker.Mock()
scalars_mock.all.return_value = [doc]
session_mock = mocker.Mock()
session_mock.scalars.return_value = scalars_mock
add_mock = session_mock.add
service._deal_document_data(cast(Dataset, dataset), session_mock)
service._deal_document_data(dataset, sqlite_session)
sqlite_session.flush()
assert doc.data_source_type == "online_document"
assert "page1" in doc.data_source_info
assert add_mock.call_count == 2 # document + log
log = sqlite_session.scalar(select(DocumentPipelineExecutionLog))
assert log is not None
assert log.document_id == doc.id
@pytest.mark.parametrize(("provider", "node_id"), [("firecrawl", "1752565402678"), ("jinareader", "1752491761974")])
def test_deal_document_data_website(mocker: MockerFixture, provider: str, node_id: str) -> None:
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog)], indirect=True)
def test_deal_document_data_website(sqlite_session: Session, provider: str, node_id: str) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(id="d1", pipeline_id="p1")
doc = SimpleNamespace(
dataset = _dataset(id="d1", pipeline_id="p1")
doc = _document(
id="doc1",
dataset_id="d1",
data_source_type="website_crawl",
data_source_info_dict={
"url": "https://example.com",
"provider": provider,
},
data_source_info=f'{{"url":"https://example.com","provider":"{provider}"}}',
name="Web Doc",
created_by="u1",
created_at=datetime.now(UTC).replace(tzinfo=None),
data_source_info=None,
)
sqlite_session.add(doc)
sqlite_session.commit()
scalars_mock = mocker.Mock()
scalars_mock.all.return_value = [doc]
session_mock = mocker.Mock()
session_mock.scalars.return_value = scalars_mock
add_mock = session_mock.add
service._deal_document_data(cast(Dataset, dataset), session_mock)
service._deal_document_data(dataset, sqlite_session)
sqlite_session.flush()
assert doc.data_source_type == "website_crawl"
assert "example.com" in doc.data_source_info
# Check if correct node id was used in log
log = add_mock.call_args_list[1][0][0]
log = sqlite_session.scalar(select(DocumentPipelineExecutionLog))
assert log is not None
assert log.datasource_node_id == node_id
# --- transform_dataset complex flow ---
def test_transform_dataset_full_flow(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_full_flow(mocker: MockerFixture, sqlite_session: Session) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
dataset = _dataset(
id="d1",
tenant_id="t1",
name="D",
@ -355,20 +341,18 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None:
provider="vendor",
data_source_type="upload_file",
indexing_technique="high_quality",
doc_form="text_model",
chunk_structure="text_model",
retrieval_model={"search_method": "semantic_search", "top_k": 3},
embedding_model="m1",
embedding_model_provider="p1",
summary_index_setting=None,
chunk_structure=None,
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
sqlite_session.add(dataset)
sqlite_session.commit()
mocker.patch.object(service, "_deal_dependencies")
mocker.patch.object(service, "_deal_document_data")
session_mock.commit = mocker.Mock()
# Mock current_user to have the same tenant_id as dataset
mock_current_user = SimpleNamespace(current_tenant_id="t1")
@ -377,16 +361,19 @@ def test_transform_dataset_full_flow(mocker: MockerFixture) -> None:
pipeline = SimpleNamespace(id="p-new")
mocker.patch.object(service, "_create_pipeline", return_value=pipeline)
result = service.transform_dataset("d1", session_mock)
result = service.transform_dataset("d1", sqlite_session)
assert result["pipeline_id"] == "p-new"
assert dataset.runtime_mode == "rag_pipeline"
assert dataset.chunk_structure == "text_model"
def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create(
mocker: MockerFixture, sqlite_session: Session
) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
dataset = _dataset(
id="d1",
tenant_id="t1",
name="D",
@ -396,22 +383,25 @@ def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create
provider="vendor",
data_source_type="upload_file",
indexing_technique="high_quality",
doc_form="unsupported",
chunk_structure="unsupported",
retrieval_model=None,
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
sqlite_session.add(dataset)
sqlite_session.commit()
mocker.patch.object(service, "_get_transform_yaml", return_value={"workflow": {"graph": {"nodes": []}}})
mocker.patch.object(service, "_deal_dependencies")
mocker.patch.object(service, "_create_pipeline", return_value=SimpleNamespace(id="p-new"))
with pytest.raises(ValueError, match="Unsupported doc form"):
service.transform_dataset("d1", session_mock)
service.transform_dataset("d1", sqlite_session)
def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
def test_transform_dataset_raises_when_transform_yaml_missing_workflow(
mocker: MockerFixture, sqlite_session: Session
) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(
dataset = _dataset(
id="d1",
tenant_id="t1",
name="D",
@ -421,53 +411,62 @@ def test_transform_dataset_raises_when_transform_yaml_missing_workflow(mocker: M
provider="vendor",
data_source_type="upload_file",
indexing_technique="high_quality",
doc_form="text_model",
chunk_structure="text_model",
retrieval_model=None,
)
session_mock = mocker.Mock()
session_mock.get.return_value = dataset
sqlite_session.add(dataset)
sqlite_session.commit()
mocker.patch.object(service, "_get_transform_yaml", return_value={})
mocker.patch.object(service, "_deal_dependencies")
with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"):
service.transform_dataset("d1", session_mock)
service.transform_dataset("d1", sqlite_session)
def test_create_pipeline_raises_when_workflow_data_missing(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
def test_create_pipeline_raises_when_workflow_data_missing(sqlite_session: Session) -> None:
service = RagPipelineTransformService()
session = mocker.Mock()
with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"):
service._create_pipeline({"rag_pipeline": {"name": "N"}}, session=session)
service._create_pipeline({"rag_pipeline": {"name": "N"}}, session=sqlite_session)
def test_deal_document_data_upload_file_with_existing_file(mocker: MockerFixture) -> None:
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog, UploadFile)], indirect=True)
def test_deal_document_data_upload_file_with_existing_file(sqlite_session: Session) -> None:
service = RagPipelineTransformService()
dataset = SimpleNamespace(id="d1", pipeline_id="p1")
document = SimpleNamespace(
dataset = _dataset(id="d1", pipeline_id="p1")
document = _document(
id="doc-1",
dataset_id="d1",
data_source_type="upload_file",
data_source_info_dict={"upload_file_id": "file-1"},
data_source_info='{"upload_file_id":"file-1"}',
name="Doc",
created_by="u1",
created_at=datetime.now(UTC).replace(tzinfo=None),
data_source_info=None,
)
upload_file = SimpleNamespace(name="f.txt", size=10, extension="txt", mime_type="text/plain")
upload_file = UploadFile(
tenant_id="tenant-1",
storage_type=StorageType.LOCAL,
key="files/f.txt",
name="f.txt",
size=10,
extension="txt",
mime_type="text/plain",
created_by_role=CreatorUserRole.ACCOUNT,
created_by="user-1",
created_at=datetime.now(UTC).replace(tzinfo=None),
used=False,
)
upload_file.id = "file-1"
sqlite_session.add_all([document, upload_file])
sqlite_session.commit()
scalars_mock = mocker.Mock()
scalars_mock.all.return_value = [document]
session_mock = mocker.Mock()
session_mock.scalars.return_value = scalars_mock
session_mock.get.return_value = upload_file
add_mock = session_mock.add
service._deal_document_data(cast(Dataset, dataset), session_mock)
service._deal_document_data(dataset, sqlite_session)
sqlite_session.flush()
assert document.data_source_type == "local_file"
assert "real_file_id" in document.data_source_info
assert add_mock.call_count >= 2
log = sqlite_session.scalar(select(DocumentPipelineExecutionLog))
assert log is not None
assert log.document_id == document.id
def _make_service():