mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix(api): bind remaining RAG resources to owners (#40503)
This commit is contained in:
+1
-1
@@ -36,7 +36,7 @@ def test_export_customized_pipeline_template_from_database(
|
||||
db_session_with_containers.expire_all()
|
||||
|
||||
with flask_app_with_containers.test_request_context("/"):
|
||||
response, status = method(api, template.id)
|
||||
response, status = method(api, db_session_with_containers, template.tenant_id, template.id)
|
||||
|
||||
assert status == 200
|
||||
assert response == {"data": "yaml-data"}
|
||||
|
||||
+193
-98
@@ -1,15 +1,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterator
|
||||
from inspect import unwrap
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import PropertyMock, patch
|
||||
from inspect import getclosurevars, unwrap
|
||||
from unittest.mock import ANY, PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import Engine
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import NotFound
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.datasets.rag_pipeline import rag_pipeline as module
|
||||
@@ -22,10 +21,12 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline import (
|
||||
PipelineTemplateListQuery,
|
||||
PublishCustomizedPipelineTemplateApi,
|
||||
)
|
||||
from models.account import Account
|
||||
from models.dataset import PipelineCustomizedTemplate
|
||||
from models.account import Account, TenantAccountRole
|
||||
from models.dataset import Pipeline, PipelineCustomizedTemplate
|
||||
from models.engine import db
|
||||
from services.entities.knowledge_entities.rag_pipeline_entities import PipelineTemplateInfoEntity
|
||||
from services.errors.account import NoPermissionError
|
||||
from services.errors.rag_pipeline import RagPipelineResourceNotFoundError
|
||||
|
||||
|
||||
def _template_item() -> dict[str, object]:
|
||||
@@ -65,6 +66,12 @@ def _account() -> Account:
|
||||
return account
|
||||
|
||||
|
||||
def _pipeline() -> Pipeline:
|
||||
pipeline = Pipeline(tenant_id="tenant-1", name="Pipeline")
|
||||
pipeline.id = "pipeline-1"
|
||||
return pipeline
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def database_app() -> Iterator[Flask]:
|
||||
app = Flask(__name__)
|
||||
@@ -136,11 +143,13 @@ class TestPipelineTemplateDetailApi:
|
||||
def test_get_serializes_template_detail(self, app: Flask, sqlite_engine: Engine) -> None:
|
||||
api = PipelineTemplateDetailApi()
|
||||
method = unwrap(api.get)
|
||||
service_calls: list[tuple[str, str]] = []
|
||||
service_calls: list[tuple[str, str, str]] = []
|
||||
|
||||
def get_pipeline_template_detail(template_id: str, type: str, *, session) -> dict[str, object]:
|
||||
def get_pipeline_template_detail(
|
||||
template_id: str, current_tenant_id: str, type: str, *, session
|
||||
) -> dict[str, object]:
|
||||
del session
|
||||
service_calls.append((template_id, type))
|
||||
service_calls.append((template_id, current_tenant_id, type))
|
||||
return _template_detail()
|
||||
|
||||
with (
|
||||
@@ -152,18 +161,20 @@ class TestPipelineTemplateDetailApi:
|
||||
side_effect=get_pipeline_template_detail,
|
||||
),
|
||||
):
|
||||
response, status = method(api, PipelineTemplateDetailQuery(type="customized"), session, "template-1")
|
||||
response, status = method(
|
||||
api, PipelineTemplateDetailQuery(type="customized"), session, "tenant-1", "template-1"
|
||||
)
|
||||
|
||||
assert status == 200
|
||||
assert response == {**_template_detail(), "created_by": None}
|
||||
assert service_calls == [("template-1", "customized")]
|
||||
assert service_calls == [("template-1", "tenant-1", "customized")]
|
||||
|
||||
def test_get_raises_not_found_without_custom_response_body(self, app: Flask, sqlite_engine: Engine) -> None:
|
||||
api = PipelineTemplateDetailApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
def get_pipeline_template_detail(template_id: str, type: str, *, session) -> None:
|
||||
del template_id, type, session
|
||||
def get_pipeline_template_detail(template_id: str, current_tenant_id: str, type: str, *, session) -> None:
|
||||
del template_id, current_tenant_id, type, session
|
||||
|
||||
with (
|
||||
Session(sqlite_engine) as session,
|
||||
@@ -175,7 +186,7 @@ class TestPipelineTemplateDetailApi:
|
||||
),
|
||||
pytest.raises(NotFound),
|
||||
):
|
||||
method(api, PipelineTemplateDetailQuery(), session, "missing")
|
||||
method(api, PipelineTemplateDetailQuery(), session, "tenant-1", "missing")
|
||||
|
||||
|
||||
class TestCustomizedPipelineTemplateApi:
|
||||
@@ -286,9 +297,7 @@ class TestCustomizedPipelineTemplateApi:
|
||||
assert deleted_templates == [("template-1", tenant_id)]
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(PipelineCustomizedTemplate,)], indirect=True)
|
||||
def test_post_exports_yaml_from_orm_template(
|
||||
self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
|
||||
) -> None:
|
||||
def test_post_exports_yaml_from_orm_template(self, app: Flask, sqlite_session: Session) -> None:
|
||||
api = CustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
template = PipelineCustomizedTemplate(
|
||||
@@ -306,115 +315,201 @@ class TestCustomizedPipelineTemplateApi:
|
||||
template.id = "template-1"
|
||||
sqlite_session.add(template)
|
||||
sqlite_session.commit()
|
||||
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine))
|
||||
|
||||
with app.test_request_context("/rag/pipeline/customized/templates/template-1", method="POST"):
|
||||
response, status = method(api, "template-1")
|
||||
response, status = method(
|
||||
api,
|
||||
sqlite_session,
|
||||
"00000000-0000-0000-0000-000000000001",
|
||||
"template-1",
|
||||
)
|
||||
|
||||
assert status == 200
|
||||
assert response == {"data": "dsl: value"}
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(PipelineCustomizedTemplate,)], indirect=True)
|
||||
def test_post_raises_when_template_is_missing(
|
||||
self, app: Flask, monkeypatch: pytest.MonkeyPatch, sqlite_engine: Engine, sqlite_session: Session
|
||||
) -> None:
|
||||
def test_post_returns_not_found_for_other_tenant(self, app: Flask, sqlite_session: Session) -> None:
|
||||
api = CustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
assert sqlite_session.get(PipelineCustomizedTemplate, "missing") is None
|
||||
monkeypatch.setattr(module, "db", SimpleNamespace(engine=sqlite_engine))
|
||||
template = PipelineCustomizedTemplate(
|
||||
tenant_id="00000000-0000-0000-0000-000000000002",
|
||||
name="Other tenant template",
|
||||
description="Description",
|
||||
chunk_structure="general",
|
||||
icon={},
|
||||
position=1,
|
||||
yaml_content="secret: value",
|
||||
install_count=0,
|
||||
language="en-US",
|
||||
created_by="00000000-0000-0000-0000-000000000003",
|
||||
)
|
||||
template.id = "template-1"
|
||||
sqlite_session.add(template)
|
||||
sqlite_session.commit()
|
||||
|
||||
with app.test_request_context("/rag/pipeline/customized/templates/missing", method="POST"):
|
||||
with pytest.raises(ValueError, match="Customized pipeline template not found"):
|
||||
method(api, "missing")
|
||||
with (
|
||||
app.test_request_context("/rag/pipeline/customized/templates/template-1", method="POST"),
|
||||
pytest.raises(NotFound, match="Customized pipeline template not found"),
|
||||
):
|
||||
method(
|
||||
api,
|
||||
sqlite_session,
|
||||
"00000000-0000-0000-0000-000000000001",
|
||||
"template-1",
|
||||
)
|
||||
|
||||
|
||||
class TestPublishCustomizedPipelineTemplateApi:
|
||||
def test_post_validates_payload_and_returns_empty_204(self, app: Flask) -> None:
|
||||
def test_post_uses_pipeline_release_rbac_scene(self) -> None:
|
||||
method = PublishCustomizedPipelineTemplateApi.post
|
||||
while "scene" not in getclosurevars(method).nonlocals:
|
||||
method = method.__wrapped__
|
||||
|
||||
assert getclosurevars(method).nonlocals["scene"] == module.RBACPermission.DATASET_PIPELINE_RELEASE
|
||||
|
||||
def test_post_validates_payload_and_returns_empty_204(self) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
payload = _payload()
|
||||
account = _account()
|
||||
tenant_id = "tenant-1"
|
||||
service_calls: list[tuple[str, dict[str, object], Account, str]] = []
|
||||
|
||||
class Service:
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
def publish_customized_pipeline_template(
|
||||
self,
|
||||
pipeline_id: str,
|
||||
data: dict[str, object],
|
||||
current_user: Account,
|
||||
current_tenant_id: str,
|
||||
*,
|
||||
session,
|
||||
) -> None:
|
||||
del session
|
||||
service_calls.append((pipeline_id, data, current_user, current_tenant_id))
|
||||
pipeline = _pipeline()
|
||||
dataset = object()
|
||||
|
||||
with (
|
||||
app.test_request_context("/rag/pipelines/pipeline-1/customized/publish", method="POST", json=payload),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch.object(module, "RagPipelineService", Service),
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", True),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=dataset),
|
||||
patch.object(module.DatasetService, "check_dataset_permission") as legacy_acl,
|
||||
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
|
||||
):
|
||||
response, status = method(
|
||||
api, CustomizedPipelineTemplatePayload.model_validate(payload), tenant_id, account, "pipeline-1"
|
||||
)
|
||||
response, status = method(api, CustomizedPipelineTemplatePayload.model_validate(payload), account, pipeline)
|
||||
|
||||
assert (response, status) == ("", 204)
|
||||
assert service_calls == [("pipeline-1", payload, account, tenant_id)]
|
||||
publish.assert_called_once_with(pipeline, dataset, payload, account, session=ANY)
|
||||
legacy_acl.assert_not_called()
|
||||
|
||||
def test_post_allows_missing_icon_info_for_publish_service_fallback(self, app: Flask) -> None:
|
||||
@pytest.mark.parametrize(
|
||||
("payload", "expected_icon_info"),
|
||||
[
|
||||
(
|
||||
{"name": "Published template", "description": "Description"},
|
||||
{"icon": "", "icon_background": None, "icon_type": None, "icon_url": None},
|
||||
),
|
||||
({"name": "Published template", "description": "Description", "icon_info": {}}, {}),
|
||||
],
|
||||
)
|
||||
def test_post_preserves_valid_icon_info(
|
||||
self,
|
||||
payload: dict[str, object],
|
||||
expected_icon_info: dict[str, object | None],
|
||||
) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
payload: dict[str, object] = {
|
||||
"name": "Published template",
|
||||
"description": "Description",
|
||||
}
|
||||
account = _account()
|
||||
tenant_id = "tenant-1"
|
||||
service_calls: list[tuple[str, dict[str, object], Account, str]] = []
|
||||
|
||||
class Service:
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
def publish_customized_pipeline_template(
|
||||
self,
|
||||
pipeline_id: str,
|
||||
data: dict[str, object],
|
||||
current_user: Account,
|
||||
current_tenant_id: str,
|
||||
*,
|
||||
session,
|
||||
) -> None:
|
||||
del session
|
||||
service_calls.append((pipeline_id, data, current_user, current_tenant_id))
|
||||
pipeline = _pipeline()
|
||||
dataset = object()
|
||||
|
||||
with (
|
||||
app.test_request_context("/rag/pipelines/pipeline-1/customized/publish", method="POST", json=payload),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch.object(module, "RagPipelineService", Service),
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", True),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=dataset),
|
||||
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
|
||||
):
|
||||
response, status = method(
|
||||
api, CustomizedPipelineTemplatePayload.model_validate(payload), tenant_id, account, "pipeline-1"
|
||||
)
|
||||
response, status = method(api, CustomizedPipelineTemplatePayload.model_validate(payload), account, pipeline)
|
||||
|
||||
assert (response, status) == ("", 204)
|
||||
assert service_calls == [
|
||||
(
|
||||
"pipeline-1",
|
||||
{
|
||||
**payload,
|
||||
"icon_info": {
|
||||
"icon": "",
|
||||
"icon_background": None,
|
||||
"icon_type": None,
|
||||
"icon_url": None,
|
||||
},
|
||||
},
|
||||
account,
|
||||
tenant_id,
|
||||
)
|
||||
]
|
||||
publish.assert_called_once_with(pipeline, dataset, ANY, account, session=ANY)
|
||||
assert publish.call_args.args[2]["icon_info"] == expected_icon_info
|
||||
|
||||
def test_post_translates_missing_owned_resource_to_not_found(self) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
payload = _payload()
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", True),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=object()),
|
||||
patch.object(
|
||||
module.RagPipelineService,
|
||||
"publish_customized_pipeline_template",
|
||||
side_effect=RagPipelineResourceNotFoundError("Workflow not found"),
|
||||
),
|
||||
pytest.raises(NotFound, match="Workflow not found"),
|
||||
):
|
||||
method(api, CustomizedPipelineTemplatePayload.model_validate(payload), _account(), _pipeline())
|
||||
|
||||
def test_post_allows_legacy_dataset_operator_after_dataset_acl(self) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
account = _account()
|
||||
account.role = TenantAccountRole.DATASET_OPERATOR
|
||||
pipeline = _pipeline()
|
||||
dataset = object()
|
||||
payload = _payload()
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=dataset),
|
||||
patch.object(module.DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
|
||||
):
|
||||
response = method(api, CustomizedPipelineTemplatePayload.model_validate(payload), account, pipeline)
|
||||
|
||||
assert response == ("", 204)
|
||||
assert check_permission.call_args.args[:2] == (dataset, account)
|
||||
publish.assert_called_once_with(pipeline, dataset, payload, account, session=ANY)
|
||||
|
||||
def test_post_rejects_legacy_non_editor_before_dataset_acl(self) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
account = _account()
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=object()),
|
||||
patch.object(module.DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
|
||||
pytest.raises(Forbidden),
|
||||
):
|
||||
method(api, CustomizedPipelineTemplatePayload.model_validate(_payload()), account, _pipeline())
|
||||
|
||||
check_permission.assert_not_called()
|
||||
publish.assert_not_called()
|
||||
|
||||
def test_post_rejects_legacy_dataset_acl_before_publish(self) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
account = _account()
|
||||
account.role = TenantAccountRole.EDITOR
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=object()),
|
||||
patch.object(
|
||||
module.DatasetService,
|
||||
"check_dataset_permission",
|
||||
side_effect=NoPermissionError("Dataset is private"),
|
||||
),
|
||||
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
|
||||
pytest.raises(Forbidden, match="Dataset is private"),
|
||||
):
|
||||
method(api, CustomizedPipelineTemplatePayload.model_validate(_payload()), account, _pipeline())
|
||||
|
||||
publish.assert_not_called()
|
||||
|
||||
def test_post_rejects_missing_legacy_dataset_before_publish(self) -> None:
|
||||
api = PublishCustomizedPipelineTemplateApi()
|
||||
method = unwrap(api.post)
|
||||
account = _account()
|
||||
account.role = TenantAccountRole.EDITOR
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(Pipeline, "retrieve_dataset", return_value=None),
|
||||
patch.object(module.DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(module.RagPipelineService, "publish_customized_pipeline_template") as publish,
|
||||
pytest.raises(NotFound, match="Dataset not found"),
|
||||
):
|
||||
method(api, CustomizedPipelineTemplatePayload.model_validate(_payload()), account, _pipeline())
|
||||
|
||||
check_permission.assert_not_called()
|
||||
publish.assert_not_called()
|
||||
|
||||
+129
-10
@@ -13,7 +13,7 @@ import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import Engine
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import Forbidden
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
from controllers.console.datasets.rag_pipeline import rag_pipeline_workflow as module
|
||||
from controllers.console.datasets.rag_pipeline.rag_pipeline_workflow import (
|
||||
@@ -25,18 +25,21 @@ from controllers.console.datasets.rag_pipeline.rag_pipeline_workflow import (
|
||||
WorkflowUpdatePayload,
|
||||
)
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from models.account import Account, TenantAccountRole
|
||||
from models.dataset import Pipeline
|
||||
from models.account import Account, Tenant, TenantAccountRole
|
||||
from models.dataset import Dataset, Pipeline
|
||||
from models.engine import db
|
||||
from models.enums import PermissionEnum
|
||||
from models.tools import WorkflowToolProvider
|
||||
from models.workflow import Workflow, WorkflowType
|
||||
from services.errors.llm import InvokeRateLimitError
|
||||
from services.errors.rag_pipeline import RagPipelineResourceNotFoundError
|
||||
from services.rag_pipeline.rag_pipeline import RagPipelineService
|
||||
|
||||
DEFAULT_WORKFLOW_TENANT_ID = "00000000-0000-0000-0000-000000000001"
|
||||
DEFAULT_WORKFLOW_APP_ID = "00000000-0000-0000-0000-000000000002"
|
||||
DEFAULT_WORKFLOW_CREATED_BY = "00000000-0000-0000-0000-000000000003"
|
||||
DEFAULT_WORKFLOW_ID = "00000000-0000-0000-0000-000000000004"
|
||||
DEFAULT_DATASET_ID = "44444444-4444-4444-4444-444444444444"
|
||||
|
||||
|
||||
def _make_workflow(**overrides: object) -> Workflow:
|
||||
@@ -67,6 +70,9 @@ def _account() -> Account:
|
||||
account = Account(name="Alice", email="alice@example.com")
|
||||
account.id = DEFAULT_WORKFLOW_CREATED_BY
|
||||
account.role = TenantAccountRole.EDITOR
|
||||
tenant = Tenant(name="Tenant")
|
||||
tenant.id = DEFAULT_WORKFLOW_TENANT_ID
|
||||
account._current_tenant = tenant
|
||||
return account
|
||||
|
||||
|
||||
@@ -76,6 +82,18 @@ def _pipeline() -> Pipeline:
|
||||
return pipeline
|
||||
|
||||
|
||||
def _dataset(*, tenant_id: str = DEFAULT_WORKFLOW_TENANT_ID, maintainer: str = DEFAULT_WORKFLOW_CREATED_BY) -> Dataset:
|
||||
return Dataset(
|
||||
id=DEFAULT_DATASET_ID,
|
||||
tenant_id=tenant_id,
|
||||
name="Dataset",
|
||||
created_by=maintainer,
|
||||
maintainer=maintainer,
|
||||
permission=PermissionEnum.ONLY_ME,
|
||||
provider="vendor",
|
||||
)
|
||||
|
||||
|
||||
def _persist_workflow(workflow: Workflow) -> None:
|
||||
db.session.add(workflow)
|
||||
db.session.commit()
|
||||
@@ -231,18 +249,119 @@ def test_rag_pipeline_recommended_plugins_serializes_known_envelope(database_app
|
||||
assert response == recommended_plugins
|
||||
|
||||
|
||||
def test_rag_pipeline_transform_rejects_read_only_member(app: Flask, sqlite_engine: Engine) -> None:
|
||||
def test_rag_pipeline_transform_rejects_read_only_member(sqlite_engine: Engine) -> None:
|
||||
account = _account()
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
api = module.RagPipelineTransformApi()
|
||||
handler = unwrap_all(api.post)
|
||||
|
||||
with (
|
||||
Session(sqlite_engine) as session,
|
||||
app.test_request_context("/"),
|
||||
pytest.raises(Forbidden),
|
||||
):
|
||||
handler(api, session, account, UUID("44444444-4444-4444-4444-444444444444"))
|
||||
with Session(sqlite_engine) as session:
|
||||
session.add(_dataset())
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
pytest.raises(Forbidden),
|
||||
):
|
||||
handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID))
|
||||
|
||||
|
||||
def test_rag_pipeline_transform_rejects_dataset_from_another_tenant_before_service_call(
|
||||
sqlite_engine: Engine,
|
||||
) -> None:
|
||||
api = module.RagPipelineTransformApi()
|
||||
handler = unwrap_all(api.post)
|
||||
|
||||
with Session(sqlite_engine) as session:
|
||||
session.add(_dataset(tenant_id="00000000-0000-0000-0000-000000000099"))
|
||||
|
||||
with (
|
||||
patch.object(module.RagPipelineTransformService, "transform_dataset") as transform_dataset,
|
||||
pytest.raises(NotFound),
|
||||
):
|
||||
handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, _account(), UUID(DEFAULT_DATASET_ID))
|
||||
|
||||
transform_dataset.assert_not_called()
|
||||
|
||||
|
||||
def test_rag_pipeline_transform_enforces_legacy_dataset_permission_before_service_call(
|
||||
sqlite_engine: Engine,
|
||||
) -> None:
|
||||
api = module.RagPipelineTransformApi()
|
||||
handler = unwrap_all(api.post)
|
||||
|
||||
with Session(sqlite_engine) as session:
|
||||
session.add(_dataset(maintainer="00000000-0000-0000-0000-000000000099"))
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(module.RagPipelineTransformService, "transform_dataset") as transform_dataset,
|
||||
pytest.raises(Forbidden),
|
||||
):
|
||||
handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, _account(), UUID(DEFAULT_DATASET_ID))
|
||||
|
||||
transform_dataset.assert_not_called()
|
||||
|
||||
|
||||
def test_rag_pipeline_transform_passes_authorized_dataset_and_account_to_service(
|
||||
sqlite_engine: Engine,
|
||||
) -> None:
|
||||
api = module.RagPipelineTransformApi()
|
||||
handler = unwrap_all(api.post)
|
||||
account = _account()
|
||||
expected = {"pipeline_id": "pipeline-1", "dataset_id": DEFAULT_DATASET_ID, "status": "success"}
|
||||
|
||||
with Session(sqlite_engine) as session:
|
||||
dataset = _dataset()
|
||||
session.add(dataset)
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(module.RagPipelineTransformService, "transform_dataset", return_value=expected) as transform,
|
||||
):
|
||||
response = handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID))
|
||||
|
||||
transform.assert_called_once_with(dataset, account.id, session)
|
||||
|
||||
assert response == expected
|
||||
|
||||
|
||||
def test_rag_pipeline_transform_maps_missing_pipeline_to_not_found(sqlite_engine: Engine) -> None:
|
||||
api = module.RagPipelineTransformApi()
|
||||
handler = unwrap_all(api.post)
|
||||
|
||||
with Session(sqlite_engine) as session:
|
||||
session.add(_dataset())
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", False),
|
||||
patch.object(
|
||||
module.RagPipelineTransformService,
|
||||
"transform_dataset",
|
||||
side_effect=RagPipelineResourceNotFoundError("Pipeline not found"),
|
||||
),
|
||||
pytest.raises(NotFound, match="Pipeline not found"),
|
||||
):
|
||||
handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, _account(), UUID(DEFAULT_DATASET_ID))
|
||||
|
||||
|
||||
def test_rag_pipeline_transform_skips_legacy_acl_when_rbac_is_enabled(sqlite_engine: Engine) -> None:
|
||||
api = module.RagPipelineTransformApi()
|
||||
handler = unwrap_all(api.post)
|
||||
account = _account()
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
expected = {"pipeline_id": "pipeline-1", "dataset_id": DEFAULT_DATASET_ID, "status": "success"}
|
||||
|
||||
with Session(sqlite_engine) as session:
|
||||
session.add(_dataset(maintainer="00000000-0000-0000-0000-000000000099"))
|
||||
|
||||
with (
|
||||
patch.object(module.dify_config, "RBAC_ENABLED", True),
|
||||
patch.object(module.RagPipelineTransformService, "transform_dataset", return_value=expected) as transform,
|
||||
):
|
||||
response = handler(api, session, DEFAULT_WORKFLOW_TENANT_ID, account, UUID(DEFAULT_DATASET_ID))
|
||||
|
||||
assert response == expected
|
||||
transform.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
||||
@@ -40,7 +40,7 @@ from libs.login import AccountWithTenant
|
||||
from machinery.context import RequestContext
|
||||
from models import Account
|
||||
from models.account import AccountStatus, TenantAccountRole
|
||||
from models.dataset import RateLimitLog
|
||||
from models.dataset import Dataset, RateLimitLog
|
||||
from services.entities.feature_entities import LicenseStatus
|
||||
|
||||
|
||||
@@ -396,6 +396,36 @@ class TestRbacPermissionRequired:
|
||||
request.view_args = {"resource_id": "dataset-1"}
|
||||
assert _extract_resource_id(RBACResourceScope.DATASET, "tenant-1") == "dataset-1"
|
||||
|
||||
def test_extract_resource_id_scopes_pipeline_resolution_to_the_calling_tenant(self, sqlite_session: Session):
|
||||
app = Flask(__name__)
|
||||
pipeline_id = "00000000-0000-0000-0000-000000000001"
|
||||
current_tenant_id = "00000000-0000-0000-0000-000000000002"
|
||||
foreign_dataset = Dataset(
|
||||
id="00000000-0000-0000-0000-000000000003",
|
||||
tenant_id="00000000-0000-0000-0000-000000000004",
|
||||
name="Foreign decoy",
|
||||
created_by="00000000-0000-0000-0000-000000000005",
|
||||
pipeline_id=pipeline_id,
|
||||
)
|
||||
current_dataset = Dataset(
|
||||
id="00000000-0000-0000-0000-000000000006",
|
||||
tenant_id=current_tenant_id,
|
||||
name="Current tenant dataset",
|
||||
created_by="00000000-0000-0000-0000-000000000007",
|
||||
pipeline_id=pipeline_id,
|
||||
)
|
||||
sqlite_session.add_all([foreign_dataset, current_dataset])
|
||||
|
||||
unscoped_dataset = sqlite_session.scalar(select(Dataset).where(Dataset.pipeline_id == pipeline_id))
|
||||
assert unscoped_dataset is foreign_dataset
|
||||
|
||||
with (
|
||||
app.test_request_context("/rag/pipelines/pipeline-1"),
|
||||
patch("controllers.common.wraps.db", SimpleNamespace(session=sqlite_session)),
|
||||
):
|
||||
request.view_args = {"pipeline_id": pipeline_id}
|
||||
assert _extract_resource_id(RBACResourceScope.DATASET, current_tenant_id) == current_dataset.id
|
||||
|
||||
def test_extract_resource_id_resolves_agent_to_its_authz_app(self):
|
||||
app = Flask(__name__)
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from flask import Flask, current_app
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.app_config.entities import (
|
||||
DatasetEntity,
|
||||
@@ -4961,6 +4962,7 @@ class TestSingleAndMultipleRetrieveCoverage:
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].provider == "external"
|
||||
session.scalar.assert_called_once()
|
||||
mock_end.assert_called_once()
|
||||
assert retrieval.llm_usage.total_tokens == 2
|
||||
|
||||
@@ -5037,6 +5039,95 @@ class TestSingleAndMultipleRetrieveCoverage:
|
||||
)
|
||||
assert results == []
|
||||
|
||||
def test_single_retrieve_rejects_dataset_outside_available_datasets(self, retrieval: DatasetRetrieval) -> None:
|
||||
available_dataset = _dataset(id="ds-1", name="Available DS", description=None)
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = _dataset(
|
||||
id="ds-2",
|
||||
name="Foreign DS",
|
||||
provider="external",
|
||||
tenant_id="tenant-2",
|
||||
retrieval_model={},
|
||||
)
|
||||
|
||||
with (
|
||||
patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls,
|
||||
patch(
|
||||
"core.rag.retrieval.dataset_retrieval.ExternalDatasetService.fetch_external_knowledge_retrieval",
|
||||
return_value=[],
|
||||
) as mock_external_retrieve,
|
||||
patch.object(retrieval, "_on_query") as mock_on_query,
|
||||
):
|
||||
mock_router_cls.return_value.invoke.return_value = ("ds-2", LLMUsage.empty_usage())
|
||||
results = retrieval.single_retrieve(
|
||||
session,
|
||||
app_id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
user_id="user-1",
|
||||
user_from="workflow",
|
||||
query="python",
|
||||
available_datasets=[available_dataset],
|
||||
model_instance=Mock(),
|
||||
model_config=Mock(),
|
||||
planning_strategy=PlanningStrategy.REACT_ROUTER,
|
||||
)
|
||||
|
||||
assert results == []
|
||||
session.scalar.assert_not_called()
|
||||
mock_external_retrieve.assert_not_called()
|
||||
mock_on_query.assert_not_called()
|
||||
|
||||
def test_single_retrieve_rejects_allowlisted_dataset_owned_by_another_tenant(
|
||||
self, retrieval: DatasetRetrieval, sqlite_session: Session
|
||||
) -> None:
|
||||
dataset_id = str(uuid4())
|
||||
caller_tenant_id = str(uuid4())
|
||||
foreign_dataset = Dataset(
|
||||
id=dataset_id,
|
||||
tenant_id=str(uuid4()),
|
||||
name="Foreign DS",
|
||||
provider="external",
|
||||
indexing_technique="high_quality",
|
||||
retrieval_model={},
|
||||
created_by=str(uuid4()),
|
||||
)
|
||||
sqlite_session.add(foreign_dataset)
|
||||
available_dataset = _dataset(
|
||||
id=dataset_id,
|
||||
tenant_id=caller_tenant_id,
|
||||
name="Available DS",
|
||||
description=None,
|
||||
)
|
||||
|
||||
with (
|
||||
patch("core.rag.retrieval.dataset_retrieval.ReactMultiDatasetRouter") as mock_router_cls,
|
||||
patch(
|
||||
"core.rag.retrieval.dataset_retrieval.ExternalDatasetService.fetch_external_knowledge_retrieval",
|
||||
) as mock_external_retrieve,
|
||||
patch(
|
||||
"core.rag.retrieval.dataset_retrieval.RetrievalService.retrieve",
|
||||
) as mock_internal_retrieve,
|
||||
patch.object(retrieval, "_on_query") as mock_on_query,
|
||||
):
|
||||
mock_router_cls.return_value.invoke.return_value = (dataset_id, LLMUsage.empty_usage())
|
||||
results = retrieval.single_retrieve(
|
||||
sqlite_session,
|
||||
app_id="app-1",
|
||||
tenant_id=caller_tenant_id,
|
||||
user_id="user-1",
|
||||
user_from="workflow",
|
||||
query="python",
|
||||
available_datasets=[available_dataset],
|
||||
model_instance=Mock(),
|
||||
model_config=Mock(),
|
||||
planning_strategy=PlanningStrategy.REACT_ROUTER,
|
||||
)
|
||||
|
||||
assert results == []
|
||||
mock_internal_retrieve.assert_not_called()
|
||||
mock_external_retrieve.assert_not_called()
|
||||
mock_on_query.assert_not_called()
|
||||
|
||||
def test_single_retrieve_respects_metadata_filter_shortcuts(self, retrieval: DatasetRetrieval) -> None:
|
||||
dataset = _dataset(
|
||||
id="ds-1",
|
||||
|
||||
@@ -18,6 +18,7 @@ from urllib.parse import parse_qs, urlparse
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.rag.entities import ParentMode
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
|
||||
@@ -112,6 +113,32 @@ class TestDatasetModelValidation:
|
||||
session.get.assert_called_once_with(Account, dataset.created_by)
|
||||
assert session.scalar.call_count == 2
|
||||
|
||||
def test_get_doc_form_ignores_foreign_tenant_document(self, sqlite_session: Session) -> None:
|
||||
dataset_id = str(uuid4())
|
||||
tenant_id = str(uuid4())
|
||||
created_by = str(uuid4())
|
||||
dataset = Dataset(
|
||||
id=dataset_id,
|
||||
tenant_id=tenant_id,
|
||||
name="Dataset",
|
||||
data_source_type=DataSourceType.UPLOAD_FILE,
|
||||
created_by=created_by,
|
||||
)
|
||||
foreign_document = Document(
|
||||
tenant_id=str(uuid4()),
|
||||
dataset_id=dataset_id,
|
||||
position=1,
|
||||
data_source_type=DataSourceType.UPLOAD_FILE,
|
||||
batch="foreign",
|
||||
name="Foreign",
|
||||
created_from=DocumentCreatedFrom.WEB,
|
||||
created_by=created_by,
|
||||
doc_form=IndexStructureType.PARENT_CHILD_INDEX,
|
||||
)
|
||||
sqlite_session.add_all([dataset, foreign_document])
|
||||
|
||||
assert dataset.get_doc_form(session=sqlite_session) is None
|
||||
|
||||
def test_get_dataset_keyword_table_uses_caller_session(self):
|
||||
dataset = Dataset(
|
||||
tenant_id=str(uuid4()),
|
||||
|
||||
+2
-2
@@ -45,7 +45,7 @@ def test_get_pipeline_template_detail(mocker: MockerFixture, sqlite_session: Ses
|
||||
)
|
||||
retrieval = BuiltInPipelineTemplateRetrieval()
|
||||
|
||||
detail = retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session)
|
||||
detail = retrieval.get_pipeline_template_detail("tpl-1", "tenant-1", session=sqlite_session)
|
||||
|
||||
assert detail == {"id": "tpl-1", "name": "Template 1"}
|
||||
assert not sqlite_session.in_transaction()
|
||||
@@ -79,7 +79,7 @@ def test_get_pipeline_template_detail_returns_none_for_unknown_id(
|
||||
)
|
||||
retrieval = BuiltInPipelineTemplateRetrieval()
|
||||
|
||||
result = retrieval.get_pipeline_template_detail("nonexistent-id", session=sqlite_session)
|
||||
result = retrieval.get_pipeline_template_detail("nonexistent-id", "tenant-1", session=sqlite_session)
|
||||
|
||||
assert result is None
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
+5
-3
@@ -81,7 +81,7 @@ def test_get_pipeline_template_detail_returns_detail(monkeypatch: pytest.MonkeyP
|
||||
monkeypatch.setattr("models.dataset.db", SimpleNamespace(session=sqlite_session))
|
||||
retrieval = CustomizedPipelineTemplateRetrieval()
|
||||
|
||||
detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session)
|
||||
detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, TENANT_ID, session=sqlite_session)
|
||||
|
||||
assert detail == {
|
||||
"id": TEMPLATE_ID,
|
||||
@@ -97,10 +97,12 @@ def test_get_pipeline_template_detail_returns_detail(monkeypatch: pytest.MonkeyP
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(PipelineCustomizedTemplate,)], indirect=True)
|
||||
def test_get_pipeline_template_detail_returns_none_when_not_found(sqlite_session: Session) -> None:
|
||||
def test_get_pipeline_template_detail_rejects_other_tenant(sqlite_session: Session) -> None:
|
||||
sqlite_session.add(_template(tenant_id=OTHER_TENANT_ID))
|
||||
sqlite_session.commit()
|
||||
retrieval = CustomizedPipelineTemplateRetrieval()
|
||||
|
||||
result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session)
|
||||
result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, TENANT_ID, session=sqlite_session)
|
||||
|
||||
assert result is None
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
+2
-2
@@ -68,7 +68,7 @@ def test_get_pipeline_template_detail_returns_detail(sqlite_session: Session) ->
|
||||
sqlite_session.commit()
|
||||
retrieval = DatabasePipelineTemplateRetrieval()
|
||||
|
||||
detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session)
|
||||
detail = retrieval.get_pipeline_template_detail(TEMPLATE_ID, "tenant-1", session=sqlite_session)
|
||||
|
||||
assert detail == {
|
||||
"id": TEMPLATE_ID,
|
||||
@@ -86,7 +86,7 @@ def test_get_pipeline_template_detail_returns_detail(sqlite_session: Session) ->
|
||||
def test_get_pipeline_template_detail_returns_none_when_not_found(sqlite_session: Session) -> None:
|
||||
retrieval = DatabasePipelineTemplateRetrieval()
|
||||
|
||||
result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, session=sqlite_session)
|
||||
result = retrieval.get_pipeline_template_detail(TEMPLATE_ID, "tenant-1", session=sqlite_session)
|
||||
|
||||
assert result is None
|
||||
assert sqlite_session.in_transaction()
|
||||
|
||||
+3
-3
@@ -9,8 +9,8 @@ class DummyRetrieval(PipelineTemplateRetrievalBase):
|
||||
del session
|
||||
return {"language": language}
|
||||
|
||||
def get_pipeline_template_detail(self, template_id: str, *, session) -> dict | None:
|
||||
del session
|
||||
def get_pipeline_template_detail(self, template_id: str, current_tenant_id: str, *, session) -> dict | None:
|
||||
del current_tenant_id, session
|
||||
return {"id": template_id}
|
||||
|
||||
def get_type(self) -> str:
|
||||
@@ -22,6 +22,6 @@ def test_pipeline_template_retrieval_base_concrete_implementation(sqlite_session
|
||||
retrieval = DummyRetrieval()
|
||||
|
||||
assert retrieval.get_pipeline_templates("en-US", session=sqlite_session) == {"language": "en-US"}
|
||||
assert retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session) == {"id": "tpl-1"}
|
||||
assert retrieval.get_pipeline_template_detail("tpl-1", "tenant-1", session=sqlite_session) == {"id": "tpl-1"}
|
||||
assert retrieval.get_type() == "dummy"
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
+1
-1
@@ -46,7 +46,7 @@ def test_get_pipeline_template_detail_fallbacks_to_database_on_error(
|
||||
)
|
||||
retrieval = RemotePipelineTemplateRetrieval()
|
||||
|
||||
result = retrieval.get_pipeline_template_detail("tpl-1", session=sqlite_session)
|
||||
result = retrieval.get_pipeline_template_detail("tpl-1", "tenant-1", session=sqlite_session)
|
||||
|
||||
assert result == {"id": "db-1"}
|
||||
fetch_mock.assert_called_once_with("tpl-1")
|
||||
|
||||
@@ -30,6 +30,7 @@ from models.dataset import (
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
|
||||
from models.workflow import Workflow
|
||||
from services.entities.knowledge_entities.rag_pipeline_entities import IconInfo, PipelineTemplateInfoEntity
|
||||
from services.errors.rag_pipeline import RagPipelineResourceNotFoundError
|
||||
from services.rag_pipeline.rag_pipeline import RagPipelineService
|
||||
from services.workflow_ref_service import WorkflowRef
|
||||
|
||||
@@ -95,6 +96,10 @@ def _make_pipeline(
|
||||
return pipeline
|
||||
|
||||
|
||||
def _make_template_args(name: str = "New Template") -> dict[str, object]:
|
||||
return {"name": name, "description": "Desc", "icon_info": {"icon": "star"}}
|
||||
|
||||
|
||||
def _make_workflow(
|
||||
*,
|
||||
workflow_id: str = "wf-1",
|
||||
@@ -103,13 +108,14 @@ def _make_workflow(
|
||||
graph: dict[str, object] | None = None,
|
||||
features: dict[str, object] | None = None,
|
||||
created_by: str = "u1",
|
||||
version: str = Workflow.VERSION_DRAFT,
|
||||
) -> Workflow:
|
||||
workflow = Workflow(
|
||||
id=workflow_id,
|
||||
tenant_id=tenant_id,
|
||||
app_id=app_id,
|
||||
type="workflow",
|
||||
version="draft",
|
||||
version=version,
|
||||
marked_name="",
|
||||
marked_comment="",
|
||||
graph=json.dumps(graph or {"nodes": []}),
|
||||
@@ -267,12 +273,12 @@ def test_get_pipeline_template_detail_uses_expected_mode(
|
||||
factory_mock = mocker.patch("services.rag_pipeline.rag_pipeline.PipelineTemplateRetrievalFactory")
|
||||
factory_mock.get_pipeline_template_factory.return_value.return_value = retrieval
|
||||
|
||||
result = RagPipelineService.get_pipeline_template_detail("tpl-1", type=template_type, session=session)
|
||||
result = RagPipelineService.get_pipeline_template_detail("tpl-1", "tenant-1", type=template_type, session=session)
|
||||
|
||||
assert result == {"id": "tpl-1"}
|
||||
expected_mode = "remote" if template_type == "built-in" else "customized"
|
||||
factory_mock.get_pipeline_template_factory.assert_called_with(expected_mode)
|
||||
retrieval.get_pipeline_template_detail.assert_called_once_with("tpl-1", session=session)
|
||||
retrieval.get_pipeline_template_detail.assert_called_once_with("tpl-1", "tenant-1", session=session)
|
||||
|
||||
|
||||
def test_get_published_workflow_returns_none_when_pipeline_has_no_workflow_id(
|
||||
@@ -884,8 +890,9 @@ def test_publish_customized_pipeline_template_success(
|
||||
|
||||
account = _make_account(account_id="user-123")
|
||||
|
||||
args = {"name": "New Template", "description": "Desc", "icon_info": {"icon": "star"}, "tags": ["tag1"]}
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template("p1", args, account, "t1", session=session)
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template(
|
||||
pipeline, dataset, _make_template_args(), account, session=session
|
||||
)
|
||||
|
||||
mock_dsl_service.export_rag_pipeline_dsl.assert_called_once_with(pipeline=pipeline, include_secret=True)
|
||||
templates = session.query(PipelineCustomizedTemplate).all()
|
||||
@@ -1847,24 +1854,15 @@ def test_run_datasource_node_preview_raises_for_unsupported_provider(
|
||||
)
|
||||
|
||||
|
||||
def test_publish_customized_pipeline_template_raises_for_missing_pipeline(
|
||||
rag_pipeline_service: RagPipelineServiceTestContext,
|
||||
) -> None:
|
||||
with pytest.raises(ValueError, match="Pipeline not found"):
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template(
|
||||
"p1", {}, _make_account(), "t1", session=rag_pipeline_service.session
|
||||
)
|
||||
|
||||
|
||||
def test_publish_customized_pipeline_template_raises_for_missing_workflow_id(
|
||||
rag_pipeline_service: RagPipelineServiceTestContext,
|
||||
) -> None:
|
||||
pipeline = _make_pipeline(workflow_id=None)
|
||||
_persist(rag_pipeline_service.session, pipeline)
|
||||
|
||||
with pytest.raises(ValueError, match="Pipeline workflow not found"):
|
||||
with pytest.raises(RagPipelineResourceNotFoundError, match="Pipeline workflow not found"):
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template(
|
||||
"p1", {"name": "template-name"}, _make_account(), "t1", session=rag_pipeline_service.session
|
||||
pipeline, _make_dataset(), _make_template_args(), _make_account(), session=rag_pipeline_service.session
|
||||
)
|
||||
|
||||
|
||||
@@ -2153,30 +2151,69 @@ def test_run_free_workflow_node_delegates_to_handle_result(
|
||||
handle.assert_called_once()
|
||||
|
||||
|
||||
def test_publish_customized_pipeline_template_raises_when_workflow_missing(
|
||||
@pytest.mark.parametrize(("workflow_tenant_id", "workflow_app_id"), [("t2", "p1"), ("t1", "p2")])
|
||||
def test_publish_customized_pipeline_template_rejects_unowned_workflow_before_export(
|
||||
mocker: MockerFixture,
|
||||
rag_pipeline_service: RagPipelineServiceTestContext,
|
||||
workflow_tenant_id: str,
|
||||
workflow_app_id: str,
|
||||
) -> None:
|
||||
pipeline = _make_pipeline(workflow_id="wf-1")
|
||||
_persist(rag_pipeline_service.session, pipeline)
|
||||
|
||||
with pytest.raises(ValueError, match="Workflow not found"):
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template(
|
||||
"p1", {}, _make_account(), "t1", session=rag_pipeline_service.session
|
||||
)
|
||||
|
||||
|
||||
def test_publish_customized_pipeline_template_raises_when_dataset_missing(
|
||||
rag_pipeline_service: RagPipelineServiceTestContext,
|
||||
) -> None:
|
||||
pipeline = _make_pipeline(workflow_id="wf-1")
|
||||
workflow = _make_workflow(workflow_id="wf-1")
|
||||
workflow = _make_workflow(workflow_id="wf-1", tenant_id=workflow_tenant_id, app_id=workflow_app_id)
|
||||
_persist(rag_pipeline_service.session, pipeline, workflow)
|
||||
dsl_service = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.RagPipelineDslService")
|
||||
|
||||
with pytest.raises(ValueError, match="Dataset not found"):
|
||||
with pytest.raises(RagPipelineResourceNotFoundError, match="Workflow not found"):
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template(
|
||||
"p1", {}, _make_account(), "t1", session=rag_pipeline_service.session
|
||||
pipeline,
|
||||
_make_dataset(),
|
||||
_make_template_args(),
|
||||
_make_account(),
|
||||
session=rag_pipeline_service.session,
|
||||
)
|
||||
|
||||
dsl_service.assert_not_called()
|
||||
assert rag_pipeline_service.session.query(PipelineCustomizedTemplate).count() == 0
|
||||
|
||||
|
||||
def test_pipeline_retrieve_dataset_rejects_unowned_dataset(
|
||||
rag_pipeline_service: RagPipelineServiceTestContext,
|
||||
) -> None:
|
||||
pipeline = _make_pipeline(workflow_id="wf-1")
|
||||
other_tenant_dataset = _make_dataset(tenant_id="t2")
|
||||
_persist(rag_pipeline_service.session, pipeline, other_tenant_dataset)
|
||||
|
||||
assert pipeline.retrieve_dataset(session=rag_pipeline_service.session) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("draft_tenant_id", "draft_app_id"),
|
||||
[(None, None), ("t2", "p1"), ("t1", "p2")],
|
||||
)
|
||||
def test_publish_customized_pipeline_template_rejects_missing_or_unowned_draft_before_side_effects(
|
||||
mocker: MockerFixture,
|
||||
rag_pipeline_service: RagPipelineServiceTestContext,
|
||||
draft_tenant_id: str | None,
|
||||
draft_app_id: str | None,
|
||||
) -> None:
|
||||
session = rag_pipeline_service.session
|
||||
pipeline = _make_pipeline(workflow_id="wf-published")
|
||||
published_workflow = _make_workflow(workflow_id="wf-published", version="published")
|
||||
dataset = _make_dataset()
|
||||
resources = [pipeline, published_workflow, dataset]
|
||||
if draft_tenant_id and draft_app_id:
|
||||
resources.append(_make_workflow(workflow_id="wf-draft", tenant_id=draft_tenant_id, app_id=draft_app_id))
|
||||
_persist(session, *resources)
|
||||
dsl_service = mocker.patch("services.rag_pipeline.rag_pipeline_dsl_service.RagPipelineDslService")
|
||||
|
||||
with pytest.raises(RagPipelineResourceNotFoundError, match="Draft workflow not found"):
|
||||
rag_pipeline_service.service.publish_customized_pipeline_template(
|
||||
pipeline, dataset, _make_template_args(), _make_account(), session=session
|
||||
)
|
||||
|
||||
dsl_service.assert_not_called()
|
||||
assert session.query(PipelineCustomizedTemplate).count() == 0
|
||||
|
||||
|
||||
def test_get_recommended_plugins_skips_manifest_when_missing(
|
||||
mocker: MockerFixture, rag_pipeline_service: RagPipelineServiceTestContext
|
||||
|
||||
+103
-48
@@ -13,6 +13,7 @@ from models.dataset import Dataset, Document, DocumentPipelineExecutionLog, Pipe
|
||||
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom
|
||||
from models.model import UploadFile
|
||||
from services.entities.knowledge_entities.rag_pipeline_entities import KnowledgeConfiguration
|
||||
from services.errors.rag_pipeline import RagPipelineResourceNotFoundError
|
||||
from services.rag_pipeline.rag_pipeline_transform_service import RagPipelineTransformService
|
||||
|
||||
|
||||
@@ -46,6 +47,24 @@ def _document(**overrides: object) -> Document:
|
||||
return Document(**values)
|
||||
|
||||
|
||||
def _upload_file(*, file_id: str = "file-1", tenant_id: str = "tenant-1") -> UploadFile:
|
||||
upload_file = UploadFile(
|
||||
tenant_id=tenant_id,
|
||||
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_id
|
||||
return upload_file
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("doc_form", "datasource_type", "indexing_technique"),
|
||||
[
|
||||
@@ -128,20 +147,14 @@ def test_deal_dependencies_installs_missing_marketplace_plugins(mocker: MockerFi
|
||||
|
||||
|
||||
@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:
|
||||
def test_transform_to_empty_pipeline_updates_dataset_and_commits(sqlite_session: Session) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_transform_service.current_user",
|
||||
SimpleNamespace(id="user-1"),
|
||||
)
|
||||
|
||||
dataset = _dataset()
|
||||
sqlite_session.add(dataset)
|
||||
sqlite_session.commit()
|
||||
|
||||
result = service._transform_to_empty_pipeline(dataset, session=sqlite_session)
|
||||
result = service._transform_to_empty_pipeline(dataset, account_id="user-1", session=sqlite_session)
|
||||
|
||||
pipeline = sqlite_session.get(Pipeline, result["pipeline_id"])
|
||||
assert pipeline is not None
|
||||
@@ -155,33 +168,54 @@ def test_transform_to_empty_pipeline_updates_dataset_and_commits(
|
||||
# --- transform_dataset ---
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
|
||||
def test_transform_dataset_returns_early_when_pipeline_exists(sqlite_session: Session) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
dataset = _dataset(id="d1", pipeline_id="p1", runtime_mode="rag_pipeline")
|
||||
sqlite_session.add(dataset)
|
||||
sqlite_session.commit()
|
||||
pipeline = Pipeline(tenant_id="tenant-1", name="Pipeline", description="")
|
||||
pipeline.id = "p1"
|
||||
sqlite_session.add_all([dataset, pipeline])
|
||||
|
||||
result = service.transform_dataset("d1", sqlite_session)
|
||||
result = service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
assert result == {"pipeline_id": "p1", "dataset_id": "d1", "status": "success"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
|
||||
def test_transform_dataset_raises_for_dataset_not_found(sqlite_session: Session) -> None:
|
||||
@pytest.mark.parametrize("pipeline_tenant_id", [None, "tenant-2"])
|
||||
def test_transform_dataset_rejects_missing_or_foreign_pipeline_before_side_effects(
|
||||
mocker: MockerFixture,
|
||||
sqlite_session: Session,
|
||||
pipeline_tenant_id: str | None,
|
||||
) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
with pytest.raises(ValueError, match="Dataset not found"):
|
||||
service.transform_dataset("d1", sqlite_session)
|
||||
dataset = _dataset(id="d1", pipeline_id="p1", runtime_mode="rag_pipeline")
|
||||
sqlite_session.add(dataset)
|
||||
if pipeline_tenant_id is not None:
|
||||
pipeline = Pipeline(tenant_id=pipeline_tenant_id, name="Pipeline", description="")
|
||||
pipeline.id = "p1"
|
||||
sqlite_session.add(pipeline)
|
||||
install_plugins = mocker.patch(
|
||||
"services.rag_pipeline.rag_pipeline_transform_service.PluginService.install_from_marketplace_pkg"
|
||||
)
|
||||
create_pipeline = mocker.patch.object(service, "_create_pipeline")
|
||||
transform_empty = mocker.patch.object(service, "_transform_to_empty_pipeline")
|
||||
|
||||
with pytest.raises(RagPipelineResourceNotFoundError, match="Pipeline not found"):
|
||||
service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
install_plugins.assert_not_called()
|
||||
create_pipeline.assert_not_called()
|
||||
transform_empty.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
|
||||
def test_transform_dataset_raises_for_external_dataset(sqlite_session: Session) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
sqlite_session.add(_dataset(id="d1", provider="external"))
|
||||
dataset = _dataset(id="d1", provider="external")
|
||||
sqlite_session.add(dataset)
|
||||
sqlite_session.commit()
|
||||
|
||||
with pytest.raises(ValueError, match="External dataset is not supported"):
|
||||
service.transform_dataset("d1", sqlite_session)
|
||||
service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
|
||||
@@ -189,13 +223,14 @@ def test_transform_dataset_calls_empty_pipeline_when_no_datasource(
|
||||
mocker: MockerFixture, sqlite_session: Session
|
||||
) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
sqlite_session.add(_dataset(id="d1", data_source_type=None, indexing_technique=None))
|
||||
dataset = _dataset(id="d1", data_source_type=None, indexing_technique=None)
|
||||
sqlite_session.add(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", sqlite_session)
|
||||
result = service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
assert result == empty_result
|
||||
|
||||
@@ -205,15 +240,14 @@ def test_transform_dataset_calls_empty_pipeline_when_no_doc_form(
|
||||
mocker: MockerFixture, sqlite_session: Session
|
||||
) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
sqlite_session.add(
|
||||
_dataset(id="d1", data_source_type="upload_file", indexing_technique="high_quality", chunk_structure=None)
|
||||
)
|
||||
dataset = _dataset(id="d1", data_source_type="upload_file", indexing_technique="high_quality", chunk_structure=None)
|
||||
sqlite_session.add(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", sqlite_session)
|
||||
result = service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
assert result == empty_result
|
||||
|
||||
@@ -354,18 +388,19 @@ def test_transform_dataset_full_flow(mocker: MockerFixture, sqlite_session: Sess
|
||||
mocker.patch.object(service, "_deal_dependencies")
|
||||
mocker.patch.object(service, "_deal_document_data")
|
||||
|
||||
# Mock current_user to have the same tenant_id as dataset
|
||||
mock_current_user = SimpleNamespace(current_tenant_id="t1")
|
||||
mocker.patch("services.rag_pipeline.rag_pipeline_transform_service.current_user", mock_current_user)
|
||||
|
||||
pipeline = SimpleNamespace(id="p-new")
|
||||
mocker.patch.object(service, "_create_pipeline", return_value=pipeline)
|
||||
create_pipeline = mocker.patch.object(service, "_create_pipeline", return_value=pipeline)
|
||||
|
||||
result = service.transform_dataset("d1", sqlite_session)
|
||||
result = service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
assert result["pipeline_id"] == "p-new"
|
||||
assert dataset.runtime_mode == "rag_pipeline"
|
||||
assert dataset.chunk_structure == "text_model"
|
||||
assert create_pipeline.call_args.kwargs == {
|
||||
"tenant_id": "t1",
|
||||
"account_id": "user-1",
|
||||
"session": sqlite_session,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
|
||||
@@ -393,7 +428,7 @@ def test_transform_dataset_raises_for_unsupported_doc_form_after_pipeline_create
|
||||
mocker.patch.object(service, "_create_pipeline", return_value=SimpleNamespace(id="p-new"))
|
||||
|
||||
with pytest.raises(ValueError, match="Unsupported doc form"):
|
||||
service.transform_dataset("d1", sqlite_session)
|
||||
service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset,)], indirect=True)
|
||||
@@ -420,7 +455,7 @@ def test_transform_dataset_raises_when_transform_yaml_missing_workflow(
|
||||
mocker.patch.object(service, "_deal_dependencies")
|
||||
|
||||
with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"):
|
||||
service.transform_dataset("d1", sqlite_session)
|
||||
service.transform_dataset(dataset, "user-1", sqlite_session)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [()], indirect=True)
|
||||
@@ -428,7 +463,12 @@ def test_create_pipeline_raises_when_workflow_data_missing(sqlite_session: Sessi
|
||||
service = RagPipelineTransformService()
|
||||
|
||||
with pytest.raises(ValueError, match="Missing workflow data for rag pipeline"):
|
||||
service._create_pipeline({"rag_pipeline": {"name": "N"}}, session=sqlite_session)
|
||||
service._create_pipeline(
|
||||
{"rag_pipeline": {"name": "N"}},
|
||||
tenant_id="tenant-1",
|
||||
account_id="user-1",
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog, UploadFile)], indirect=True)
|
||||
@@ -442,21 +482,7 @@ def test_deal_document_data_upload_file_with_existing_file(sqlite_session: Sessi
|
||||
data_source_info='{"upload_file_id":"file-1"}',
|
||||
name="Doc",
|
||||
)
|
||||
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.add_all([document, _upload_file()])
|
||||
sqlite_session.commit()
|
||||
|
||||
service._deal_document_data(dataset, sqlite_session)
|
||||
@@ -469,6 +495,35 @@ def test_deal_document_data_upload_file_with_existing_file(sqlite_session: Sessi
|
||||
assert log.document_id == document.id
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("document_tenant_id", "upload_file_tenant_id"),
|
||||
[("tenant-2", "tenant-1"), ("tenant-1", "tenant-2")],
|
||||
)
|
||||
@pytest.mark.parametrize("sqlite_session", [(Document, DocumentPipelineExecutionLog, UploadFile)], indirect=True)
|
||||
def test_deal_document_data_scopes_documents_and_upload_files_to_dataset_tenant(
|
||||
sqlite_session: Session,
|
||||
document_tenant_id: str,
|
||||
upload_file_tenant_id: str,
|
||||
) -> None:
|
||||
service = RagPipelineTransformService()
|
||||
dataset = _dataset(id="d1", tenant_id="tenant-1", pipeline_id="p1")
|
||||
document = _document(
|
||||
id="doc-1",
|
||||
tenant_id=document_tenant_id,
|
||||
dataset_id="d1",
|
||||
data_source_type="upload_file",
|
||||
data_source_info='{"upload_file_id":"file-1"}',
|
||||
)
|
||||
sqlite_session.add_all([document, _upload_file(tenant_id=upload_file_tenant_id)])
|
||||
sqlite_session.commit()
|
||||
|
||||
service._deal_document_data(dataset, sqlite_session)
|
||||
sqlite_session.flush()
|
||||
|
||||
assert document.data_source_type == DataSourceType.UPLOAD_FILE
|
||||
assert sqlite_session.scalar(select(DocumentPipelineExecutionLog)) is None
|
||||
|
||||
|
||||
def _make_service():
|
||||
return RagPipelineTransformService.__new__(RagPipelineTransformService)
|
||||
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
"""Unit tests for DatasetService and dataset-related collaborators."""
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from models.dataset import DatasetPermission
|
||||
|
||||
from .dataset_service_test_helpers import (
|
||||
DatasetNameDuplicateError,
|
||||
DatasetPermissionEnum,
|
||||
@@ -45,6 +49,29 @@ class TestDatasetServiceValidation:
|
||||
with pytest.raises(ValueError, match="doc_form is different"):
|
||||
DatasetService.check_doc_form(dataset, "text_model", session=session)
|
||||
|
||||
@pytest.mark.parametrize("operator_check", [False, True])
|
||||
def test_dataset_permission_checks_ignore_foreign_tenant_binding(
|
||||
self, sqlite_session: Session, operator_check: bool
|
||||
) -> None:
|
||||
dataset = DatasetServiceUnitDataFactory.create_dataset_mock(
|
||||
dataset_id="dataset-1",
|
||||
tenant_id="tenant-1",
|
||||
permission=DatasetPermissionEnum.PARTIAL_TEAM,
|
||||
maintainer="owner-1",
|
||||
)
|
||||
user = DatasetServiceUnitDataFactory.create_user_mock(
|
||||
user_id="user-1",
|
||||
tenant_id="tenant-1",
|
||||
role=TenantAccountRole.NORMAL,
|
||||
)
|
||||
sqlite_session.add(DatasetPermission(dataset_id=dataset.id, account_id=user.id, tenant_id="tenant-2"))
|
||||
|
||||
with pytest.raises(NoPermissionError):
|
||||
if operator_check:
|
||||
DatasetService.check_dataset_operator_permission(user, dataset, session=sqlite_session)
|
||||
else:
|
||||
DatasetService.check_dataset_permission(dataset, user, sqlite_session)
|
||||
|
||||
def test_check_dataset_model_setting_skips_non_high_quality_datasets(self):
|
||||
dataset = DatasetServiceUnitDataFactory.create_dataset_mock(indexing_technique="economy")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user