mirror of
https://github.com/langgenius/dify.git
synced 2026-09-21 13:20:52 +08:00
fix(api): bind dataset operations to owners (#40149)
This commit is contained in:
+13
-38
@@ -11,13 +11,13 @@ from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import NotFound
|
||||
|
||||
from core.rag.index_processor.constant.index_type import IndexTechniqueType
|
||||
from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus
|
||||
from models.dataset import AppDatasetJoin, Dataset, DatasetPermissionEnum
|
||||
from models.enums import DataSourceType
|
||||
from models.model import App
|
||||
from services.dataset_ref_service import DatasetRefService
|
||||
from services.dataset_service import DatasetService
|
||||
from services.errors.account import NoPermissionError
|
||||
|
||||
@@ -228,9 +228,10 @@ class TestDatasetServiceDatasetUseCheck:
|
||||
dataset = DatasetUpdateDeleteTestDataFactory.create_dataset(db_session_with_containers, tenant.id, owner.id)
|
||||
app = DatasetUpdateDeleteTestDataFactory.create_app(db_session_with_containers, tenant.id, owner.id)
|
||||
DatasetUpdateDeleteTestDataFactory.create_app_dataset_join(db_session_with_containers, app.id, dataset.id)
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
|
||||
# Act
|
||||
result = DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers)
|
||||
result = DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers)
|
||||
|
||||
# Assert
|
||||
assert result is True
|
||||
@@ -252,9 +253,10 @@ class TestDatasetServiceDatasetUseCheck:
|
||||
db_session_with_containers, role=TenantAccountRole.OWNER
|
||||
)
|
||||
dataset = DatasetUpdateDeleteTestDataFactory.create_dataset(db_session_with_containers, tenant.id, owner.id)
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
|
||||
# Act
|
||||
result = DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers)
|
||||
result = DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers)
|
||||
|
||||
# Assert
|
||||
assert result is False
|
||||
@@ -288,11 +290,8 @@ class TestDatasetServiceUpdateDatasetApiStatus:
|
||||
current_time = datetime.datetime(2023, 1, 1, 12, 0, 0)
|
||||
|
||||
# Act
|
||||
with (
|
||||
patch("services.dataset_service.current_user", owner),
|
||||
patch("services.dataset_service.naive_utc_now", return_value=current_time),
|
||||
):
|
||||
DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers)
|
||||
with patch("services.dataset_service.naive_utc_now", return_value=current_time):
|
||||
DatasetService.update_dataset_api_status(dataset, True, owner, session=db_session_with_containers)
|
||||
|
||||
# Assert
|
||||
db_session_with_containers.refresh(dataset)
|
||||
@@ -323,36 +322,14 @@ class TestDatasetServiceUpdateDatasetApiStatus:
|
||||
current_time = datetime.datetime(2023, 1, 1, 12, 0, 0)
|
||||
|
||||
# Act
|
||||
with (
|
||||
patch("services.dataset_service.current_user", owner),
|
||||
patch("services.dataset_service.naive_utc_now", return_value=current_time),
|
||||
):
|
||||
DatasetService.update_dataset_api_status(dataset.id, False, session=db_session_with_containers)
|
||||
with patch("services.dataset_service.naive_utc_now", return_value=current_time):
|
||||
DatasetService.update_dataset_api_status(dataset, False, owner, session=db_session_with_containers)
|
||||
|
||||
# Assert
|
||||
db_session_with_containers.refresh(dataset)
|
||||
assert dataset.enable_api is False
|
||||
assert dataset.updated_by == owner.id
|
||||
|
||||
def test_update_dataset_api_status_not_found_error(self, db_session_with_containers: Session):
|
||||
"""
|
||||
Test error handling when dataset is not found.
|
||||
|
||||
Verifies that when the dataset ID doesn't exist, a NotFound
|
||||
exception is raised.
|
||||
|
||||
This test ensures:
|
||||
- NotFound exception is raised
|
||||
- No updates are performed
|
||||
- Error message is appropriate
|
||||
"""
|
||||
# Arrange
|
||||
dataset_id = str(uuid4())
|
||||
|
||||
# Act & Assert
|
||||
with pytest.raises(NotFound, match="Dataset not found"):
|
||||
DatasetService.update_dataset_api_status(dataset_id, True, session=db_session_with_containers)
|
||||
|
||||
def test_update_dataset_api_status_missing_current_user_error(self, db_session_with_containers: Session):
|
||||
"""
|
||||
Test error handling when current_user is missing.
|
||||
@@ -372,13 +349,11 @@ class TestDatasetServiceUpdateDatasetApiStatus:
|
||||
dataset = DatasetUpdateDeleteTestDataFactory.create_dataset(
|
||||
db_session_with_containers, tenant.id, owner.id, enable_api=False
|
||||
)
|
||||
|
||||
actor = Account(name="missing-id", email="missing-id@example.com")
|
||||
actor.id = ""
|
||||
# Act & Assert
|
||||
with (
|
||||
patch("services.dataset_service.current_user", None),
|
||||
pytest.raises(ValueError, match="Current user or current user id not found"),
|
||||
):
|
||||
DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers)
|
||||
with pytest.raises(ValueError, match="Current user or current user id not found"):
|
||||
DatasetService.update_dataset_api_status(dataset, True, actor, session=db_session_with_containers)
|
||||
|
||||
# Verify no commit was attempted
|
||||
db_session_with_containers.rollback()
|
||||
|
||||
+11
-6
@@ -142,7 +142,9 @@ def test_get_document_queries_by_dataset_and_document_id(db_session_with_contain
|
||||
def test_get_documents_by_ids_returns_empty_for_empty_input(db_session_with_containers: Session):
|
||||
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
|
||||
|
||||
result = DocumentService.get_documents_by_ids(dataset.id, [], session=db_session_with_containers)
|
||||
result = DocumentService.get_documents_by_ids(
|
||||
DatasetRefService.create_dataset_ref(dataset), [], session=db_session_with_containers
|
||||
)
|
||||
|
||||
assert result == []
|
||||
|
||||
@@ -157,7 +159,9 @@ def test_get_documents_by_ids_uses_single_batch_query(db_session_with_containers
|
||||
position=2,
|
||||
)
|
||||
|
||||
result = DocumentService.get_documents_by_ids(dataset.id, [doc_a.id, doc_b.id], db_session_with_containers)
|
||||
result = DocumentService.get_documents_by_ids(
|
||||
DatasetRefService.create_dataset_ref(dataset), [doc_a.id, doc_b.id], db_session_with_containers
|
||||
)
|
||||
|
||||
assert {document.id for document in result} == {doc_a.id, doc_b.id}
|
||||
|
||||
@@ -319,7 +323,7 @@ def test_get_upload_files_by_document_id_for_zip_download_raises_for_missing_doc
|
||||
)
|
||||
|
||||
|
||||
def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_access(
|
||||
def test_get_upload_files_by_document_id_for_zip_download_hides_cross_tenant_documents(
|
||||
db_session_with_containers: Session,
|
||||
):
|
||||
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
|
||||
@@ -335,7 +339,7 @@ def test_get_upload_files_by_document_id_for_zip_download_rejects_cross_tenant_a
|
||||
data_source_info={"upload_file_id": upload_file.id},
|
||||
)
|
||||
|
||||
with pytest.raises(Forbidden, match="No permission"):
|
||||
with pytest.raises(NotFound, match="Document not found"):
|
||||
DocumentService._get_upload_files_by_document_id_for_zip_download(
|
||||
dataset_id=dataset.id,
|
||||
document_ids=[document.id],
|
||||
@@ -527,7 +531,7 @@ def test_get_working_documents_by_dataset_id_returns_completed_enabled_unarchive
|
||||
assert [document.id for document in result] == [available_document.id]
|
||||
|
||||
|
||||
def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db_session_with_containers: Session):
|
||||
def test_get_error_documents_by_dataset_ref_returns_error_and_paused_documents(db_session_with_containers: Session):
|
||||
dataset = DocumentServiceIntegrationFactory.create_dataset(db_session_with_containers)
|
||||
error_document = DocumentServiceIntegrationFactory.create_document(
|
||||
db_session_with_containers,
|
||||
@@ -547,7 +551,8 @@ def test_get_error_documents_by_dataset_id_returns_error_and_paused_documents(db
|
||||
indexing_status=IndexingStatus.COMPLETED,
|
||||
)
|
||||
|
||||
result = DocumentService.get_error_documents_by_dataset_id(dataset.id, session=db_session_with_containers)
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
result = DocumentService.get_error_documents_by_dataset_ref(dataset_ref, session=db_session_with_containers)
|
||||
|
||||
assert {document.id for document in result} == {error_document.id, paused_document.id}
|
||||
|
||||
|
||||
+19
-29
@@ -8,7 +8,6 @@ from uuid import uuid4
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import NotFound
|
||||
|
||||
from core.rag.index_processor.constant.index_type import IndexTechniqueType
|
||||
from models.account import Account, Tenant, TenantAccountJoin, TenantAccountRole
|
||||
@@ -21,6 +20,7 @@ from models.dataset import (
|
||||
DatasetPermissionEnum,
|
||||
)
|
||||
from models.enums import DataSourceType
|
||||
from services.dataset_ref_service import DatasetRef, DatasetRefService
|
||||
from services.dataset_service import DatasetCollectionBindingService, DatasetPermissionService, DatasetService
|
||||
from services.errors.account import NoPermissionError
|
||||
|
||||
@@ -213,7 +213,9 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
dataset_id=dataset.id,
|
||||
)
|
||||
|
||||
assert DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) is True
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
|
||||
assert DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers) is True
|
||||
|
||||
def test_dataset_use_check_returns_false_when_join_missing(self, db_session_with_containers: Session):
|
||||
owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers)
|
||||
@@ -223,7 +225,9 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
created_by=owner.id,
|
||||
)
|
||||
|
||||
assert DatasetService.dataset_use_check(dataset.id, session=db_session_with_containers) is False
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
|
||||
assert DatasetService.dataset_use_check(dataset_ref, session=db_session_with_containers) is False
|
||||
|
||||
def test_check_dataset_permission_rejects_cross_tenant_access(self, db_session_with_containers: Session):
|
||||
owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers)
|
||||
@@ -371,13 +375,6 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
user=operator, dataset=dataset, session=db_session_with_containers
|
||||
)
|
||||
|
||||
def test_update_dataset_api_status_raises_not_found_for_missing_dataset(
|
||||
self, flask_app_with_containers: Flask, db_session_with_containers: Session
|
||||
):
|
||||
with flask_app_with_containers.app_context():
|
||||
with pytest.raises(NotFound, match="Dataset not found"):
|
||||
DatasetService.update_dataset_api_status(str(uuid4()), True, session=db_session_with_containers)
|
||||
|
||||
def test_update_dataset_api_status_requires_current_user_id(self, db_session_with_containers: Session):
|
||||
owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers)
|
||||
dataset = DatasetPermissionIntegrationFactory.create_dataset(
|
||||
@@ -386,10 +383,10 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
created_by=owner.id,
|
||||
enable_api=False,
|
||||
)
|
||||
|
||||
with patch("services.dataset_service.current_user", SimpleNamespace(id=None)):
|
||||
with pytest.raises(ValueError, match="Current user or current user id not found"):
|
||||
DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers)
|
||||
actor = Account(name="missing-id", email="missing-id@example.com")
|
||||
actor.id = ""
|
||||
with pytest.raises(ValueError, match="Current user or current user id not found"):
|
||||
DatasetService.update_dataset_api_status(dataset, True, actor, session=db_session_with_containers)
|
||||
|
||||
def test_update_dataset_api_status_updates_fields_and_commits(self, db_session_with_containers: Session):
|
||||
owner, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant(db_session_with_containers)
|
||||
@@ -401,11 +398,8 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
)
|
||||
now = datetime(2026, 4, 14, 18, 0, 0)
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.current_user", owner),
|
||||
patch("services.dataset_service.naive_utc_now", return_value=now),
|
||||
):
|
||||
DatasetService.update_dataset_api_status(dataset.id, True, session=db_session_with_containers)
|
||||
with patch("services.dataset_service.naive_utc_now", return_value=now):
|
||||
DatasetService.update_dataset_api_status(dataset, True, owner, session=db_session_with_containers)
|
||||
|
||||
db_session_with_containers.refresh(dataset)
|
||||
assert dataset.enable_api is True
|
||||
@@ -419,12 +413,10 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
features = SimpleNamespace(
|
||||
billing=SimpleNamespace(enabled=False, subscription=SimpleNamespace(plan="professional"))
|
||||
)
|
||||
dataset_ref = DatasetRef(tenant_id=tenant.id, dataset_id=str(uuid4()))
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.current_user", owner),
|
||||
patch("services.dataset_service.FeatureService.get_features", return_value=features),
|
||||
):
|
||||
result = DatasetService.get_dataset_auto_disable_logs(str(uuid4()), session=db_session_with_containers)
|
||||
with patch("services.dataset_service.FeatureService.get_features", return_value=features):
|
||||
result = DatasetService.get_dataset_auto_disable_logs(dataset_ref, session=db_session_with_containers)
|
||||
|
||||
assert result == {"document_ids": [], "count": 0}
|
||||
|
||||
@@ -450,12 +442,10 @@ class TestDatasetServicePermissionsAndLifecycle:
|
||||
features = SimpleNamespace(
|
||||
billing=SimpleNamespace(enabled=True, subscription=SimpleNamespace(plan="professional"))
|
||||
)
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.current_user", owner),
|
||||
patch("services.dataset_service.FeatureService.get_features", return_value=features),
|
||||
):
|
||||
result = DatasetService.get_dataset_auto_disable_logs(dataset.id, session=db_session_with_containers)
|
||||
with patch("services.dataset_service.FeatureService.get_features", return_value=features):
|
||||
result = DatasetService.get_dataset_auto_disable_logs(dataset_ref, session=db_session_with_containers)
|
||||
|
||||
assert result["count"] == 2
|
||||
assert len(result["document_ids"]) == 2
|
||||
|
||||
+117
-14
@@ -1,16 +1,18 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from threading import Barrier
|
||||
from unittest.mock import patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from models import Account, Tenant
|
||||
from models.dataset import Dataset, DatasetMetadataBinding, Document
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom
|
||||
from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document
|
||||
from models.enums import DatasetMetadataType, DataSourceType, DocumentCreatedFrom
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
MetadataDetail,
|
||||
@@ -54,6 +56,19 @@ def _create_document(
|
||||
return document
|
||||
|
||||
|
||||
def _create_metadata(db_session: Session, *, dataset: Dataset, created_by: str, name: str) -> DatasetMetadata:
|
||||
metadata = DatasetMetadata(
|
||||
tenant_id=dataset.tenant_id,
|
||||
dataset_id=dataset.id,
|
||||
type=DatasetMetadataType.STRING,
|
||||
name=name,
|
||||
created_by=created_by,
|
||||
)
|
||||
db_session.add(metadata)
|
||||
db_session.commit()
|
||||
return metadata
|
||||
|
||||
|
||||
class TestMetadataPartialUpdate:
|
||||
@pytest.fixture
|
||||
def tenant_id(self) -> str:
|
||||
@@ -87,10 +102,12 @@ class TestMetadataPartialUpdate:
|
||||
doc_metadata={"existing_key": "existing_value"},
|
||||
)
|
||||
|
||||
meta_id = str(uuid4())
|
||||
metadata = _create_metadata(
|
||||
db_session_with_containers, dataset=dataset, created_by=current_account.id, name="new_key"
|
||||
)
|
||||
operation = DocumentMetadataOperation(
|
||||
document_id=document.id,
|
||||
metadata_list=[MetadataDetail(id=meta_id, name="new_key", value="new_value")],
|
||||
metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="new_value")],
|
||||
partial_update=True,
|
||||
)
|
||||
metadata_args = MetadataOperationData(operation_data=[operation])
|
||||
@@ -105,6 +122,86 @@ class TestMetadataPartialUpdate:
|
||||
assert updated_doc.doc_metadata["existing_key"] == "existing_value"
|
||||
assert updated_doc.doc_metadata["new_key"] == "new_value"
|
||||
|
||||
def test_concurrent_partial_updates_keep_both_values(
|
||||
self,
|
||||
db_session_with_containers: Session,
|
||||
tenant_id: str,
|
||||
current_account: Account,
|
||||
) -> None:
|
||||
dataset = _create_dataset(db_session_with_containers, tenant_id=tenant_id)
|
||||
document = _create_document(
|
||||
db_session_with_containers,
|
||||
dataset_id=dataset.id,
|
||||
tenant_id=tenant_id,
|
||||
doc_metadata={},
|
||||
)
|
||||
metadatas = [
|
||||
_create_metadata(
|
||||
db_session_with_containers,
|
||||
dataset=dataset,
|
||||
created_by=current_account.id,
|
||||
name=name,
|
||||
)
|
||||
for name in ("author", "region")
|
||||
]
|
||||
operations = [
|
||||
MetadataOperationData(
|
||||
operation_data=[
|
||||
DocumentMetadataOperation(
|
||||
document_id=document.id,
|
||||
metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value=value)],
|
||||
partial_update=True,
|
||||
)
|
||||
]
|
||||
)
|
||||
for metadata, value in zip(metadatas, ("Alice", "EU"), strict=True)
|
||||
]
|
||||
dataset_id = dataset.id
|
||||
document_id = document.id
|
||||
session_factory = sessionmaker(bind=db_session_with_containers.get_bind())
|
||||
engine = db_session_with_containers.get_bind()
|
||||
locked_select_barrier = Barrier(2)
|
||||
locked_selects: list[None] = []
|
||||
|
||||
def synchronize_locked_select(
|
||||
_conn: object,
|
||||
_cursor: object,
|
||||
statement: str,
|
||||
_parameters: object,
|
||||
_context: object,
|
||||
_executemany: bool,
|
||||
) -> None:
|
||||
if "FROM documents" in statement and "FOR UPDATE" in statement:
|
||||
locked_selects.append(None)
|
||||
locked_select_barrier.wait(timeout=10)
|
||||
|
||||
def update(metadata_args: MetadataOperationData) -> None:
|
||||
with session_factory() as session:
|
||||
owned_dataset = session.get(Dataset, dataset_id)
|
||||
assert owned_dataset is not None
|
||||
MetadataService.update_documents_metadata(
|
||||
owned_dataset, metadata_args, current_account, session=session
|
||||
)
|
||||
|
||||
event.listen(engine, "before_cursor_execute", synchronize_locked_select)
|
||||
try:
|
||||
with (
|
||||
patch.object(MetadataService, "knowledge_base_metadata_lock_check"),
|
||||
patch("services.metadata_service.redis_client.delete"),
|
||||
ThreadPoolExecutor(max_workers=2) as executor,
|
||||
):
|
||||
futures = [executor.submit(update, operation) for operation in operations]
|
||||
for future in futures:
|
||||
future.result(timeout=20)
|
||||
finally:
|
||||
event.remove(engine, "before_cursor_execute", synchronize_locked_select)
|
||||
|
||||
db_session_with_containers.expire_all()
|
||||
updated_doc = db_session_with_containers.get(Document, document_id)
|
||||
assert updated_doc is not None
|
||||
assert len(locked_selects) == 2
|
||||
assert updated_doc.doc_metadata == {"author": "Alice", "region": "EU"}
|
||||
|
||||
def test_full_update_replaces_metadata(
|
||||
self,
|
||||
flask_app_with_containers: Flask,
|
||||
@@ -120,10 +217,12 @@ class TestMetadataPartialUpdate:
|
||||
doc_metadata={"existing_key": "existing_value"},
|
||||
)
|
||||
|
||||
meta_id = str(uuid4())
|
||||
metadata = _create_metadata(
|
||||
db_session_with_containers, dataset=dataset, created_by=current_account.id, name="new_key"
|
||||
)
|
||||
operation = DocumentMetadataOperation(
|
||||
document_id=document.id,
|
||||
metadata_list=[MetadataDetail(id=meta_id, name="new_key", value="new_value")],
|
||||
metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="new_value")],
|
||||
partial_update=False,
|
||||
)
|
||||
metadata_args = MetadataOperationData(operation_data=[operation])
|
||||
@@ -154,12 +253,14 @@ class TestMetadataPartialUpdate:
|
||||
doc_metadata={"existing_key": "existing_value"},
|
||||
)
|
||||
|
||||
meta_id = str(uuid4())
|
||||
metadata = _create_metadata(
|
||||
db_session_with_containers, dataset=dataset, created_by=current_account.id, name="existing_key"
|
||||
)
|
||||
existing_binding = DatasetMetadataBinding(
|
||||
tenant_id=tenant_id,
|
||||
dataset_id=dataset.id,
|
||||
document_id=document.id,
|
||||
metadata_id=meta_id,
|
||||
metadata_id=metadata.id,
|
||||
created_by=user_id,
|
||||
)
|
||||
db_session_with_containers.add(existing_binding)
|
||||
@@ -167,7 +268,7 @@ class TestMetadataPartialUpdate:
|
||||
|
||||
operation = DocumentMetadataOperation(
|
||||
document_id=document.id,
|
||||
metadata_list=[MetadataDetail(id=meta_id, name="existing_key", value="existing_value")],
|
||||
metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="existing_value")],
|
||||
partial_update=True,
|
||||
)
|
||||
metadata_args = MetadataOperationData(operation_data=[operation])
|
||||
@@ -180,7 +281,7 @@ class TestMetadataPartialUpdate:
|
||||
bindings = db_session_with_containers.scalars(
|
||||
select(DatasetMetadataBinding).where(
|
||||
DatasetMetadataBinding.document_id == document.id,
|
||||
DatasetMetadataBinding.metadata_id == meta_id,
|
||||
DatasetMetadataBinding.metadata_id == metadata.id,
|
||||
)
|
||||
).all()
|
||||
assert len(bindings) == 1
|
||||
@@ -200,10 +301,12 @@ class TestMetadataPartialUpdate:
|
||||
doc_metadata={"existing_key": "existing_value"},
|
||||
)
|
||||
|
||||
meta_id = str(uuid4())
|
||||
metadata = _create_metadata(
|
||||
db_session_with_containers, dataset=dataset, created_by=current_account.id, name="key"
|
||||
)
|
||||
operation = DocumentMetadataOperation(
|
||||
document_id=document.id,
|
||||
metadata_list=[MetadataDetail(id=meta_id, name="key", value="value")],
|
||||
metadata_list=[MetadataDetail(id=metadata.id, name=metadata.name, value="value")],
|
||||
partial_update=True,
|
||||
)
|
||||
metadata_args = MetadataOperationData(operation_data=[operation])
|
||||
|
||||
@@ -10,8 +10,14 @@ from core.rag.index_processor.constant.built_in_field import BuiltInField
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType
|
||||
from models import Account, AccountStatus, Tenant, TenantAccountJoin, TenantAccountRole, TenantStatus
|
||||
from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom
|
||||
from services.entities.knowledge_entities.knowledge_entities import MetadataArgs
|
||||
from models.enums import DataSourceType, DocumentCreatedFrom, IndexingStatus
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
MetadataArgs,
|
||||
MetadataDetail,
|
||||
MetadataOperationData,
|
||||
)
|
||||
from services.errors.metadata import MetadataResourceNotFoundError
|
||||
from services.metadata_service import MetadataService
|
||||
|
||||
|
||||
@@ -300,7 +306,7 @@ class TestMetadataService:
|
||||
# Act: Execute the method under test
|
||||
new_name = "new_name"
|
||||
result = MetadataService.update_metadata_name(
|
||||
dataset.id, metadata.id, new_name, account, tenant.id, session=db_session_with_containers
|
||||
dataset, metadata.id, new_name, account, session=db_session_with_containers
|
||||
)
|
||||
|
||||
# Assert: Verify the expected outcomes
|
||||
@@ -340,7 +346,7 @@ class TestMetadataService:
|
||||
# Act & Assert: Verify proper error handling
|
||||
with pytest.raises(ValueError, match="Metadata name cannot exceed 255 characters."):
|
||||
MetadataService.update_metadata_name(
|
||||
dataset.id, metadata.id, long_name, account, tenant.id, session=db_session_with_containers
|
||||
dataset, metadata.id, long_name, account, session=db_session_with_containers
|
||||
)
|
||||
|
||||
def test_update_metadata_name_already_exists(
|
||||
@@ -371,7 +377,7 @@ class TestMetadataService:
|
||||
# Try to update first metadata with second metadata's name
|
||||
with pytest.raises(ValueError, match="Metadata name already exists."):
|
||||
MetadataService.update_metadata_name(
|
||||
dataset.id, first_metadata.id, "second_metadata", account, tenant.id, session=db_session_with_containers
|
||||
dataset, first_metadata.id, "second_metadata", account, session=db_session_with_containers
|
||||
)
|
||||
|
||||
def test_update_metadata_name_conflicts_with_built_in_field(
|
||||
@@ -399,7 +405,7 @@ class TestMetadataService:
|
||||
|
||||
with pytest.raises(ValueError, match="Metadata name already exists in Built-in fields."):
|
||||
MetadataService.update_metadata_name(
|
||||
dataset.id, metadata.id, built_in_field_name, account, tenant.id, session=db_session_with_containers
|
||||
dataset, metadata.id, built_in_field_name, account, session=db_session_with_containers
|
||||
)
|
||||
|
||||
def test_update_metadata_name_not_found(
|
||||
@@ -424,7 +430,7 @@ class TestMetadataService:
|
||||
|
||||
# Act: Execute the method under test
|
||||
result = MetadataService.update_metadata_name(
|
||||
dataset.id, fake_metadata_id, new_name, account, tenant.id, session=db_session_with_containers
|
||||
dataset, fake_metadata_id, new_name, account, session=db_session_with_containers
|
||||
)
|
||||
|
||||
# Assert: Verify the method returns None when metadata is not found
|
||||
@@ -451,7 +457,7 @@ class TestMetadataService:
|
||||
)
|
||||
|
||||
# Act: Execute the method under test
|
||||
result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers)
|
||||
result = MetadataService.delete_metadata(dataset, metadata.id, session=db_session_with_containers)
|
||||
|
||||
# Assert: Verify the expected outcomes
|
||||
assert result is not None
|
||||
@@ -482,7 +488,7 @@ class TestMetadataService:
|
||||
fake_metadata_id = str(uuid.uuid4()) # Use valid UUID format
|
||||
|
||||
# Act: Execute the method under test
|
||||
result = MetadataService.delete_metadata(dataset.id, fake_metadata_id, session=db_session_with_containers)
|
||||
result = MetadataService.delete_metadata(dataset, fake_metadata_id, session=db_session_with_containers)
|
||||
|
||||
# Assert: Verify the method returns None when metadata is not found
|
||||
assert result is None
|
||||
@@ -528,7 +534,7 @@ class TestMetadataService:
|
||||
db_session_with_containers.commit()
|
||||
|
||||
# Act: Execute the method under test
|
||||
result = MetadataService.delete_metadata(dataset.id, metadata.id, session=db_session_with_containers)
|
||||
result = MetadataService.delete_metadata(dataset, metadata.id, session=db_session_with_containers)
|
||||
|
||||
# Assert: Verify the expected outcomes
|
||||
assert result is not None
|
||||
@@ -540,6 +546,69 @@ class TestMetadataService:
|
||||
# Note: The service attempts to update document metadata but may not succeed
|
||||
# due to mock configuration. The main functionality (metadata deletion) is verified.
|
||||
|
||||
@pytest.mark.parametrize("operation", ["rename", "delete"])
|
||||
@pytest.mark.parametrize("binding_owner", ["metadata", "document"])
|
||||
def test_metadata_changes_ignore_historical_foreign_document_binding(
|
||||
self,
|
||||
operation: str,
|
||||
binding_owner: str,
|
||||
db_session_with_containers: Session,
|
||||
mock_external_service_dependencies: MetadataServiceDeps,
|
||||
) -> None:
|
||||
account, tenant = self._create_test_account_and_tenant(
|
||||
db_session_with_containers, mock_external_service_dependencies
|
||||
)
|
||||
dataset = self._create_test_dataset(
|
||||
db_session_with_containers, mock_external_service_dependencies, account, tenant
|
||||
)
|
||||
foreign_account, foreign_tenant = self._create_test_account_and_tenant(
|
||||
db_session_with_containers, mock_external_service_dependencies
|
||||
)
|
||||
foreign_dataset = self._create_test_dataset(
|
||||
db_session_with_containers,
|
||||
mock_external_service_dependencies,
|
||||
foreign_account,
|
||||
foreign_tenant,
|
||||
)
|
||||
foreign_document = self._create_test_document(
|
||||
db_session_with_containers,
|
||||
mock_external_service_dependencies,
|
||||
foreign_dataset,
|
||||
foreign_account,
|
||||
)
|
||||
foreign_document.enabled = True
|
||||
foreign_document.archived = False
|
||||
foreign_document.indexing_status = IndexingStatus.COMPLETED
|
||||
foreign_document.doc_metadata = {"old_name": "foreign-value"}
|
||||
|
||||
metadata = MetadataService.create_metadata(
|
||||
dataset.id,
|
||||
MetadataArgs(type="string", name="old_name"),
|
||||
account,
|
||||
tenant.id,
|
||||
session=db_session_with_containers,
|
||||
)
|
||||
db_session_with_containers.add(
|
||||
DatasetMetadataBinding(
|
||||
tenant_id=dataset.tenant_id if binding_owner == "metadata" else foreign_dataset.tenant_id,
|
||||
dataset_id=dataset.id if binding_owner == "metadata" else foreign_dataset.id,
|
||||
metadata_id=metadata.id,
|
||||
document_id=foreign_document.id,
|
||||
created_by=account.id,
|
||||
)
|
||||
)
|
||||
db_session_with_containers.commit()
|
||||
|
||||
if operation == "rename":
|
||||
MetadataService.update_metadata_name(
|
||||
dataset, metadata.id, "new_name", account, session=db_session_with_containers
|
||||
)
|
||||
else:
|
||||
MetadataService.delete_metadata(dataset, metadata.id, session=db_session_with_containers)
|
||||
|
||||
db_session_with_containers.refresh(foreign_document)
|
||||
assert foreign_document.doc_metadata == {"old_name": "foreign-value"}
|
||||
|
||||
def test_get_built_in_fields_success(
|
||||
self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps
|
||||
) -> None:
|
||||
@@ -799,13 +868,6 @@ class TestMetadataService:
|
||||
# Mock DocumentService.get_document
|
||||
mock_external_service_dependencies["document_service"].get_document.return_value = document
|
||||
|
||||
# Create metadata operation data
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
MetadataDetail,
|
||||
MetadataOperationData,
|
||||
)
|
||||
|
||||
metadata_detail = MetadataDetail(id=metadata.id, name=metadata.name, value="test_value")
|
||||
|
||||
operation = DocumentMetadataOperation(document_id=document.id, metadata_list=[metadata_detail])
|
||||
@@ -833,6 +895,77 @@ class TestMetadataService:
|
||||
assert binding.tenant_id == tenant.id
|
||||
assert binding.dataset_id == dataset.id
|
||||
|
||||
@pytest.mark.parametrize("foreign_resource", ["metadata", "document"])
|
||||
def test_update_documents_metadata_rejects_foreign_owner_before_writes(
|
||||
self,
|
||||
foreign_resource: str,
|
||||
db_session_with_containers: Session,
|
||||
mock_external_service_dependencies: MetadataServiceDeps,
|
||||
) -> None:
|
||||
account, tenant = self._create_test_account_and_tenant(
|
||||
db_session_with_containers, mock_external_service_dependencies
|
||||
)
|
||||
dataset = self._create_test_dataset(
|
||||
db_session_with_containers, mock_external_service_dependencies, account, tenant
|
||||
)
|
||||
document = self._create_test_document(
|
||||
db_session_with_containers, mock_external_service_dependencies, dataset, account
|
||||
)
|
||||
metadata = MetadataService.create_metadata(
|
||||
dataset.id,
|
||||
MetadataArgs(type="string", name="owned"),
|
||||
account,
|
||||
tenant.id,
|
||||
session=db_session_with_containers,
|
||||
)
|
||||
|
||||
foreign_account, foreign_tenant = self._create_test_account_and_tenant(
|
||||
db_session_with_containers, mock_external_service_dependencies
|
||||
)
|
||||
foreign_dataset = self._create_test_dataset(
|
||||
db_session_with_containers,
|
||||
mock_external_service_dependencies,
|
||||
foreign_account,
|
||||
foreign_tenant,
|
||||
)
|
||||
foreign_document = self._create_test_document(
|
||||
db_session_with_containers,
|
||||
mock_external_service_dependencies,
|
||||
foreign_dataset,
|
||||
foreign_account,
|
||||
)
|
||||
foreign_metadata = MetadataService.create_metadata(
|
||||
foreign_dataset.id,
|
||||
MetadataArgs(type="string", name="foreign"),
|
||||
foreign_account,
|
||||
foreign_tenant.id,
|
||||
session=db_session_with_containers,
|
||||
)
|
||||
operation = DocumentMetadataOperation(
|
||||
document_id=foreign_document.id if foreign_resource == "document" else document.id,
|
||||
metadata_list=[
|
||||
MetadataDetail(
|
||||
id=foreign_metadata.id if foreign_resource == "metadata" else metadata.id,
|
||||
name="ignored",
|
||||
value="value",
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(MetadataResourceNotFoundError, match=f"{foreign_resource.capitalize()} not found"):
|
||||
MetadataService.update_documents_metadata(
|
||||
dataset,
|
||||
MetadataOperationData(operation_data=[operation]),
|
||||
account,
|
||||
session=db_session_with_containers,
|
||||
)
|
||||
|
||||
db_session_with_containers.refresh(document)
|
||||
db_session_with_containers.refresh(foreign_document)
|
||||
assert document.doc_metadata is None
|
||||
assert foreign_document.doc_metadata is None
|
||||
assert db_session_with_containers.query(DatasetMetadataBinding).count() == 0
|
||||
|
||||
def test_update_documents_metadata_with_built_in_fields_enabled(
|
||||
self, db_session_with_containers: Session, mock_external_service_dependencies: MetadataServiceDeps
|
||||
) -> None:
|
||||
@@ -865,13 +998,6 @@ class TestMetadataService:
|
||||
# Mock DocumentService.get_document
|
||||
mock_external_service_dependencies["document_service"].get_document.return_value = document
|
||||
|
||||
# Create metadata operation data
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
MetadataDetail,
|
||||
MetadataOperationData,
|
||||
)
|
||||
|
||||
metadata_detail = MetadataDetail(id=metadata.id, name=metadata.name, value="test_value")
|
||||
|
||||
operation = DocumentMetadataOperation(document_id=document.id, metadata_list=[metadata_detail])
|
||||
@@ -911,13 +1037,6 @@ class TestMetadataService:
|
||||
dataset.id, metadata_args, account, tenant.id, session=db_session_with_containers
|
||||
)
|
||||
|
||||
# Create metadata operation data
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
DocumentMetadataOperation,
|
||||
MetadataDetail,
|
||||
MetadataOperationData,
|
||||
)
|
||||
|
||||
metadata_detail = MetadataDetail(id=metadata.id, name=metadata.name, value="test_value")
|
||||
|
||||
# Use a valid UUID format that does not exist in the database
|
||||
@@ -927,9 +1046,7 @@ class TestMetadataService:
|
||||
|
||||
operation_data = MetadataOperationData(operation_data=[operation])
|
||||
|
||||
# Act & Assert: The method should raise ValueError("Document not found.")
|
||||
# because the exception is now re-raised after rollback
|
||||
with pytest.raises(ValueError, match="Document not found"):
|
||||
with pytest.raises(MetadataResourceNotFoundError, match="Document not found"):
|
||||
MetadataService.update_documents_metadata(
|
||||
dataset, operation_data, account, session=db_session_with_containers
|
||||
)
|
||||
|
||||
@@ -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