mirror of
https://github.com/langgenius/dify.git
synced 2026-07-30 08:49:31 +08:00
test: move service dataset controller coverage to unit tests (#38941)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Byron Wang <byron@dify.ai>
This commit is contained in:
parent
dbabddf1ab
commit
62791579ab
File diff suppressed because it is too large
Load Diff
@ -0,0 +1,738 @@
|
||||
"""Unit tests for Service API dataset controller behavior.
|
||||
|
||||
Service boundaries stay mocked, while ORM collaborators are concrete model instances
|
||||
persisted in one in-memory SQLite session. The controller's ``db.session`` and the
|
||||
session passed to unwrapped ``@with_session`` endpoints both use that same session,
|
||||
so model properties and service call contracts exercise real SQLAlchemy behavior.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from inspect import unwrap
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy.orm import Session, scoped_session, sessionmaker
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
import services
|
||||
from controllers.service_api.dataset.error import DatasetInUseError, DatasetNameDuplicateError, InvalidActionError
|
||||
from extensions.ext_database import db
|
||||
from models.account import Account, Tenant, TenantAccountRole
|
||||
from models.dataset import AppDatasetJoin, Dataset, DatasetMetadata, Document
|
||||
from models.enums import PermissionEnum
|
||||
from models.model import App, Tag, TagBinding
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
DATASET_MODEL_TABLES = (
|
||||
Account,
|
||||
Tenant,
|
||||
Dataset,
|
||||
Document,
|
||||
App,
|
||||
AppDatasetJoin,
|
||||
DatasetMetadata,
|
||||
Tag,
|
||||
TagBinding,
|
||||
)
|
||||
pytestmark = pytest.mark.parametrize("sqlite_session", [DATASET_MODEL_TABLES], indirect=True)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def controller_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session:
|
||||
"""Route controller and model database access through the test's SQLite session."""
|
||||
|
||||
# Flask-SQLAlchemy exposes a callable registry that also proxies Session methods.
|
||||
# Seed that registry with this fixture's Session so both access styles share one transaction.
|
||||
existing_session_factory = cast(sessionmaker[Session], lambda: sqlite_session)
|
||||
session_registry = scoped_session(existing_session_factory)
|
||||
monkeypatch.setattr(db, "session", session_registry)
|
||||
return sqlite_session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tenant(controller_session: Session) -> Tenant:
|
||||
tenant = Tenant(name="Dataset API Tenant")
|
||||
controller_session.add(tenant)
|
||||
controller_session.flush()
|
||||
return tenant
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def account(controller_session: Session, tenant: Tenant, monkeypatch: pytest.MonkeyPatch) -> Account:
|
||||
account = Account(name="Dataset API User", email=f"dataset-api-{uuid.uuid4()}@example.com")
|
||||
account.role = TenantAccountRole.OWNER
|
||||
account._current_tenant = tenant
|
||||
controller_session.add(account)
|
||||
controller_session.flush()
|
||||
|
||||
# Inject the concrete account at the controller boundary without relying on Flask-Login globals.
|
||||
from controllers.service_api.dataset import dataset as dataset_module
|
||||
|
||||
monkeypatch.setattr(dataset_module, "current_user", account)
|
||||
return account
|
||||
|
||||
|
||||
def make_dataset(
|
||||
session: Session,
|
||||
tenant: Tenant,
|
||||
account: Account,
|
||||
**overrides: object,
|
||||
) -> Dataset:
|
||||
"""Create and flush a real dataset so its database-backed properties can be serialized."""
|
||||
|
||||
base: dict[str, object] = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"tenant_id": tenant.id,
|
||||
"name": "Dataset",
|
||||
"description": "desc",
|
||||
"provider": "vendor",
|
||||
"permission": PermissionEnum.ONLY_ME,
|
||||
"data_source_type": None,
|
||||
"indexing_technique": "economy",
|
||||
"created_by": account.id,
|
||||
"created_at": datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC),
|
||||
"updated_by": None,
|
||||
"updated_at": datetime(2024, 1, 1, 12, 0, 0, tzinfo=UTC),
|
||||
"embedding_model": None,
|
||||
"embedding_model_provider": None,
|
||||
"retrieval_model": None,
|
||||
"summary_index_setting": None,
|
||||
"built_in_field_enabled": False,
|
||||
"pipeline_id": None,
|
||||
"runtime_mode": "general",
|
||||
"chunk_structure": None,
|
||||
"icon_info": None,
|
||||
"enable_api": False,
|
||||
"is_multimodal": False,
|
||||
}
|
||||
base.update(overrides)
|
||||
dataset = Dataset(**base)
|
||||
session.add(dataset)
|
||||
session.flush()
|
||||
return dataset
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def dataset(controller_session: Session, tenant: Tenant, account: Account) -> Dataset:
|
||||
return make_dataset(controller_session, tenant, account)
|
||||
|
||||
|
||||
DATASET_DETAIL_KEYS = {
|
||||
"id",
|
||||
"name",
|
||||
"description",
|
||||
"provider",
|
||||
"permission",
|
||||
"data_source_type",
|
||||
"indexing_technique",
|
||||
"app_count",
|
||||
"document_count",
|
||||
"word_count",
|
||||
"created_by",
|
||||
"author_name",
|
||||
"created_at",
|
||||
"updated_by",
|
||||
"updated_at",
|
||||
"embedding_model",
|
||||
"embedding_model_provider",
|
||||
"embedding_available",
|
||||
"retrieval_model_dict",
|
||||
"summary_index_setting",
|
||||
"tags",
|
||||
"doc_form",
|
||||
"external_knowledge_info",
|
||||
"external_retrieval_model",
|
||||
"doc_metadata",
|
||||
"built_in_field_enabled",
|
||||
"pipeline_id",
|
||||
"runtime_mode",
|
||||
"chunk_structure",
|
||||
"icon_info",
|
||||
"is_published",
|
||||
"total_documents",
|
||||
"total_available_documents",
|
||||
"enable_api",
|
||||
"is_multimodal",
|
||||
"maintainer",
|
||||
}
|
||||
|
||||
|
||||
def assert_dataset_detail_shape(response: dict[str, object], *, with_partial_members: bool = False) -> None:
|
||||
expected_keys = set(DATASET_DETAIL_KEYS)
|
||||
if with_partial_members:
|
||||
expected_keys.add("partial_member_list")
|
||||
assert set(response) == expected_keys
|
||||
assert isinstance(response["created_at"], int)
|
||||
assert isinstance(response["updated_at"], int)
|
||||
retrieval_model = response["retrieval_model_dict"]
|
||||
assert isinstance(retrieval_model, dict)
|
||||
assert set(retrieval_model) == {
|
||||
"search_method",
|
||||
"reranking_enable",
|
||||
"reranking_mode",
|
||||
"reranking_model",
|
||||
"weights",
|
||||
"top_k",
|
||||
"score_threshold_enabled",
|
||||
"score_threshold",
|
||||
}
|
||||
external_retrieval_model = response["external_retrieval_model"]
|
||||
if external_retrieval_model is not None:
|
||||
assert isinstance(external_retrieval_model, dict)
|
||||
assert set(external_retrieval_model) == {
|
||||
"top_k",
|
||||
"score_threshold",
|
||||
"score_threshold_enabled",
|
||||
}
|
||||
if not with_partial_members:
|
||||
assert "partial_member_list" not in response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API endpoint tests — DatasetListApi
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDatasetListApiGet:
|
||||
"""Test suite for DatasetListApi.get() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_list_datasets_success(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_provider_mgr: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetListApi
|
||||
|
||||
mock_dataset_svc.get_datasets.return_value = ([make_dataset(controller_session, tenant, account)], 1)
|
||||
mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]()
|
||||
|
||||
with app.test_request_context("/datasets?page=1&limit=20", method="GET"):
|
||||
api = DatasetListApi()
|
||||
response, status = unwrap(api.get)(api, controller_session, tenant_id=tenant.id)
|
||||
|
||||
assert status == 200
|
||||
assert set(response) == {"data", "has_more", "limit", "total", "page"}
|
||||
assert response["has_more"] is False
|
||||
assert response["limit"] == 20
|
||||
assert response["total"] == 1
|
||||
assert response["page"] == 1
|
||||
assert len(response["data"]) == 1
|
||||
assert_dataset_detail_shape(response["data"][0])
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_list_datasets_preserves_repeated_tag_ids(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_provider_mgr: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetListApi
|
||||
|
||||
mock_dataset_svc.get_datasets.return_value = ([make_dataset(controller_session, tenant, account)], 1)
|
||||
mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]()
|
||||
|
||||
with app.test_request_context("/datasets?tag_ids=tag-a&tag_ids=tag-b", method="GET"):
|
||||
api = DatasetListApi()
|
||||
response, status = unwrap(api.get)(api, controller_session, tenant_id=tenant.id)
|
||||
page, limit, session, tenant_id, user, keyword, tag_ids, include_all = (
|
||||
mock_dataset_svc.get_datasets.call_args.args
|
||||
)
|
||||
assert user is account
|
||||
|
||||
assert status == 200
|
||||
assert response["total"] == 1
|
||||
assert (page, limit, session, tenant_id, keyword, tag_ids, include_all) == (
|
||||
1,
|
||||
20,
|
||||
controller_session,
|
||||
tenant.id,
|
||||
None,
|
||||
["tag-a", "tag-b"],
|
||||
False,
|
||||
)
|
||||
|
||||
|
||||
class TestDatasetListApiPost:
|
||||
"""Test suite for DatasetListApi.post() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_create_dataset_success(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetListApi
|
||||
|
||||
mock_dataset_svc.create_empty_dataset.return_value = make_dataset(
|
||||
controller_session, tenant, account, name="New Dataset"
|
||||
)
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets",
|
||||
method="POST",
|
||||
json={"name": "New Dataset"},
|
||||
):
|
||||
api = DatasetListApi()
|
||||
response, status = unwrap(api.post)(api, controller_session, tenant_id=tenant.id)
|
||||
|
||||
assert status == 200
|
||||
assert_dataset_detail_shape(response)
|
||||
assert response["name"] == "New Dataset"
|
||||
mock_dataset_svc.create_empty_dataset.assert_called_once()
|
||||
|
||||
@pytest.mark.usefixtures("account")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_create_dataset_duplicate_name(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetListApi
|
||||
|
||||
mock_dataset_svc.create_empty_dataset.side_effect = services.errors.dataset.DatasetNameDuplicateError()
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets",
|
||||
method="POST",
|
||||
json={"name": "Existing Dataset"},
|
||||
):
|
||||
api = DatasetListApi()
|
||||
with pytest.raises(DatasetNameDuplicateError):
|
||||
unwrap(api.post)(api, controller_session, tenant_id=tenant.id)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API endpoint tests — DatasetApi
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDatasetApiGet:
|
||||
"""Test suite for DatasetApi.get() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_get_dataset_success(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_provider_mgr: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]()
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="GET",
|
||||
):
|
||||
api = DatasetApi()
|
||||
response, status = unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
assert status == 200
|
||||
assert_dataset_detail_shape(response)
|
||||
assert response["embedding_available"] is True
|
||||
assert response["retrieval_model_dict"]["search_method"] == "keyword_search"
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetPermissionService")
|
||||
@patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_get_dataset_partial_members_shape(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_provider_mgr: MagicMock,
|
||||
mock_perm_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
dataset.permission = PermissionEnum.PARTIAL_TEAM
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_perm_svc.get_dataset_partial_member_list.return_value = ["user-1", "user-2"]
|
||||
mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]()
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="GET",
|
||||
):
|
||||
api = DatasetApi()
|
||||
response, status = unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
assert status == 200
|
||||
assert_dataset_detail_shape(response, with_partial_members=True)
|
||||
assert response["partial_member_list"] == ["user-1", "user-2"]
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.create_plugin_provider_manager")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_get_dataset_uses_default_external_retrieval_model(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_provider_mgr: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
dataset.retrieval_model = None
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_provider_mgr.return_value.get_configurations.return_value.get_models.return_value = list[object]()
|
||||
|
||||
with app.test_request_context(f"/datasets/{dataset.id}", method="GET"):
|
||||
api = DatasetApi()
|
||||
response, status = unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
assert status == 200
|
||||
assert_dataset_detail_shape(response)
|
||||
assert response["external_retrieval_model"] == {
|
||||
"top_k": 2,
|
||||
"score_threshold": 0.0,
|
||||
"score_threshold_enabled": None,
|
||||
}
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_get_dataset_not_found(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="GET",
|
||||
):
|
||||
api = DatasetApi()
|
||||
with pytest.raises(NotFound):
|
||||
unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_get_dataset_no_permission(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.side_effect = services.errors.account.NoPermissionError()
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="GET",
|
||||
):
|
||||
api = DatasetApi()
|
||||
with pytest.raises(Forbidden):
|
||||
unwrap(api.get)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
|
||||
class TestDatasetApiPatch:
|
||||
"""Test suite for DatasetApi.patch() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetPermissionService")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_patch_dataset_success_shape(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_perm_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
dataset.name = "Updated Dataset"
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.update_dataset.return_value = dataset
|
||||
mock_perm_svc.check_permission.return_value = None
|
||||
mock_perm_svc.get_dataset_partial_member_list.return_value = ["user-1"]
|
||||
|
||||
payload = {
|
||||
"name": "Updated Dataset",
|
||||
"permission": "partial_members",
|
||||
"partial_member_list": [{"user_id": "user-1", "role": "editor"}],
|
||||
}
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="PATCH",
|
||||
json=payload,
|
||||
):
|
||||
api = DatasetApi()
|
||||
response, status = unwrap(api.patch)(
|
||||
api,
|
||||
controller_session,
|
||||
_=dataset.tenant_id,
|
||||
dataset_id=dataset.id,
|
||||
)
|
||||
|
||||
assert status == 200
|
||||
assert_dataset_detail_shape(response, with_partial_members=True)
|
||||
assert response["name"] == "Updated Dataset"
|
||||
assert response["partial_member_list"] == ["user-1"]
|
||||
mock_dataset_svc.update_dataset.assert_called_once()
|
||||
_, update_data, _ = mock_dataset_svc.update_dataset.call_args.args
|
||||
session = mock_dataset_svc.update_dataset.call_args.kwargs["session"]
|
||||
assert session is controller_session
|
||||
assert update_data["name"] == "Updated Dataset"
|
||||
assert update_data["permission"] == "partial_members"
|
||||
mock_perm_svc.update_partial_member_list.assert_called_once_with(
|
||||
dataset.tenant_id,
|
||||
dataset.id,
|
||||
[{"user_id": "user-1", "role": "editor"}],
|
||||
controller_session,
|
||||
)
|
||||
|
||||
|
||||
class TestDatasetApiDelete:
|
||||
"""Test suite for DatasetApi.delete() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetPermissionService")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_delete_dataset_success(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_perm_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
mock_dataset_svc.delete_dataset.return_value = True
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="DELETE",
|
||||
):
|
||||
api = DatasetApi()
|
||||
result = unwrap(api.delete)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
assert result == ("", 204)
|
||||
mock_perm_svc.clear_partial_member_list.assert_called_once_with(dataset.id, controller_session)
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_delete_dataset_not_found(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
mock_dataset_svc.delete_dataset.return_value = False
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="DELETE",
|
||||
):
|
||||
api = DatasetApi()
|
||||
with pytest.raises(NotFound):
|
||||
unwrap(api.delete)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_delete_dataset_in_use(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
dataset: Dataset,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetApi
|
||||
|
||||
mock_dataset_svc.delete_dataset.side_effect = services.errors.dataset.DatasetInUseError()
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}",
|
||||
method="DELETE",
|
||||
):
|
||||
api = DatasetApi()
|
||||
with pytest.raises(DatasetInUseError):
|
||||
unwrap(api.delete)(api, controller_session, _=dataset.tenant_id, dataset_id=dataset.id)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API endpoint tests — DocumentStatusApi
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDocumentStatusApiPatch:
|
||||
"""Test suite for DocumentStatusApi.patch() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DocumentService")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_batch_update_status_success(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_doc_svc: MagicMock,
|
||||
app: Flask,
|
||||
tenant: Tenant,
|
||||
dataset: Dataset,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DocumentStatusApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_dataset_svc.check_dataset_model_setting.return_value = None
|
||||
mock_doc_svc.batch_update_document_status.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}/documents/status/enable",
|
||||
method="PATCH",
|
||||
json={"document_ids": ["doc-1", "doc-2"]},
|
||||
):
|
||||
api = DocumentStatusApi()
|
||||
response, status = api.patch(
|
||||
tenant_id=tenant.id,
|
||||
dataset_id=dataset.id,
|
||||
action="enable",
|
||||
)
|
||||
|
||||
assert status == 200
|
||||
assert response["result"] == "success"
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_batch_update_status_dataset_not_found(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
tenant: Tenant,
|
||||
dataset: Dataset,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DocumentStatusApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}/documents/status/enable",
|
||||
method="PATCH",
|
||||
json={"document_ids": ["doc-1"]},
|
||||
):
|
||||
api = DocumentStatusApi()
|
||||
with pytest.raises(NotFound):
|
||||
api.patch(
|
||||
tenant_id=tenant.id,
|
||||
dataset_id=dataset.id,
|
||||
action="enable",
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_batch_update_status_permission_error(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
app: Flask,
|
||||
tenant: Tenant,
|
||||
dataset: Dataset,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DocumentStatusApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.side_effect = services.errors.account.NoPermissionError(
|
||||
"No permission"
|
||||
)
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}/documents/status/enable",
|
||||
method="PATCH",
|
||||
json={"document_ids": ["doc-1"]},
|
||||
):
|
||||
api = DocumentStatusApi()
|
||||
with pytest.raises(Forbidden):
|
||||
api.patch(
|
||||
tenant_id=tenant.id,
|
||||
dataset_id=dataset.id,
|
||||
action="enable",
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DocumentService")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_batch_update_status_indexing_error(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_doc_svc: MagicMock,
|
||||
app: Flask,
|
||||
tenant: Tenant,
|
||||
dataset: Dataset,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DocumentStatusApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_dataset_svc.check_dataset_model_setting.return_value = None
|
||||
mock_doc_svc.batch_update_document_status.side_effect = services.errors.document.DocumentIndexingError()
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}/documents/status/enable",
|
||||
method="PATCH",
|
||||
json={"document_ids": ["doc-1"]},
|
||||
):
|
||||
api = DocumentStatusApi()
|
||||
with pytest.raises(InvalidActionError):
|
||||
api.patch(
|
||||
tenant_id=tenant.id,
|
||||
dataset_id=dataset.id,
|
||||
action="enable",
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.DocumentService")
|
||||
@patch("controllers.service_api.dataset.dataset.DatasetService")
|
||||
def test_batch_update_status_value_error(
|
||||
self,
|
||||
mock_dataset_svc: MagicMock,
|
||||
mock_doc_svc: MagicMock,
|
||||
app: Flask,
|
||||
tenant: Tenant,
|
||||
dataset: Dataset,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DocumentStatusApi
|
||||
|
||||
mock_dataset_svc.get_dataset.return_value = dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_dataset_svc.check_dataset_model_setting.return_value = None
|
||||
mock_doc_svc.batch_update_document_status.side_effect = ValueError("Invalid action")
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{dataset.id}/documents/status/enable",
|
||||
method="PATCH",
|
||||
json={"document_ids": ["doc-1"]},
|
||||
):
|
||||
api = DocumentStatusApi()
|
||||
with pytest.raises(InvalidActionError):
|
||||
api.patch(
|
||||
tenant_id=tenant.id,
|
||||
dataset_id=dataset.id,
|
||||
action="enable",
|
||||
)
|
||||
@ -0,0 +1,207 @@
|
||||
"""Unit tests for Service API dataset request payloads."""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
import pytest
|
||||
|
||||
from controllers.service_api.dataset.dataset import (
|
||||
DatasetCreatePayload,
|
||||
DatasetListQuery,
|
||||
DatasetUpdatePayload,
|
||||
TagBindingPayload,
|
||||
TagCreatePayload,
|
||||
TagDeletePayload,
|
||||
TagUnbindingPayload,
|
||||
TagUpdatePayload,
|
||||
)
|
||||
from models.dataset import DatasetPermissionEnum
|
||||
|
||||
|
||||
class TestDatasetCreatePayload:
|
||||
"""Test suite for DatasetCreatePayload Pydantic model."""
|
||||
|
||||
def test_payload_with_required_name(self) -> None:
|
||||
payload = DatasetCreatePayload(name="Test Dataset")
|
||||
assert payload.name == "Test Dataset"
|
||||
assert payload.description == ""
|
||||
assert payload.permission == DatasetPermissionEnum.ONLY_ME
|
||||
|
||||
def test_payload_with_all_fields(self) -> None:
|
||||
payload = DatasetCreatePayload(
|
||||
name="Full Dataset",
|
||||
description="A comprehensive dataset description",
|
||||
indexing_technique="high_quality",
|
||||
permission=DatasetPermissionEnum.ALL_TEAM,
|
||||
provider="vendor",
|
||||
embedding_model="text-embedding-ada-002",
|
||||
embedding_model_provider="openai",
|
||||
)
|
||||
assert payload.name == "Full Dataset"
|
||||
assert payload.description == "A comprehensive dataset description"
|
||||
assert payload.indexing_technique == "high_quality"
|
||||
assert payload.permission == DatasetPermissionEnum.ALL_TEAM
|
||||
assert payload.provider == "vendor"
|
||||
assert payload.embedding_model == "text-embedding-ada-002"
|
||||
assert payload.embedding_model_provider == "openai"
|
||||
|
||||
def test_payload_name_length_validation_min(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
DatasetCreatePayload(name="")
|
||||
|
||||
def test_payload_name_length_validation_max(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
DatasetCreatePayload(name="A" * 41)
|
||||
|
||||
def test_payload_description_max_length(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
DatasetCreatePayload(name="Dataset", description="A" * 401)
|
||||
|
||||
@pytest.mark.parametrize("technique", ["high_quality", "economy"])
|
||||
def test_payload_valid_indexing_techniques(self, technique: Literal["high_quality", "economy"]) -> None:
|
||||
payload = DatasetCreatePayload(name="Dataset", indexing_technique=technique)
|
||||
assert payload.indexing_technique == technique
|
||||
|
||||
def test_payload_with_external_knowledge_settings(self) -> None:
|
||||
payload = DatasetCreatePayload(
|
||||
name="External Dataset", external_knowledge_api_id="api_123", external_knowledge_id="knowledge_456"
|
||||
)
|
||||
assert payload.external_knowledge_api_id == "api_123"
|
||||
assert payload.external_knowledge_id == "knowledge_456"
|
||||
|
||||
|
||||
class TestDatasetUpdatePayload:
|
||||
"""Test suite for DatasetUpdatePayload Pydantic model."""
|
||||
|
||||
def test_payload_all_optional(self) -> None:
|
||||
payload = DatasetUpdatePayload()
|
||||
assert payload.name is None
|
||||
assert payload.description is None
|
||||
assert payload.permission is None
|
||||
|
||||
def test_payload_with_partial_update(self) -> None:
|
||||
payload = DatasetUpdatePayload(name="Updated Name", description="Updated description")
|
||||
assert payload.name == "Updated Name"
|
||||
assert payload.description == "Updated description"
|
||||
|
||||
def test_payload_with_permission_change(self) -> None:
|
||||
payload = DatasetUpdatePayload(
|
||||
permission=DatasetPermissionEnum.PARTIAL_TEAM,
|
||||
partial_member_list=[{"user_id": "user_123", "role": "editor"}],
|
||||
)
|
||||
assert payload.permission == DatasetPermissionEnum.PARTIAL_TEAM
|
||||
assert payload.partial_member_list is not None
|
||||
assert len(payload.partial_member_list) == 1
|
||||
|
||||
def test_payload_name_length_validation(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
DatasetUpdatePayload(name="")
|
||||
with pytest.raises(ValueError):
|
||||
DatasetUpdatePayload(name="A" * 41)
|
||||
|
||||
|
||||
class TestDatasetListQuery:
|
||||
"""Test suite for DatasetListQuery Pydantic model."""
|
||||
|
||||
def test_query_with_defaults(self) -> None:
|
||||
query = DatasetListQuery()
|
||||
assert query.page == 1
|
||||
assert query.limit == 20
|
||||
assert query.keyword is None
|
||||
assert query.include_all is False
|
||||
assert query.tag_ids == []
|
||||
|
||||
def test_query_with_all_filters(self) -> None:
|
||||
query = DatasetListQuery(
|
||||
page=3, limit=50, keyword="machine learning", include_all=True, tag_ids=["tag1", "tag2", "tag3"]
|
||||
)
|
||||
assert query.page == 3
|
||||
assert query.limit == 50
|
||||
assert query.keyword == "machine learning"
|
||||
assert query.include_all is True
|
||||
assert len(query.tag_ids) == 3
|
||||
|
||||
def test_query_with_tag_filter(self) -> None:
|
||||
query = DatasetListQuery(tag_ids=["tag_abc", "tag_def"])
|
||||
assert query.tag_ids == ["tag_abc", "tag_def"]
|
||||
|
||||
|
||||
class TestTagCreatePayload:
|
||||
"""Test suite for TagCreatePayload Pydantic model."""
|
||||
|
||||
def test_payload_with_name(self) -> None:
|
||||
payload = TagCreatePayload(name="New Tag")
|
||||
assert payload.name == "New Tag"
|
||||
|
||||
def test_payload_name_length_min(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
TagCreatePayload(name="")
|
||||
|
||||
def test_payload_name_length_max(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
TagCreatePayload(name="A" * 51)
|
||||
|
||||
def test_payload_with_unicode_name(self) -> None:
|
||||
payload = TagCreatePayload(name="标签 🏷️ Тег")
|
||||
assert payload.name == "标签 🏷️ Тег"
|
||||
|
||||
|
||||
class TestTagUpdatePayload:
|
||||
"""Test suite for TagUpdatePayload Pydantic model."""
|
||||
|
||||
def test_payload_with_name_and_id(self) -> None:
|
||||
payload = TagUpdatePayload(name="Updated Tag", tag_id="tag_123")
|
||||
assert payload.name == "Updated Tag"
|
||||
assert payload.tag_id == "tag_123"
|
||||
|
||||
def test_payload_requires_tag_id(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
TagUpdatePayload.model_validate({"name": "Updated Tag"})
|
||||
|
||||
|
||||
class TestTagDeletePayload:
|
||||
"""Test suite for TagDeletePayload Pydantic model."""
|
||||
|
||||
def test_payload_with_tag_id(self) -> None:
|
||||
payload = TagDeletePayload(tag_id="tag_to_delete")
|
||||
assert payload.tag_id == "tag_to_delete"
|
||||
|
||||
def test_payload_requires_tag_id(self) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
TagDeletePayload.model_validate({})
|
||||
|
||||
|
||||
class TestTagBindingPayload:
|
||||
"""Test suite for TagBindingPayload Pydantic model."""
|
||||
|
||||
def test_payload_with_valid_data(self) -> None:
|
||||
payload = TagBindingPayload(tag_ids=["tag1", "tag2"], target_id="dataset_123")
|
||||
assert len(payload.tag_ids) == 2
|
||||
assert payload.target_id == "dataset_123"
|
||||
|
||||
def test_payload_rejects_empty_tag_ids(self) -> None:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
TagBindingPayload(tag_ids=[], target_id="dataset_123")
|
||||
assert "Tag IDs is required" in str(exc_info.value)
|
||||
|
||||
def test_payload_single_tag_id(self) -> None:
|
||||
payload = TagBindingPayload(tag_ids=["single_tag"], target_id="dataset_456")
|
||||
assert payload.tag_ids == ["single_tag"]
|
||||
|
||||
|
||||
class TestTagUnbindingPayload:
|
||||
"""Test suite for TagUnbindingPayload Pydantic model."""
|
||||
|
||||
def test_payload_with_valid_data(self) -> None:
|
||||
payload = TagUnbindingPayload(tag_ids=["tag_123"], target_id="dataset_456")
|
||||
assert payload.tag_ids == ["tag_123"]
|
||||
assert payload.target_id == "dataset_456"
|
||||
|
||||
def test_payload_normalizes_legacy_tag_id(self) -> None:
|
||||
payload = TagUnbindingPayload(tag_id="tag_123", target_id="dataset_456")
|
||||
assert payload.tag_ids == ["tag_123"]
|
||||
assert payload.target_id == "dataset_456"
|
||||
|
||||
def test_payload_rejects_empty_tag_ids(self) -> None:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
TagUnbindingPayload(tag_ids=[], target_id="dataset_456")
|
||||
assert "Tag IDs is required" in str(exc_info.value)
|
||||
@ -0,0 +1,380 @@
|
||||
"""Unit tests for Service API dataset tag controller behavior.
|
||||
|
||||
Service boundaries stay mocked, while users, tenants, and tags are real ORM objects
|
||||
persisted in SQLite. Controller database calls share that SQLite session so assertions
|
||||
cover the concrete objects and session passed across the controller boundary.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from inspect import unwrap
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy.orm import Session, scoped_session, sessionmaker
|
||||
from werkzeug.exceptions import Forbidden
|
||||
|
||||
from extensions.ext_database import db
|
||||
from models.account import Account, Tenant, TenantAccountRole
|
||||
from models.enums import TagType
|
||||
from models.model import Tag
|
||||
|
||||
TAG_MODEL_TABLES = (Account, Tenant, Tag)
|
||||
pytestmark = pytest.mark.parametrize("sqlite_session", [TAG_MODEL_TABLES], indirect=True)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def controller_session(sqlite_session: Session, monkeypatch: pytest.MonkeyPatch) -> Session:
|
||||
"""Route controller database access through the test's SQLite session."""
|
||||
|
||||
# Flask-SQLAlchemy exposes a callable registry that also proxies Session methods.
|
||||
# Seed that registry with this fixture's Session so both access styles share one transaction.
|
||||
existing_session_factory = cast(sessionmaker[Session], lambda: sqlite_session)
|
||||
session_registry = scoped_session(existing_session_factory)
|
||||
monkeypatch.setattr(db, "session", session_registry)
|
||||
return sqlite_session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tenant(controller_session: Session) -> Tenant:
|
||||
tenant = Tenant(name="Dataset Tag API Tenant")
|
||||
controller_session.add(tenant)
|
||||
controller_session.flush()
|
||||
return tenant
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def account(controller_session: Session, tenant: Tenant, monkeypatch: pytest.MonkeyPatch) -> Account:
|
||||
account = Account(name="Dataset Tag API User", email=f"dataset-tag-api-{uuid.uuid4()}@example.com")
|
||||
account.role = TenantAccountRole.OWNER
|
||||
account._current_tenant = tenant
|
||||
controller_session.add(account)
|
||||
controller_session.flush()
|
||||
|
||||
# Inject the concrete account at the controller boundary without relying on Flask-Login globals.
|
||||
from controllers.service_api.dataset import dataset as dataset_module
|
||||
|
||||
monkeypatch.setattr(dataset_module, "current_user", account)
|
||||
return account
|
||||
|
||||
|
||||
def make_tag(
|
||||
session: Session,
|
||||
tenant: Tenant,
|
||||
account: Account,
|
||||
*,
|
||||
id: str,
|
||||
name: str,
|
||||
binding_count: int | None = None,
|
||||
) -> Tag:
|
||||
"""Create and flush a real tag, optionally adding the aggregate count returned by TagService."""
|
||||
|
||||
tag = Tag(tenant_id=tenant.id, type=TagType.KNOWLEDGE, name=name, created_by=account.id)
|
||||
tag.id = id
|
||||
session.add(tag)
|
||||
session.flush()
|
||||
if binding_count is not None:
|
||||
tag.__dict__["binding_count"] = binding_count
|
||||
return tag
|
||||
|
||||
|
||||
class TestDatasetTagsApiGet:
|
||||
"""Test suite for DatasetTagsApi.get() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
def test_list_tags_success(
|
||||
self,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsApi
|
||||
|
||||
tag = make_tag(controller_session, tenant, account, id="tag-1", name="Test Tag", binding_count=0)
|
||||
mock_tag_svc.get_tags.return_value = [tag]
|
||||
|
||||
with app.test_request_context("/datasets/tags", method="GET"):
|
||||
api = DatasetTagsApi()
|
||||
response, status = unwrap(api.get)(api, controller_session, _=None)
|
||||
|
||||
assert status == 200
|
||||
assert response == [{"id": "tag-1", "name": "Test Tag", "type": "knowledge", "binding_count": "0"}]
|
||||
mock_tag_svc.get_tags.assert_called_once_with("knowledge", tenant.id, session=controller_session)
|
||||
|
||||
|
||||
class TestDatasetTagsApiPost:
|
||||
"""Test suite for DatasetTagsApi.post() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
def test_create_tag_success(
|
||||
self,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsApi
|
||||
|
||||
tag = make_tag(controller_session, tenant, account, id="tag-new", name="New Tag")
|
||||
mock_tag_svc.save_tags.return_value = tag
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags",
|
||||
method="POST",
|
||||
json={"name": "New Tag"},
|
||||
):
|
||||
api = DatasetTagsApi()
|
||||
response, status = unwrap(api.post)(api, controller_session, _=None)
|
||||
|
||||
assert status == 200
|
||||
assert response == {"id": "tag-new", "name": "New Tag", "type": "knowledge", "binding_count": "0"}
|
||||
mock_tag_svc.save_tags.assert_called_once()
|
||||
|
||||
def test_create_tag_forbidden(self, app: Flask, account: Account) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsApi
|
||||
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags",
|
||||
method="POST",
|
||||
json={"name": "New Tag"},
|
||||
):
|
||||
api = DatasetTagsApi()
|
||||
with pytest.raises(Forbidden):
|
||||
api.post(_=None)
|
||||
|
||||
|
||||
class TestDatasetTagsApiPatch:
|
||||
"""Test suite for DatasetTagsApi.patch() endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
@patch("controllers.service_api.dataset.dataset.service_api_ns")
|
||||
def test_update_tag_success(
|
||||
self,
|
||||
mock_service_api_ns: MagicMock,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsApi
|
||||
|
||||
tag = make_tag(controller_session, tenant, account, id="tag-1", name="Updated Tag")
|
||||
mock_tag_svc.update_tags.return_value = tag
|
||||
mock_tag_svc.get_tag_binding_count.return_value = 5
|
||||
mock_service_api_ns.payload = {"name": "Updated Tag", "tag_id": "tag-1"}
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags",
|
||||
method="PATCH",
|
||||
json={"name": "Updated Tag", "tag_id": "tag-1"},
|
||||
):
|
||||
api = DatasetTagsApi()
|
||||
response, status = unwrap(api.patch)(api, controller_session, _=None)
|
||||
|
||||
assert status == 200
|
||||
assert response == {"id": "tag-1", "name": "Updated Tag", "type": "knowledge", "binding_count": "5"}
|
||||
mock_tag_svc.update_tags.assert_called_once()
|
||||
update_payload, tag_id, session = mock_tag_svc.update_tags.call_args.args
|
||||
assert update_payload.name == "Updated Tag"
|
||||
assert tag_id == "tag-1"
|
||||
assert session is controller_session
|
||||
|
||||
def test_update_tag_forbidden(self, app: Flask, account: Account) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsApi
|
||||
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags",
|
||||
method="PATCH",
|
||||
json={"name": "Updated Tag", "tag_id": "tag-1"},
|
||||
):
|
||||
api = DatasetTagsApi()
|
||||
with pytest.raises(Forbidden):
|
||||
api.patch(_=None)
|
||||
|
||||
|
||||
class TestDatasetTagsApiDelete:
|
||||
"""Test suite for DatasetTagsApi.delete() endpoint."""
|
||||
|
||||
@pytest.mark.usefixtures("account")
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
@patch("controllers.service_api.dataset.dataset.service_api_ns")
|
||||
def test_delete_tag_success(
|
||||
self,
|
||||
mock_service_api_ns: MagicMock,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsApi
|
||||
|
||||
mock_tag_svc.delete_tag.return_value = None
|
||||
mock_service_api_ns.payload = {"tag_id": "tag-1"}
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags",
|
||||
method="DELETE",
|
||||
json={"tag_id": "tag-1"},
|
||||
):
|
||||
api = DatasetTagsApi()
|
||||
result = unwrap(api.delete)(api, controller_session, _=None)
|
||||
|
||||
assert result == ("", 204)
|
||||
mock_tag_svc.delete_tag.assert_called_once_with("tag-1", controller_session, tag_type=TagType.KNOWLEDGE)
|
||||
|
||||
|
||||
class TestDatasetTagsBindingStatusApi:
|
||||
"""Test suite for DatasetTagsBindingStatusApi endpoints."""
|
||||
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
def test_get_dataset_tags_binding_status(
|
||||
self,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
account: Account,
|
||||
tenant: Tenant,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagsBindingStatusApi
|
||||
|
||||
tag = make_tag(controller_session, tenant, account, id="tag_1", name="Test Tag")
|
||||
mock_tag_svc.get_tags_by_target_id.return_value = [tag]
|
||||
|
||||
with app.test_request_context("/", method="GET"):
|
||||
api = DatasetTagsBindingStatusApi()
|
||||
response, status_code = unwrap(api.get)(api, controller_session, tenant.id, dataset_id="dataset_123")
|
||||
|
||||
assert status_code == 200
|
||||
assert response["data"] == [{"id": "tag_1", "name": "Test Tag"}]
|
||||
assert response["total"] == 1
|
||||
mock_tag_svc.get_tags_by_target_id.assert_called_once_with(
|
||||
"knowledge", tenant.id, "dataset_123", controller_session
|
||||
)
|
||||
|
||||
|
||||
class TestDatasetTagBindingApiPost:
|
||||
"""Test suite for DatasetTagBindingApi.post() endpoint."""
|
||||
|
||||
@pytest.mark.usefixtures("account")
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
def test_bind_tags_success(
|
||||
self,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagBindingApi
|
||||
|
||||
mock_tag_svc.save_tag_binding.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags/binding",
|
||||
method="POST",
|
||||
json={"tag_ids": ["tag-1"], "target_id": "ds-1"},
|
||||
):
|
||||
api = DatasetTagBindingApi()
|
||||
result = unwrap(api.post)(api, controller_session, _=None)
|
||||
|
||||
assert result == ("", 204)
|
||||
from services.tag_service import TagBindingCreatePayload
|
||||
|
||||
mock_tag_svc.save_tag_binding.assert_called_once_with(
|
||||
TagBindingCreatePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE),
|
||||
controller_session,
|
||||
)
|
||||
|
||||
def test_bind_tags_forbidden(self, app: Flask, account: Account) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagBindingApi
|
||||
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags/binding",
|
||||
method="POST",
|
||||
json={"tag_ids": ["tag-1"], "target_id": "ds-1"},
|
||||
):
|
||||
api = DatasetTagBindingApi()
|
||||
with pytest.raises(Forbidden):
|
||||
api.post(_=None)
|
||||
|
||||
|
||||
class TestDatasetTagUnbindingApiPost:
|
||||
"""Test suite for DatasetTagUnbindingApi.post() endpoint."""
|
||||
|
||||
@pytest.mark.usefixtures("account")
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
def test_unbind_tag_success(
|
||||
self,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi
|
||||
|
||||
mock_tag_svc.delete_tag_binding.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags/unbinding",
|
||||
method="POST",
|
||||
json={"tag_ids": ["tag-1"], "target_id": "ds-1"},
|
||||
):
|
||||
api = DatasetTagUnbindingApi()
|
||||
result = unwrap(api.post)(api, controller_session, _=None)
|
||||
|
||||
assert result == ("", 204)
|
||||
from services.tag_service import TagBindingDeletePayload
|
||||
|
||||
mock_tag_svc.delete_tag_binding.assert_called_once_with(
|
||||
TagBindingDeletePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE),
|
||||
controller_session,
|
||||
)
|
||||
|
||||
@pytest.mark.usefixtures("account")
|
||||
@patch("controllers.service_api.dataset.dataset.TagService")
|
||||
def test_unbind_legacy_tag_id_success(
|
||||
self,
|
||||
mock_tag_svc: MagicMock,
|
||||
app: Flask,
|
||||
controller_session: Session,
|
||||
) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi
|
||||
|
||||
mock_tag_svc.delete_tag_binding.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags/unbinding",
|
||||
method="POST",
|
||||
json={"tag_id": "tag-1", "target_id": "ds-1"},
|
||||
):
|
||||
api = DatasetTagUnbindingApi()
|
||||
result = unwrap(api.post)(api, controller_session, _=None)
|
||||
|
||||
assert result == ("", 204)
|
||||
from services.tag_service import TagBindingDeletePayload
|
||||
|
||||
mock_tag_svc.delete_tag_binding.assert_called_once_with(
|
||||
TagBindingDeletePayload(tag_ids=["tag-1"], target_id="ds-1", type=TagType.KNOWLEDGE),
|
||||
controller_session,
|
||||
)
|
||||
|
||||
def test_unbind_tag_forbidden(self, app: Flask, account: Account) -> None:
|
||||
from controllers.service_api.dataset.dataset import DatasetTagUnbindingApi
|
||||
|
||||
account.role = TenantAccountRole.NORMAL
|
||||
|
||||
with app.test_request_context(
|
||||
"/datasets/tags/unbinding",
|
||||
method="POST",
|
||||
json={"tag_ids": ["tag-1"], "target_id": "ds-1"},
|
||||
):
|
||||
api = DatasetTagUnbindingApi()
|
||||
with pytest.raises(Forbidden):
|
||||
api.post(_=None)
|
||||
Loading…
Reference in New Issue
Block a user