fix(api): bind dataset operations to owners (#40149)

This commit is contained in:
WH-2099
2026-08-11 16:37:14 +00:00
committed by GitHub
parent bee269afe8
commit ef8544b173
35 changed files with 1520 additions and 496 deletions
@@ -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()
@@ -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}
@@ -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
@@ -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)
@@ -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}