mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix(api): bind dataset operations to owners (#40149)
This commit is contained in:
@@ -46,6 +46,7 @@ from models.account import Account, TenantAccountRole
|
||||
from models.dataset import Dataset, DatasetQuery, Document
|
||||
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, IndexingStatus
|
||||
from models.model import ApiToken, App, AppMode, IconType, UploadFile
|
||||
from services.dataset_ref_service import DatasetRef
|
||||
from services.dataset_service import DatasetPermissionService, DatasetService
|
||||
from services.enterprise import rbac_service as enterprise_rbac_service
|
||||
|
||||
@@ -811,29 +812,65 @@ class TestDatasetApiDelete:
|
||||
|
||||
|
||||
class TestDatasetUseCheckApi:
|
||||
def test_get_use_check_true(self, app: Flask):
|
||||
@pytest.mark.parametrize("is_using", [True, False])
|
||||
def test_get_use_check(self, app: Flask, is_using: bool):
|
||||
api = DatasetUseCheckApi()
|
||||
method = unwrap(api.get)
|
||||
dataset_id = "dataset-id"
|
||||
dataset = make_dataset(id=dataset_id)
|
||||
current_user = make_account()
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context(f"/datasets/{dataset_id}/use-check"),
|
||||
patch.object(DatasetService, "dataset_use_check", return_value=True),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(DatasetService, "dataset_use_check", return_value=is_using) as dataset_use_check,
|
||||
):
|
||||
result, status = method(api, MagicMock(), dataset_id)
|
||||
result, status = method(api, session, "tenant-1", current_user, dataset_id)
|
||||
assert status == 200
|
||||
assert result == {"is_using": True}
|
||||
assert result == {"is_using": is_using}
|
||||
get_dataset.assert_called_once_with(dataset_id, "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, current_user, session)
|
||||
dataset_use_check.assert_called_once_with(DatasetRef("tenant-1", dataset_id), session)
|
||||
|
||||
def test_get_use_check_false(self, app: Flask):
|
||||
def test_get_use_check_relies_on_rbac_in_rbac_mode(self, app: Flask):
|
||||
api = DatasetUseCheckApi()
|
||||
method = unwrap(api.get)
|
||||
dataset_id = "dataset-id"
|
||||
dataset = make_dataset(id="dataset-id")
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context(f"/datasets/{dataset_id}/use-check"),
|
||||
app.test_request_context("/datasets/dataset-id/use-check"),
|
||||
patch("controllers.console.datasets.datasets.dify_config.RBAC_ENABLED", True),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(DatasetService, "dataset_use_check", return_value=False),
|
||||
):
|
||||
result, status = method(api, MagicMock(), dataset_id)
|
||||
_, status = method(api, session, "tenant-1", make_account(), "dataset-id")
|
||||
|
||||
assert status == 200
|
||||
assert result == {"is_using": False}
|
||||
check_permission.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_cls",
|
||||
[DatasetUseCheckApi, DatasetIndexingStatusApi, DatasetErrorDocs, DatasetAutoDisableLogApi],
|
||||
)
|
||||
def test_dataset_scoped_read_permission_denied(app: Flask, api_cls):
|
||||
api = api_cls()
|
||||
method = unwrap(api.get)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(
|
||||
DatasetService,
|
||||
"check_dataset_permission",
|
||||
side_effect=services.errors.account.NoPermissionError("no permission"),
|
||||
),
|
||||
):
|
||||
with pytest.raises(Forbidden, match="no permission"):
|
||||
method(api, session, "tenant-1", make_account(), "dataset-1")
|
||||
|
||||
|
||||
class TestDatasetQueryApi:
|
||||
@@ -1241,6 +1278,8 @@ class TestDatasetIndexingStatusApi:
|
||||
def test_get_success_with_documents(self, app: Flask):
|
||||
api = DatasetIndexingStatusApi()
|
||||
method = unwrap(api.get)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
current_user = make_account()
|
||||
document = MagicMock()
|
||||
document.id = "doc-1"
|
||||
document.indexing_status = "completed"
|
||||
@@ -1255,28 +1294,43 @@ class TestDatasetIndexingStatusApi:
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [document]
|
||||
session.scalar.return_value = 3
|
||||
with app.test_request_context("/"):
|
||||
response, status = method(api, session, "tenant-1", "dataset-1")
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
):
|
||||
response, status = method(api, session, "tenant-1", current_user, "dataset-1")
|
||||
assert status == 200
|
||||
assert "data" in response
|
||||
assert len(response["data"]) == 1
|
||||
item = response["data"][0]
|
||||
assert item["completed_segments"] == 3
|
||||
assert item["total_segments"] == 3
|
||||
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, current_user, session)
|
||||
assert {"dataset-1", "tenant-1"} <= set(session.scalars.call_args.args[0].compile().params.values())
|
||||
for segment_count_call in session.scalar.call_args_list:
|
||||
assert {"dataset-1", "tenant-1", "doc-1"} <= set(segment_count_call.args[0].compile().params.values())
|
||||
|
||||
def test_get_success_no_documents(self, app: Flask):
|
||||
api = DatasetIndexingStatusApi()
|
||||
method = unwrap(api.get)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
with app.test_request_context("/"):
|
||||
response, status = method(api, session, "tenant-1", "dataset-1")
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
):
|
||||
response, status = method(api, session, "tenant-1", make_account(), "dataset-1")
|
||||
assert status == 200
|
||||
assert response == {"data": []}
|
||||
|
||||
def test_segment_counts_different_values(self, app: Flask):
|
||||
api = DatasetIndexingStatusApi()
|
||||
method = unwrap(api.get)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
document = MagicMock()
|
||||
document.id = "doc-1"
|
||||
document.indexing_status = "indexing"
|
||||
@@ -1291,8 +1345,12 @@ class TestDatasetIndexingStatusApi:
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [document]
|
||||
session.scalar.side_effect = [2, 5]
|
||||
with app.test_request_context("/"):
|
||||
response, status = method(api, session, "tenant-1", "dataset-1")
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
):
|
||||
response, status = method(api, session, "tenant-1", make_account(), "dataset-1")
|
||||
assert status == 200
|
||||
item = response["data"][0]
|
||||
assert item["completed_segments"] == 2
|
||||
@@ -1388,27 +1446,47 @@ class TestDatasetApiDeleteApi:
|
||||
|
||||
|
||||
class TestDatasetEnableApiApi:
|
||||
def test_enable_api(self, app: Flask):
|
||||
@pytest.mark.parametrize(("status_value", "enabled"), [("enable", True), ("disable", False)])
|
||||
def test_update_api_status(self, app: Flask, status_value: str, enabled: bool):
|
||||
api = DatasetEnableApiApi()
|
||||
method = unwrap(api.post)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
current_user = make_account()
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", return_value=None),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(DatasetService, "update_dataset_api_status") as update_status,
|
||||
):
|
||||
response, status = method(api, MagicMock(), "dataset-1", "enable")
|
||||
response, status = method(api, session, "tenant-1", current_user, "dataset-1", status_value)
|
||||
assert status == 200
|
||||
assert response["result"] == "success"
|
||||
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, current_user, session)
|
||||
update_status.assert_called_once_with(dataset, enabled, current_user, session)
|
||||
|
||||
def test_disable_api(self, app: Flask):
|
||||
def test_rejects_non_editor(self, app: Flask):
|
||||
api = DatasetEnableApiApi()
|
||||
method = unwrap(api.post)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets.DatasetService.update_dataset_api_status", return_value=None),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
patch.object(DatasetService, "update_dataset_api_status") as update_status,
|
||||
):
|
||||
response, status = method(api, MagicMock(), "dataset-1", "disable")
|
||||
assert status == 200
|
||||
assert response["result"] == "success"
|
||||
with pytest.raises(Forbidden):
|
||||
method(
|
||||
api,
|
||||
session,
|
||||
"tenant-1",
|
||||
make_account(TenantAccountRole.NORMAL),
|
||||
"dataset-1",
|
||||
"enable",
|
||||
)
|
||||
update_status.assert_not_called()
|
||||
|
||||
|
||||
class TestDatasetApiBaseUrlApi:
|
||||
@@ -1502,27 +1580,35 @@ class TestDatasetErrorDocs:
|
||||
method = unwrap(api.get)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
error_doc = make_document_status(id="error-doc", indexing_status=IndexingStatus.ERROR, error="failed")
|
||||
current_user = make_account()
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch(
|
||||
"controllers.console.datasets.datasets.DocumentService.get_error_documents_by_dataset_id",
|
||||
"controllers.console.datasets.datasets.DocumentService.get_error_documents_by_dataset_ref",
|
||||
return_value=[error_doc],
|
||||
),
|
||||
) as get_error_documents,
|
||||
):
|
||||
response, status = method(api, MagicMock(), "dataset-1")
|
||||
response, status = method(api, session, "tenant-1", current_user, "dataset-1")
|
||||
assert status == 200
|
||||
assert response["total"] == 1
|
||||
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, current_user, session)
|
||||
get_error_documents.assert_called_once_with(DatasetRef("tenant-1", "dataset-1"), session)
|
||||
|
||||
def test_get_dataset_not_found(self, app: Flask):
|
||||
api = DatasetErrorDocs()
|
||||
method = unwrap(api.get)
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset,
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), "dataset-1")
|
||||
method(api, session, "tenant-1", make_account(), "dataset-1")
|
||||
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
|
||||
|
||||
|
||||
class TestDatasetPermissionUserListApi:
|
||||
@@ -1566,23 +1652,29 @@ class TestDatasetAutoDisableLogApi:
|
||||
method = unwrap(api.get)
|
||||
dataset = make_dataset(id="dataset-1")
|
||||
logs = {"document_ids": ["doc-1"], "count": 1}
|
||||
current_user = make_account()
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=dataset),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets.DatasetService.get_dataset_auto_disable_logs", return_value=logs
|
||||
),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(DatasetService, "get_dataset_auto_disable_logs", return_value=logs) as get_logs,
|
||||
):
|
||||
response, status = method(api, MagicMock(), "dataset-1")
|
||||
response, status = method(api, session, "tenant-1", current_user, "dataset-1")
|
||||
assert status == 200
|
||||
assert response == logs
|
||||
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, current_user, session)
|
||||
get_logs.assert_called_once_with(DatasetRef("tenant-1", "dataset-1"), session)
|
||||
|
||||
def test_get_dataset_not_found(self, app: Flask):
|
||||
api = DatasetAutoDisableLogApi()
|
||||
method = unwrap(api.get)
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.datasets.DatasetService.get_dataset", return_value=None),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset,
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), "dataset-1")
|
||||
method(api, session, "tenant-1", make_account(), "dataset-1")
|
||||
get_dataset.assert_called_once_with("dataset-1", "tenant-1", session=session)
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
import datetime
|
||||
from inspect import unwrap
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import select
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
import services
|
||||
@@ -22,13 +23,17 @@ from controllers.console.datasets.datasets_document import (
|
||||
DocumentIndexingStatusApi,
|
||||
DocumentMetadataApi,
|
||||
DocumentMetadataUpdatePayload,
|
||||
DocumentPauseApi,
|
||||
DocumentPipelineExecutionLogApi,
|
||||
DocumentProcessingApi,
|
||||
DocumentRecoverApi,
|
||||
DocumentRenameApi,
|
||||
DocumentResource,
|
||||
DocumentRetryApi,
|
||||
DocumentStatusApi,
|
||||
DocumentSummaryStatusApi,
|
||||
GetProcessRuleApi,
|
||||
WebsiteDocumentSyncApi,
|
||||
)
|
||||
from controllers.console.datasets.error import (
|
||||
DocumentAlreadyFinishedError,
|
||||
@@ -42,6 +47,7 @@ from core.rag.index_processor.constant.index_type import IndexStructureType
|
||||
from models.dataset import Dataset
|
||||
from models.dataset import Document as DatasetDocument
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
|
||||
from services.dataset_ref_service import DatasetRef, DocumentRef
|
||||
from services.vector_space_admission_service import (
|
||||
VECTOR_SPACE_ADMISSION_ERROR_CODE,
|
||||
format_vector_space_admission_error,
|
||||
@@ -191,6 +197,12 @@ def tenant_ctx():
|
||||
return (MagicMock(is_dataset_editor=True, id="u1"), "tenant-1")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def bypass_knowledge_rate_limit():
|
||||
with patch("controllers.console.datasets.datasets_document.check_knowledge_rate_limit") as check:
|
||||
yield check
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patch_tenant(tenant_ctx):
|
||||
return tenant_ctx
|
||||
@@ -222,6 +234,14 @@ def patch_dataset(dataset):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patch_scoped_dataset(dataset):
|
||||
with patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant", return_value=dataset
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patch_permission():
|
||||
with patch(
|
||||
@@ -498,6 +518,57 @@ class TestDatasetInitApi:
|
||||
assert response["batch"] == "batch-init"
|
||||
|
||||
|
||||
class TestDocumentResource:
|
||||
def test_get_document_resolves_owner_chain(self, dataset):
|
||||
api = DocumentResource()
|
||||
session = MagicMock()
|
||||
user = MagicMock()
|
||||
document = MagicMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=dataset,
|
||||
) as get_dataset,
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission"
|
||||
) as check_permission,
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetRefService.get_document_by_ref",
|
||||
return_value=document,
|
||||
) as get_document,
|
||||
):
|
||||
assert api.get_document(session, "ds-1", "doc-1", user, "tenant-1") is document
|
||||
|
||||
get_dataset.assert_called_once_with("ds-1", "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, user, session)
|
||||
get_document.assert_called_once_with(
|
||||
DocumentRef(dataset=DatasetRef(tenant_id="tenant-1", dataset_id="ds-1"), document_id="doc-1"),
|
||||
session=session,
|
||||
)
|
||||
|
||||
def test_get_document_relies_on_rbac_in_rbac_mode(self, dataset):
|
||||
api = DocumentResource()
|
||||
session = MagicMock()
|
||||
with (
|
||||
patch("controllers.console.datasets.datasets_document.dify_config.RBAC_ENABLED", True),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=dataset,
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.check_dataset_permission"
|
||||
) as check_permission,
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetRefService.get_document_by_ref",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
api.get_document(session, "ds-1", "doc-1", MagicMock(), "tenant-1")
|
||||
|
||||
check_permission.assert_not_called()
|
||||
|
||||
|
||||
class TestDocumentApi:
|
||||
def test_get_success(self, app: Flask, patch_tenant):
|
||||
api = DocumentApi()
|
||||
@@ -737,108 +808,219 @@ class TestDocumentStatusApi:
|
||||
|
||||
|
||||
class TestDocumentRetryApi:
|
||||
def test_retry_archived_document_skipped(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_retry_archived_document_skipped(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission):
|
||||
api = DocumentRetryApi()
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
doc = MagicMock(id="doc-1", indexing_status="indexing")
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [doc]
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
|
||||
return_value=[doc],
|
||||
),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=True),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_mock,
|
||||
):
|
||||
resp, status = method(api, MagicMock(), "ds-1")
|
||||
resp, status = method(api, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
retry_mock.assert_called_once_with("ds-1", [], ANY)
|
||||
retry_mock.assert_called_once_with("ds-1", [], session)
|
||||
|
||||
def test_retry_success(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_retry_success(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission):
|
||||
api = DocumentRetryApi()
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
document = MagicMock(id="doc-1", indexing_status=IndexingStatus.INDEXING, archived=False)
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [document]
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
|
||||
return_value=[document],
|
||||
),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
|
||||
) as retry_mock,
|
||||
):
|
||||
response, status = method(api, MagicMock(), "ds-1")
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
retry_mock.assert_called_once_with("ds-1", [document], ANY)
|
||||
retry_mock.assert_called_once_with("ds-1", [document], session)
|
||||
|
||||
def test_retry_loads_selected_documents_in_one_batch(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_retry_loads_selected_documents_in_one_scoped_query(
|
||||
self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission
|
||||
):
|
||||
api = DocumentRetryApi()
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1", "doc-2"]}
|
||||
first_document = MagicMock(id="doc-1", indexing_status=IndexingStatus.ERROR, archived=False)
|
||||
second_document = MagicMock(id="doc-2", indexing_status=IndexingStatus.ERROR, archived=False)
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [first_document, second_document]
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
|
||||
return_value=[first_document, second_document],
|
||||
) as get_documents_by_ids,
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
|
||||
) as retry_mock,
|
||||
):
|
||||
response, status = method(api, session, "ds-1")
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
|
||||
assert status == 204
|
||||
get_documents_by_ids.assert_called_once_with("ds-1", ["doc-1", "doc-2"], session)
|
||||
statement = session.scalars.call_args.args[0]
|
||||
assert statement.compare(
|
||||
select(DatasetDocument).where(
|
||||
DatasetDocument.tenant_id == "tenant-1",
|
||||
DatasetDocument.dataset_id == "ds-1",
|
||||
DatasetDocument.id.in_(["doc-1", "doc-2"]),
|
||||
)
|
||||
)
|
||||
retry_mock.assert_called_once_with("ds-1", [first_document, second_document], session)
|
||||
|
||||
def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission):
|
||||
api = DocumentRetryApi()
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
document = MagicMock(id="doc-1", indexing_status=IndexingStatus.COMPLETED, archived=False)
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = [document]
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
|
||||
return_value=[document],
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
|
||||
) as retry_mock,
|
||||
):
|
||||
response, status = method(api, MagicMock(), "ds-1")
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
retry_mock.assert_called_once_with("ds-1", [], ANY)
|
||||
retry_mock.assert_called_once_with("ds-1", [], session)
|
||||
|
||||
def test_retry_foreign_dataset_has_no_side_effects(self, app: Flask, patch_tenant, bypass_knowledge_rate_limit):
|
||||
api = DocumentRetryApi()
|
||||
method = unwrap(api.post)
|
||||
user, tenant_id = patch_tenant
|
||||
session = MagicMock()
|
||||
payload = {"document_ids": ["doc-1"]}
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", payload),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=None,
|
||||
),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_document,
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, session, tenant_id, user, "foreign-dataset")
|
||||
|
||||
session.scalars.assert_not_called()
|
||||
bypass_knowledge_rate_limit.assert_not_called()
|
||||
retry_document.assert_not_called()
|
||||
|
||||
|
||||
class TestDocumentPauseRecoverApi:
|
||||
@pytest.mark.parametrize(
|
||||
("api_type", "service_method"),
|
||||
[(DocumentPauseApi, "pause_document"), (DocumentRecoverApi, "recover_document")],
|
||||
)
|
||||
def test_patch_uses_scoped_document(
|
||||
self, app: Flask, patch_tenant, bypass_knowledge_rate_limit, api_type, service_method
|
||||
):
|
||||
api = api_type()
|
||||
method = unwrap(api.patch)
|
||||
user, tenant_id = patch_tenant
|
||||
session = MagicMock()
|
||||
document = MagicMock()
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(api, "get_document", return_value=document) as get_document,
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
|
||||
patch(
|
||||
f"controllers.console.datasets.datasets_document.DocumentService.{service_method}"
|
||||
) as process_document,
|
||||
):
|
||||
response, status = method(api, session, tenant_id, user, "ds-1", "doc-1")
|
||||
|
||||
assert (response, status) == ("", 204)
|
||||
get_document.assert_called_once_with(session, "ds-1", "doc-1", user, tenant_id)
|
||||
bypass_knowledge_rate_limit.assert_called_once_with()
|
||||
process_document.assert_called_once_with(document, session)
|
||||
|
||||
|
||||
class TestWebsiteDocumentSyncApi:
|
||||
def test_get_uses_scoped_dataset_and_document(self, app: Flask, patch_tenant, dataset):
|
||||
api = WebsiteDocumentSyncApi()
|
||||
method = unwrap(api.get)
|
||||
user, tenant_id = patch_tenant
|
||||
session = MagicMock()
|
||||
document = MagicMock(data_source_type="website_crawl")
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=dataset,
|
||||
) as get_dataset,
|
||||
patch.object(api, "get_document", return_value=document) as get_document,
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.sync_website_document"
|
||||
) as sync_document,
|
||||
):
|
||||
response, status = method(api, session, tenant_id, user, "ds-1", "doc-1")
|
||||
|
||||
assert status == 200
|
||||
assert response["result"] == "success"
|
||||
get_dataset.assert_called_once_with("ds-1", tenant_id, session=session)
|
||||
get_document.assert_called_once_with(session, dataset.id, "doc-1", user, tenant_id)
|
||||
sync_document.assert_called_once_with(dataset, document, session)
|
||||
|
||||
def test_get_rejects_non_editor_before_loading_document(self, app: Flask, dataset):
|
||||
api = WebsiteDocumentSyncApi()
|
||||
method = unwrap(api.get)
|
||||
user = MagicMock(is_dataset_editor=False)
|
||||
session = MagicMock()
|
||||
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=dataset,
|
||||
),
|
||||
patch.object(api, "get_document") as get_document,
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.sync_website_document"
|
||||
) as sync_document,
|
||||
):
|
||||
with pytest.raises(Forbidden):
|
||||
method(api, session, "tenant-1", user, "ds-1", "doc-1")
|
||||
|
||||
get_document.assert_not_called()
|
||||
sync_document.assert_not_called()
|
||||
|
||||
|
||||
class TestDocumentPipelineExecutionLogApi:
|
||||
def test_get_log_success(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_get_log_success(self, app: Flask, patch_tenant):
|
||||
api = DocumentPipelineExecutionLogApi()
|
||||
method = unwrap(api.get)
|
||||
user, tenant_id = patch_tenant
|
||||
log = MagicMock(datasource_info="{}", datasource_type="file", input_data={}, datasource_node_id="n1")
|
||||
document = MagicMock(id="trusted-doc")
|
||||
session = MagicMock()
|
||||
session.scalar.return_value = log
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=MagicMock()
|
||||
),
|
||||
patch.object(api, "get_document", return_value=document) as get_document,
|
||||
):
|
||||
response, status = method(api, session, "ds-1", "doc-1")
|
||||
response, status = method(api, session, tenant_id, user, "ds-1", "doc-1")
|
||||
assert status == 200
|
||||
get_document.assert_called_once_with(session, "ds-1", "doc-1", user, tenant_id)
|
||||
assert "trusted-doc" in session.scalar.call_args.args[0].compile().params.values()
|
||||
|
||||
|
||||
class TestDocumentGenerateSummaryApi:
|
||||
@@ -996,49 +1178,54 @@ class TestDocumentBatchDownloadZipApi:
|
||||
|
||||
|
||||
class TestDatasetDocumentListApiDelete:
|
||||
def test_delete_success(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_delete_success(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission):
|
||||
"""Test successful deletion of documents"""
|
||||
api = DatasetDocumentListApi()
|
||||
method = unwrap(api.delete)
|
||||
user, tenant_id = patch_tenant
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/?document_id=doc-1&document_id=doc-2"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.check_dataset_model_setting",
|
||||
return_value=None,
|
||||
),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.delete_documents", return_value=None),
|
||||
):
|
||||
response, status = method(api, MagicMock(), "ds-1")
|
||||
response, status = method(api, session, tenant_id, user, "ds-1")
|
||||
assert status == 204
|
||||
|
||||
def test_delete_indexing_error(self, app: Flask, patch_tenant, patch_dataset):
|
||||
def test_delete_indexing_error(self, app: Flask, patch_tenant, patch_scoped_dataset, patch_permission):
|
||||
"""Test deletion with indexing error"""
|
||||
api = DatasetDocumentListApi()
|
||||
method = unwrap(api.delete)
|
||||
user, tenant_id = patch_tenant
|
||||
with (
|
||||
app.test_request_context("/?document_id=doc-1"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.check_dataset_model_setting",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.delete_documents",
|
||||
side_effect=services.errors.document.DocumentIndexingError(),
|
||||
),
|
||||
):
|
||||
with pytest.raises(DocumentIndexingError):
|
||||
method(api, MagicMock(), "ds-1")
|
||||
method(api, MagicMock(), tenant_id, user, "ds-1")
|
||||
|
||||
def test_delete_dataset_not_found(self, app: Flask, patch_tenant):
|
||||
def test_delete_dataset_not_found(self, app: Flask, patch_tenant, bypass_knowledge_rate_limit):
|
||||
"""Test deletion when dataset not found"""
|
||||
api = DatasetDocumentListApi()
|
||||
method = unwrap(api.delete)
|
||||
user, tenant_id = patch_tenant
|
||||
with (
|
||||
app.test_request_context("/?document_id=doc-1"),
|
||||
patch("controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=None),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset_for_tenant",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DocumentService.delete_documents"
|
||||
) as delete_documents,
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), "ds-1")
|
||||
method(api, MagicMock(), tenant_id, user, "foreign-dataset")
|
||||
|
||||
bypass_knowledge_rate_limit.assert_not_called()
|
||||
delete_documents.assert_not_called()
|
||||
|
||||
|
||||
class TestDocumentBatchIndexingEstimateApi:
|
||||
@@ -1339,22 +1526,6 @@ class TestDocumentPermissionCases:
|
||||
assert status == 200
|
||||
assert response == {"tokens": 0, "total_price": 0, "currency": "USD", "total_segments": 0, "preview": []}
|
||||
|
||||
def test_document_tenant_mismatch(self, app: Flask):
|
||||
api = DocumentApi()
|
||||
method = unwrap(api.get)
|
||||
user = MagicMock(is_dataset_editor=True)
|
||||
document = MagicMock(tenant_id="other-tenant", dataset_process_rule=None)
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.DatasetService.get_dataset", return_value=MagicMock()
|
||||
),
|
||||
patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document),
|
||||
patch("controllers.console.datasets.datasets_document.DatasetService.get_process_rules", return_value={}),
|
||||
):
|
||||
with pytest.raises(Forbidden):
|
||||
method(api, MagicMock(), "tenant-1", user, "ds-1", "doc-1")
|
||||
|
||||
def test_process_rule_get_by_document_success(self, app: Flask, patch_tenant):
|
||||
api = GetProcessRuleApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
+14
-6
@@ -18,7 +18,7 @@ from zipfile import ZipFile
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
from werkzeug.exceptions import NotFound
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -108,7 +108,11 @@ def _wire_common_success_mocks(
|
||||
import services.dataset_service as dataset_service_module
|
||||
|
||||
# Return a dataset object and allow permission checks to pass.
|
||||
monkeypatch.setattr(module.DatasetService, "get_dataset", lambda *_args, **_kwargs: SimpleNamespace(id="ds-1"))
|
||||
monkeypatch.setattr(
|
||||
module.DatasetService,
|
||||
"get_dataset_for_tenant",
|
||||
lambda *_args, **_kwargs: SimpleNamespace(id="ds-1", tenant_id="tenant-123"),
|
||||
)
|
||||
monkeypatch.setattr(module.DatasetService, "check_dataset_permission", lambda *_args, **_kwargs: None)
|
||||
|
||||
# Return a document that will be validated inside DocumentResource.get_document.
|
||||
@@ -118,7 +122,11 @@ def _wire_common_success_mocks(
|
||||
data_source_type=data_source_type,
|
||||
upload_file_id=upload_file_id,
|
||||
)
|
||||
monkeypatch.setattr(module.DocumentService, "get_document", lambda *_args, **_kwargs: document)
|
||||
monkeypatch.setattr(
|
||||
module.DatasetRefService,
|
||||
"get_document_by_ref",
|
||||
lambda *_args, **_kwargs: document if document.tenant_id == "tenant-123" else None,
|
||||
)
|
||||
|
||||
# Mock UploadFile lookup via FileService batch helper.
|
||||
upload_files_by_id: dict[str, object] = {}
|
||||
@@ -404,10 +412,10 @@ def test_document_download_rejects_when_upload_file_record_missing(
|
||||
method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1")
|
||||
|
||||
|
||||
def test_document_download_rejects_tenant_mismatch(
|
||||
def test_document_download_rejects_document_owner_mismatch(
|
||||
app: Flask, datasets_document_module, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Ensure tenant mismatch is rejected by the shared `get_document()` permission check."""
|
||||
"""Ensure an owner mismatch is rejected by the shared document resolver."""
|
||||
|
||||
_wire_common_success_mocks(
|
||||
module=datasets_document_module,
|
||||
@@ -422,5 +430,5 @@ def test_document_download_rejects_tenant_mismatch(
|
||||
with app.test_request_context("/datasets/ds-1/documents/doc-1/download", method="GET"):
|
||||
api = datasets_document_module.DocumentDownloadApi()
|
||||
method = unwrap(api.get)
|
||||
with pytest.raises(Forbidden):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), "tenant-123", _mock_user(), dataset_id="ds-1", document_id="doc-1")
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import uuid
|
||||
from inspect import unwrap
|
||||
from unittest.mock import MagicMock, PropertyMock, patch
|
||||
from unittest.mock import ANY, MagicMock, PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from pytest_mock import MockerFixture
|
||||
from werkzeug.exceptions import NotFound
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
from controllers.common.controller_schemas import MetadataUpdatePayload
|
||||
from controllers.console import console_ns
|
||||
@@ -19,6 +19,8 @@ from controllers.console.datasets.metadata import (
|
||||
from models.account import Account
|
||||
from services.dataset_service import DatasetService
|
||||
from services.entities.knowledge_entities.knowledge_entities import MetadataArgs, MetadataOperationData
|
||||
from services.errors.account import NoPermissionError
|
||||
from services.errors.metadata import MetadataResourceNotFoundError
|
||||
from services.metadata_service import MetadataService
|
||||
|
||||
|
||||
@@ -102,12 +104,13 @@ class TestDatasetMetadataCreateApi:
|
||||
|
||||
|
||||
class TestDatasetMetadataGetApi:
|
||||
def test_get_metadata_success(self, app: Flask, dataset, dataset_id):
|
||||
def test_get_metadata_success(self, app: Flask, current_user, dataset, dataset_id):
|
||||
api = DatasetMetadataCreateApi()
|
||||
method = unwrap(api.get)
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset", return_value=dataset),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(
|
||||
MetadataService,
|
||||
"get_dataset_metadatas",
|
||||
@@ -117,17 +120,63 @@ class TestDatasetMetadataGetApi:
|
||||
},
|
||||
),
|
||||
):
|
||||
result, status = method(api, MagicMock(), dataset_id)
|
||||
session = MagicMock()
|
||||
result, status = method(api, session, "tenant-1", current_user, dataset_id)
|
||||
assert status == 200
|
||||
assert result["doc_metadata"] == [{"id": "m1", "name": "author", "type": "string", "count": 0}]
|
||||
assert result["built_in_field_enabled"] is False
|
||||
get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=session)
|
||||
check_permission.assert_called_once_with(dataset, current_user, session)
|
||||
|
||||
def test_get_metadata_dataset_not_found(self, app: Flask, dataset_id):
|
||||
def test_get_metadata_rejects_foreign_tenant_before_read(self, app: Flask, current_user, dataset_id):
|
||||
api = DatasetMetadataCreateApi()
|
||||
method = unwrap(api.get)
|
||||
with app.test_request_context("/"), patch.object(DatasetService, "get_dataset", return_value=None):
|
||||
session = MagicMock()
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=None) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(MetadataService, "get_dataset_metadatas") as get_metadata,
|
||||
):
|
||||
with pytest.raises(NotFound):
|
||||
method(api, MagicMock(), dataset_id)
|
||||
method(api, session, "tenant-1", current_user, dataset_id)
|
||||
|
||||
get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=session)
|
||||
check_permission.assert_not_called()
|
||||
get_metadata.assert_not_called()
|
||||
|
||||
def test_get_metadata_relies_on_rbac_in_rbac_mode(self, app: Flask, current_user, dataset, dataset_id):
|
||||
api = DatasetMetadataCreateApi()
|
||||
method = unwrap(api.get)
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch("controllers.console.datasets.metadata.dify_config.RBAC_ENABLED", True),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission") as check_permission,
|
||||
patch.object(
|
||||
MetadataService,
|
||||
"get_dataset_metadatas",
|
||||
return_value={"doc_metadata": [], "built_in_field_enabled": False},
|
||||
),
|
||||
):
|
||||
_, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id)
|
||||
|
||||
assert status == 200
|
||||
check_permission.assert_not_called()
|
||||
|
||||
def test_get_metadata_rejects_inaccessible_dataset(self, app: Flask, current_user, dataset, dataset_id):
|
||||
api = DatasetMetadataCreateApi()
|
||||
method = unwrap(api.get)
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission", side_effect=NoPermissionError),
|
||||
patch.object(MetadataService, "get_dataset_metadatas") as get_metadata,
|
||||
):
|
||||
with pytest.raises(Forbidden):
|
||||
method(api, MagicMock(), "tenant-1", current_user, dataset_id)
|
||||
|
||||
get_metadata.assert_not_called()
|
||||
|
||||
|
||||
class TestDatasetMetadataApi:
|
||||
@@ -138,13 +187,13 @@ class TestDatasetMetadataApi:
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch.object(DatasetService, "get_dataset", return_value=dataset),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
patch.object(
|
||||
MetadataService,
|
||||
"update_metadata_name",
|
||||
return_value={"id": "m1", "type": "string", "name": "updated-name"},
|
||||
),
|
||||
) as update_metadata,
|
||||
):
|
||||
result, status = method(
|
||||
api,
|
||||
@@ -158,19 +207,23 @@ class TestDatasetMetadataApi:
|
||||
assert status == 200
|
||||
assert result["type"] == "string"
|
||||
assert result["name"] == "updated-name"
|
||||
get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=ANY)
|
||||
update_metadata.assert_called_once_with(dataset, str(metadata_id), "updated-name", current_user, session=ANY)
|
||||
|
||||
def test_delete_metadata_success(self, app: Flask, current_user, dataset, dataset_id, metadata_id):
|
||||
api = DatasetMetadataApi()
|
||||
method = unwrap(api.delete)
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset", return_value=dataset),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset) as get_dataset,
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
patch.object(MetadataService, "delete_metadata"),
|
||||
patch.object(MetadataService, "delete_metadata") as delete_metadata,
|
||||
):
|
||||
result, status = method(api, MagicMock(), current_user, dataset_id, metadata_id)
|
||||
result, status = method(api, MagicMock(), "tenant-1", current_user, dataset_id, metadata_id)
|
||||
assert status == 204
|
||||
assert result == ""
|
||||
get_dataset.assert_called_once_with(str(dataset_id), "tenant-1", session=ANY)
|
||||
delete_metadata.assert_called_once_with(dataset, str(metadata_id), ANY)
|
||||
|
||||
|
||||
class TestDatasetMetadataBuiltInFieldApi:
|
||||
@@ -213,16 +266,38 @@ class TestDocumentMetadataEditApi:
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch.object(DatasetService, "get_dataset", return_value=dataset),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
patch.object(MetadataService, "update_documents_metadata"),
|
||||
):
|
||||
result, status = method(
|
||||
api,
|
||||
MetadataOperationData(operation_data=[{"document_id": "doc-1", "metadata_list": []}]),
|
||||
MetadataOperationData(
|
||||
operation_data=[{"document_id": "00000000-0000-0000-0000-000000000001", "metadata_list": []}]
|
||||
),
|
||||
MagicMock(),
|
||||
dataset.tenant_id,
|
||||
current_user,
|
||||
dataset_id,
|
||||
)
|
||||
assert status == 204
|
||||
assert result == ""
|
||||
|
||||
def test_update_document_metadata_translates_missing_resource(self, app: Flask, current_user, dataset, dataset_id):
|
||||
api = DocumentMetadataEditApi()
|
||||
method = unwrap(api.post)
|
||||
request = MetadataOperationData(operation_data=[])
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch.object(DatasetService, "get_dataset_for_tenant", return_value=dataset),
|
||||
patch.object(DatasetService, "check_dataset_permission"),
|
||||
patch.object(
|
||||
MetadataService,
|
||||
"update_documents_metadata",
|
||||
side_effect=MetadataResourceNotFoundError("Metadata not found."),
|
||||
),
|
||||
pytest.raises(NotFound) as exc_info,
|
||||
):
|
||||
method(api, request, MagicMock(), dataset.tenant_id, current_user, dataset_id)
|
||||
|
||||
assert exc_info.value.description == "Metadata not found."
|
||||
|
||||
@@ -47,6 +47,7 @@ from controllers.service_api.dataset.error import ArchivedDocumentImmutableError
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType
|
||||
from models.dataset import Dataset, Document, DocumentSegment
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom, DocumentDocType, IndexingStatus, SegmentStatus
|
||||
from services.dataset_ref_service import DatasetRef
|
||||
from services.dataset_service import DocumentService
|
||||
from services.entities.knowledge_entities.knowledge_entities import ProcessRule, RetrievalModel
|
||||
from services.errors.file import FileTooLargeError as FileTooLargeServiceError
|
||||
@@ -577,17 +578,22 @@ class TestDocumentServiceBatchMethods:
|
||||
doc_ids = [str(uuid.uuid4()), str(uuid.uuid4())]
|
||||
|
||||
session = sqlite_session
|
||||
session.add_all([make_serializable_document(id=document_id, dataset_id=dataset_id) for document_id in doc_ids])
|
||||
session.add_all(
|
||||
[
|
||||
make_serializable_document(id=document_id, tenant_id="tenant-id", dataset_id=dataset_id)
|
||||
for document_id in doc_ids
|
||||
]
|
||||
)
|
||||
session.flush()
|
||||
|
||||
documents = DocumentService.get_documents_by_ids(dataset_id, doc_ids, session)
|
||||
documents = DocumentService.get_documents_by_ids(DatasetRef("tenant-id", dataset_id), doc_ids, session)
|
||||
|
||||
assert len(documents) == 2
|
||||
assert {document.id for document in documents} == set(doc_ids)
|
||||
|
||||
def test_get_documents_by_ids_empty(self, sqlite_session: Session):
|
||||
"""Test batch retrieval with empty list returns empty."""
|
||||
assert DocumentService.get_documents_by_ids("ds_id", [], sqlite_session) == []
|
||||
assert DocumentService.get_documents_by_ids(DatasetRef("tenant-id", "ds_id"), [], sqlite_session) == []
|
||||
|
||||
|
||||
class TestDocumentServiceFileOperations:
|
||||
|
||||
@@ -17,7 +17,7 @@ Decorator strategy:
|
||||
|
||||
import uuid
|
||||
from inspect import unwrap
|
||||
from unittest.mock import Mock, patch
|
||||
from unittest.mock import ANY, MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
@@ -31,6 +31,7 @@ from controllers.service_api.dataset.metadata import (
|
||||
DatasetMetadataServiceApi,
|
||||
DocumentMetadataEditServiceApi,
|
||||
)
|
||||
from services.errors.metadata import MetadataResourceNotFoundError
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -224,7 +225,7 @@ class TestDatasetMetadataServiceApiPatch(_UsesSQLiteSession):
|
||||
):
|
||||
"""Test successful metadata name update."""
|
||||
metadata_id = str(uuid.uuid4())
|
||||
mock_dataset_svc.get_dataset.return_value = mock_dataset
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_meta_svc.update_metadata_name.return_value = {"id": metadata_id, "type": "string", "name": "New Name"}
|
||||
|
||||
@@ -245,7 +246,12 @@ class TestDatasetMetadataServiceApiPatch(_UsesSQLiteSession):
|
||||
|
||||
assert status == 200
|
||||
assert response == {"id": metadata_id, "type": "string", "name": "New Name"}
|
||||
mock_meta_svc.update_metadata_name.assert_called_once()
|
||||
mock_dataset_svc.get_dataset_for_tenant.assert_called_once_with(
|
||||
str(mock_dataset.id), mock_tenant.id, session=session
|
||||
)
|
||||
mock_meta_svc.update_metadata_name.assert_called_once_with(
|
||||
mock_dataset, metadata_id, "New Name", mock_current_user, session=session
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.metadata.DatasetService")
|
||||
def test_update_metadata_dataset_not_found(
|
||||
@@ -257,7 +263,7 @@ class TestDatasetMetadataServiceApiPatch(_UsesSQLiteSession):
|
||||
):
|
||||
"""Test 404 when dataset not found."""
|
||||
metadata_id = str(uuid.uuid4())
|
||||
mock_dataset_svc.get_dataset.return_value = None
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{mock_dataset.id}/metadata/{metadata_id}",
|
||||
@@ -300,7 +306,7 @@ class TestDatasetMetadataServiceApiDelete(_UsesSQLiteSession):
|
||||
):
|
||||
"""Test successful metadata deletion."""
|
||||
metadata_id = str(uuid.uuid4())
|
||||
mock_dataset_svc.get_dataset.return_value = mock_dataset
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_meta_svc.delete_metadata.return_value = None
|
||||
|
||||
@@ -319,7 +325,10 @@ class TestDatasetMetadataServiceApiDelete(_UsesSQLiteSession):
|
||||
)
|
||||
|
||||
assert response == ("", 204)
|
||||
mock_meta_svc.delete_metadata.assert_called_once()
|
||||
mock_dataset_svc.get_dataset_for_tenant.assert_called_once_with(
|
||||
str(mock_dataset.id), mock_tenant.id, session=session
|
||||
)
|
||||
mock_meta_svc.delete_metadata.assert_called_once_with(mock_dataset, metadata_id, session)
|
||||
|
||||
@patch("controllers.service_api.dataset.metadata.DatasetService")
|
||||
def test_delete_metadata_dataset_not_found(
|
||||
@@ -331,7 +340,7 @@ class TestDatasetMetadataServiceApiDelete(_UsesSQLiteSession):
|
||||
):
|
||||
"""Test 404 when dataset not found."""
|
||||
metadata_id = str(uuid.uuid4())
|
||||
mock_dataset_svc.get_dataset.return_value = None
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{mock_dataset.id}/metadata/{metadata_id}",
|
||||
@@ -521,7 +530,7 @@ class TestDocumentMetadataEditPost(_UsesSQLiteSession):
|
||||
mock_dataset,
|
||||
):
|
||||
"""Test successful documents metadata update."""
|
||||
mock_dataset_svc.get_dataset.return_value = mock_dataset
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset
|
||||
mock_dataset_svc.check_dataset_permission.return_value = None
|
||||
mock_meta_svc.update_documents_metadata.return_value = None
|
||||
|
||||
@@ -541,6 +550,12 @@ class TestDocumentMetadataEditPost(_UsesSQLiteSession):
|
||||
|
||||
assert status == 200
|
||||
assert response["result"] == "success"
|
||||
mock_meta_svc.update_documents_metadata.assert_called_once_with(
|
||||
mock_dataset,
|
||||
ANY,
|
||||
mock_current_user,
|
||||
session=session,
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.metadata.DatasetService")
|
||||
def test_update_documents_metadata_dataset_not_found(
|
||||
@@ -551,7 +566,7 @@ class TestDocumentMetadataEditPost(_UsesSQLiteSession):
|
||||
mock_dataset,
|
||||
):
|
||||
"""Test 404 when dataset not found."""
|
||||
mock_dataset_svc.get_dataset.return_value = None
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = None
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{mock_dataset.id}/documents/metadata",
|
||||
@@ -567,3 +582,34 @@ class TestDocumentMetadataEditPost(_UsesSQLiteSession):
|
||||
tenant_id=mock_tenant.id,
|
||||
dataset_id=mock_dataset.id,
|
||||
)
|
||||
|
||||
@patch("controllers.service_api.dataset.metadata.MetadataService")
|
||||
@patch("controllers.service_api.dataset.metadata.DatasetService")
|
||||
@patch("controllers.service_api.dataset.metadata.current_user")
|
||||
def test_update_documents_metadata_translates_missing_resource(
|
||||
self,
|
||||
mock_current_user,
|
||||
mock_dataset_svc,
|
||||
mock_meta_svc,
|
||||
app: Flask,
|
||||
mock_tenant,
|
||||
mock_dataset,
|
||||
):
|
||||
mock_dataset_svc.get_dataset_for_tenant.return_value = mock_dataset
|
||||
mock_meta_svc.update_documents_metadata.side_effect = MetadataResourceNotFoundError("Document not found.")
|
||||
|
||||
with app.test_request_context(
|
||||
f"/datasets/{mock_dataset.id}/documents/metadata",
|
||||
method="POST",
|
||||
json={"operation_data": []},
|
||||
):
|
||||
api = DocumentMetadataEditServiceApi()
|
||||
with pytest.raises(NotFound) as exc_info:
|
||||
self._call_post(
|
||||
api,
|
||||
MagicMock(),
|
||||
tenant_id=mock_tenant.id,
|
||||
dataset_id=mock_dataset.id,
|
||||
)
|
||||
|
||||
assert exc_info.value.description == "Document not found."
|
||||
|
||||
@@ -332,19 +332,34 @@ class TestDocumentServiceMutations:
|
||||
retry_task.delay.assert_not_called()
|
||||
|
||||
def test_sync_website_document_raises_when_sync_flag_exists(self):
|
||||
document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1")
|
||||
dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1")
|
||||
document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", dataset_id=dataset.id)
|
||||
session = MagicMock()
|
||||
|
||||
with patch("services.dataset_service.redis_client") as mock_redis:
|
||||
mock_redis.get.return_value = "1"
|
||||
|
||||
with pytest.raises(ValueError, match="being synced"):
|
||||
DocumentService.sync_website_document("dataset-1", document, session)
|
||||
DocumentService.sync_website_document(dataset, document, session)
|
||||
|
||||
def test_sync_website_document_rejects_document_outside_dataset(self):
|
||||
dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1")
|
||||
document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", dataset_id="dataset-2")
|
||||
|
||||
with (
|
||||
pytest.raises(ValueError, match="Document not found"),
|
||||
patch("services.dataset_service.redis_client") as mock_redis,
|
||||
):
|
||||
DocumentService.sync_website_document(dataset, document, MagicMock())
|
||||
|
||||
mock_redis.get.assert_not_called()
|
||||
|
||||
def test_sync_website_document_updates_status_sets_cache_and_dispatches_task(self):
|
||||
session = MagicMock()
|
||||
dataset = DatasetServiceUnitDataFactory.create_dataset_mock(dataset_id="dataset-1")
|
||||
document = DatasetServiceUnitDataFactory.create_document_mock(
|
||||
document_id="doc-1",
|
||||
dataset_id=dataset.id,
|
||||
data_source_info_dict={"mode": "crawl"},
|
||||
)
|
||||
|
||||
@@ -354,14 +369,14 @@ class TestDocumentServiceMutations:
|
||||
):
|
||||
mock_redis.get.return_value = None
|
||||
|
||||
DocumentService.sync_website_document("dataset-1", document, session)
|
||||
DocumentService.sync_website_document(dataset, document, session)
|
||||
|
||||
assert document.indexing_status == "waiting"
|
||||
assert '"mode": "scrape"' in document.data_source_info
|
||||
session.add.assert_called_once_with(document)
|
||||
session.commit.assert_called_once()
|
||||
mock_redis.setex.assert_called_once_with("document_doc-1_is_sync", 600, 1)
|
||||
sync_task.delay.assert_called_once_with("dataset-1", "doc-1")
|
||||
sync_task.delay.assert_called_once_with(dataset.id, document.id)
|
||||
|
||||
|
||||
class TestDocumentServiceSaveDocumentWithoutDatasetId:
|
||||
|
||||
@@ -58,9 +58,7 @@ class TestMetadataBugCompleteValidation:
|
||||
account = _make_account()
|
||||
none_name = cast(str, None)
|
||||
with pytest.raises(TypeError, match="object of type 'NoneType' has no len"):
|
||||
MetadataService.update_metadata_name(
|
||||
"dataset-123", "metadata-456", none_name, account, "tenant-123", session=sqlite_session
|
||||
)
|
||||
MetadataService.update_metadata_name(Mock(), "metadata-456", none_name, account, session=sqlite_session)
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
def test_3_database_constraints_verification(self) -> None:
|
||||
|
||||
@@ -54,9 +54,7 @@ class TestMetadataNullableBug:
|
||||
none_name = cast(str, None)
|
||||
# This should crash with TypeError when calling len(None)
|
||||
with pytest.raises(TypeError, match="object of type 'NoneType' has no len"):
|
||||
MetadataService.update_metadata_name(
|
||||
"dataset-123", "metadata-456", none_name, account, "tenant-123", session=sqlite_session
|
||||
)
|
||||
MetadataService.update_metadata_name(Mock(), "metadata-456", none_name, account, session=sqlite_session)
|
||||
assert not sqlite_session.in_transaction()
|
||||
|
||||
def test_api_layer_now_uses_pydantic_validation(self) -> None:
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import patch
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -12,10 +14,17 @@ from services.dataset_service import DocumentService
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
MetadataArgs,
|
||||
MetadataDetail,
|
||||
MetadataOperationData,
|
||||
)
|
||||
from services.errors.metadata import MetadataResourceNotFoundError
|
||||
from services.metadata_service import MetadataService
|
||||
|
||||
DOCUMENT_ID = "11111111-1111-1111-1111-111111111111"
|
||||
FOREIGN_DOCUMENT_ID = "22222222-2222-2222-2222-222222222222"
|
||||
METADATA_ID = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"
|
||||
FOREIGN_METADATA_ID = "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb"
|
||||
|
||||
|
||||
def _account() -> Account:
|
||||
account = Account(name="User", email="user@example.com")
|
||||
@@ -55,7 +64,7 @@ def _dataset(*, built_in_field_enabled: bool) -> Dataset:
|
||||
|
||||
def _document() -> Document:
|
||||
return Document(
|
||||
id="document-1",
|
||||
id=DOCUMENT_ID,
|
||||
tenant_id="tenant-1",
|
||||
dataset_id="dataset-1",
|
||||
position=1,
|
||||
@@ -99,20 +108,129 @@ def test_update_documents_metadata_uses_caller_session_for_uploader(sqlite_sessi
|
||||
|
||||
with (
|
||||
patch.object(MetadataService, "knowledge_base_metadata_lock_check"),
|
||||
patch.object(DocumentService, "get_document", return_value=document),
|
||||
patch("services.metadata_service.redis_client.delete"),
|
||||
):
|
||||
MetadataService.update_documents_metadata(
|
||||
dataset,
|
||||
metadata_args,
|
||||
_account(),
|
||||
"tenant-1",
|
||||
session=sqlite_session,
|
||||
)
|
||||
|
||||
assert document.doc_metadata[BuiltInField.uploader] == "User"
|
||||
|
||||
|
||||
def test_update_documents_metadata_rejects_foreign_metadata_before_writes() -> None:
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
dataset = MagicMock(id="dataset-1", tenant_id="tenant-1")
|
||||
metadata_args = MetadataOperationData(
|
||||
operation_data=[
|
||||
DocumentMetadataOperation(
|
||||
document_id=DOCUMENT_ID,
|
||||
metadata_list=[MetadataDetail(id=FOREIGN_METADATA_ID, name="spoofed", value="value")],
|
||||
partial_update=False,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
pytest.raises(MetadataResourceNotFoundError, match="Metadata not found"),
|
||||
patch.object(MetadataService, "knowledge_base_metadata_lock_check") as lock_check,
|
||||
):
|
||||
MetadataService.update_documents_metadata(dataset, metadata_args, _account(), session=session)
|
||||
|
||||
lock_check.assert_not_called()
|
||||
session.add.assert_not_called()
|
||||
session.execute.assert_not_called()
|
||||
session.commit.assert_not_called()
|
||||
|
||||
|
||||
def test_update_documents_metadata_validates_all_documents_before_writes() -> None:
|
||||
session = MagicMock()
|
||||
metadata = SimpleNamespace(id=METADATA_ID, name="canonical")
|
||||
session.scalars.return_value.all.side_effect = [[metadata], [DOCUMENT_ID]]
|
||||
dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=False)
|
||||
metadata_detail = MetadataDetail(id=metadata.id, name="spoofed", value="value")
|
||||
metadata_args = MetadataOperationData(
|
||||
operation_data=[
|
||||
DocumentMetadataOperation(document_id=DOCUMENT_ID, metadata_list=[metadata_detail], partial_update=False),
|
||||
DocumentMetadataOperation(
|
||||
document_id=FOREIGN_DOCUMENT_ID, metadata_list=[metadata_detail], partial_update=False
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
pytest.raises(MetadataResourceNotFoundError, match="Document not found"),
|
||||
patch.object(MetadataService, "knowledge_base_metadata_lock_check") as lock_check,
|
||||
patch("services.metadata_service.redis_client.delete"),
|
||||
):
|
||||
MetadataService.update_documents_metadata(dataset, metadata_args, _account(), session=session)
|
||||
|
||||
lock_check.assert_not_called()
|
||||
session.add.assert_not_called()
|
||||
session.execute.assert_not_called()
|
||||
session.commit.assert_not_called()
|
||||
|
||||
|
||||
def test_update_documents_metadata_uses_canonical_metadata_name() -> None:
|
||||
session = MagicMock()
|
||||
metadata = SimpleNamespace(id=METADATA_ID, name="canonical")
|
||||
session.scalars.return_value.all.side_effect = [[metadata], [DOCUMENT_ID]]
|
||||
dataset = MagicMock(id="dataset-1", tenant_id="tenant-1", built_in_field_enabled=False)
|
||||
document = _document()
|
||||
session.scalar.return_value = document
|
||||
metadata_args = MetadataOperationData(
|
||||
operation_data=[
|
||||
DocumentMetadataOperation(
|
||||
document_id=document.id,
|
||||
metadata_list=[MetadataDetail(id=metadata.id, name="spoofed", value="value")],
|
||||
partial_update=False,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(MetadataService, "knowledge_base_metadata_lock_check"),
|
||||
patch("services.metadata_service.redis_client.delete"),
|
||||
):
|
||||
MetadataService.update_documents_metadata(dataset, metadata_args, _account(), session=session)
|
||||
|
||||
assert document.doc_metadata == {"canonical": "value"}
|
||||
|
||||
|
||||
def test_metadata_operation_normalizes_uuid_ids() -> None:
|
||||
operation = DocumentMetadataOperation(
|
||||
document_id=DOCUMENT_ID.upper(),
|
||||
metadata_list=[MetadataDetail(id=METADATA_ID.upper(), name="ignored", value="value")],
|
||||
)
|
||||
|
||||
assert operation.document_id == DOCUMENT_ID
|
||||
assert operation.metadata_list[0].id == METADATA_ID
|
||||
|
||||
|
||||
def test_document_metadata_details_scopes_binding_to_document_owner() -> None:
|
||||
session = MagicMock()
|
||||
session.scalars.return_value.all.return_value = []
|
||||
document = MagicMock(
|
||||
id="document-1",
|
||||
tenant_id="tenant-1",
|
||||
dataset_id="dataset-1",
|
||||
doc_metadata={"canonical": "value"},
|
||||
)
|
||||
document.get_built_in_fields.return_value = []
|
||||
|
||||
assert Document.get_doc_metadata_details(document, session=session) == []
|
||||
|
||||
statement = str(session.scalars.call_args.args[0])
|
||||
assert "dataset_metadatas.tenant_id" in statement
|
||||
assert "dataset_metadatas.dataset_id" in statement
|
||||
assert "dataset_metadata_bindings.tenant_id" in statement
|
||||
assert "dataset_metadata_bindings.dataset_id" in statement
|
||||
assert "dataset_metadata_bindings.document_id" in statement
|
||||
|
||||
|
||||
def test_get_dataset_metadatas_uses_caller_session(monkeypatch, sqlite_session: Session) -> None:
|
||||
dataset = _dataset(built_in_field_enabled=False)
|
||||
sqlite_session.add_all(
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
import uuid
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType
|
||||
from models.dataset import Dataset, Document, DocumentSegment
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom
|
||||
from tasks.sync_website_document_indexing_task import sync_website_document_indexing_task
|
||||
|
||||
|
||||
def _dataset(tenant_id: str) -> Dataset:
|
||||
return Dataset(
|
||||
id=str(uuid.uuid4()),
|
||||
tenant_id=tenant_id,
|
||||
name="Website dataset",
|
||||
data_source_type=DataSourceType.WEBSITE_CRAWL,
|
||||
created_by=str(uuid.uuid4()),
|
||||
)
|
||||
|
||||
|
||||
def _document(dataset: Dataset) -> Document:
|
||||
return Document(
|
||||
id=str(uuid.uuid4()),
|
||||
tenant_id=dataset.tenant_id,
|
||||
dataset_id=dataset.id,
|
||||
position=1,
|
||||
data_source_type=DataSourceType.WEBSITE_CRAWL,
|
||||
batch="batch-1",
|
||||
name="Website document",
|
||||
created_from=DocumentCreatedFrom.WEB,
|
||||
created_by=str(uuid.uuid4()),
|
||||
doc_form=IndexStructureType.PARAGRAPH_INDEX,
|
||||
)
|
||||
|
||||
|
||||
def _segment(*, tenant_id: str, dataset_id: str, document_id: str) -> DocumentSegment:
|
||||
return DocumentSegment(
|
||||
tenant_id=tenant_id,
|
||||
dataset_id=dataset_id,
|
||||
document_id=document_id,
|
||||
position=1,
|
||||
content="content",
|
||||
word_count=1,
|
||||
tokens=1,
|
||||
created_by=str(uuid.uuid4()),
|
||||
)
|
||||
|
||||
|
||||
def test_rejects_document_outside_dataset_before_side_effects(sqlite_session: Session) -> None:
|
||||
tenant_id = str(uuid.uuid4())
|
||||
requested_dataset = _dataset(tenant_id)
|
||||
foreign_dataset = _dataset(tenant_id)
|
||||
foreign_document = _document(foreign_dataset)
|
||||
sqlite_session.add_all([requested_dataset, foreign_dataset, foreign_document])
|
||||
sqlite_session.commit()
|
||||
|
||||
with (
|
||||
patch("tasks.sync_website_document_indexing_task.FeatureService") as feature_service,
|
||||
patch("tasks.sync_website_document_indexing_task.IndexProcessorFactory") as processor_factory,
|
||||
):
|
||||
sync_website_document_indexing_task(requested_dataset.id, foreign_document.id)
|
||||
|
||||
feature_service.get_features.assert_not_called()
|
||||
processor_factory.assert_not_called()
|
||||
|
||||
|
||||
def test_cleanup_is_owner_scoped_and_skips_empty_vector_ids(sqlite_session: Session) -> None:
|
||||
tenant_id = str(uuid.uuid4())
|
||||
dataset = _dataset(tenant_id)
|
||||
document = _document(dataset)
|
||||
owned_segment = _segment(tenant_id=tenant_id, dataset_id=dataset.id, document_id=document.id)
|
||||
other_dataset = _dataset(tenant_id)
|
||||
decoy_segments = [
|
||||
_segment(tenant_id=tenant_id, dataset_id=dataset.id, document_id=str(uuid.uuid4())),
|
||||
_segment(tenant_id=tenant_id, dataset_id=other_dataset.id, document_id=document.id),
|
||||
_segment(tenant_id=str(uuid.uuid4()), dataset_id=dataset.id, document_id=document.id),
|
||||
]
|
||||
for index, segment in enumerate(decoy_segments):
|
||||
segment.index_node_id = f"decoy-node-{index}"
|
||||
sqlite_session.add_all([dataset, document, owned_segment, other_dataset, *decoy_segments])
|
||||
sqlite_session.commit()
|
||||
|
||||
features = MagicMock()
|
||||
features.billing.enabled = False
|
||||
with (
|
||||
patch("tasks.sync_website_document_indexing_task.FeatureService.get_features", return_value=features),
|
||||
patch("tasks.sync_website_document_indexing_task.IndexProcessorFactory") as processor_factory,
|
||||
patch("tasks.sync_website_document_indexing_task.IndexingRunner") as indexing_runner,
|
||||
patch("tasks.sync_website_document_indexing_task.redis_client"),
|
||||
):
|
||||
sync_website_document_indexing_task(dataset.id, document.id)
|
||||
|
||||
processor_factory.return_value.init_index_processor.return_value.clean.assert_not_called()
|
||||
indexing_runner.return_value.run.assert_called_once()
|
||||
sqlite_session.expire_all()
|
||||
assert set(sqlite_session.scalars(select(DocumentSegment.id))) == {segment.id for segment in decoy_segments}
|
||||
Reference in New Issue
Block a user