mirror of
https://github.com/langgenius/dify.git
synced 2026-09-19 10:11:30 +08:00
fix(api): bind dataset operations to owners (#40149)
This commit is contained in:
@@ -64,7 +64,7 @@ from models.model import UploadFile
|
||||
from models.provider_ids import ModelProviderID
|
||||
from models.source import DataSourceOauthBinding
|
||||
from models.workflow import Workflow
|
||||
from services.dataset_ref_service import DatasetRef, SegmentRef
|
||||
from services.dataset_ref_service import DatasetRef, DatasetRefService, SegmentRef
|
||||
from services.document_indexing_proxy.document_indexing_task_proxy import DocumentIndexingTaskProxy
|
||||
from services.document_indexing_proxy.duplicate_document_indexing_task_proxy import DuplicateDocumentIndexingTaskProxy
|
||||
from services.enterprise import rbac_service as enterprise_rbac_service
|
||||
@@ -1365,8 +1365,8 @@ class DatasetService:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def dataset_use_check(dataset_id, session: Session) -> bool:
|
||||
stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_id))
|
||||
def dataset_use_check(dataset_ref: DatasetRef, session: Session) -> bool:
|
||||
stmt = select(exists().where(AppDatasetJoin.dataset_id == dataset_ref.dataset_id))
|
||||
return session.execute(stmt).scalar_one()
|
||||
|
||||
@staticmethod
|
||||
@@ -1432,22 +1432,17 @@ class DatasetService:
|
||||
).all()
|
||||
|
||||
@staticmethod
|
||||
def update_dataset_api_status(dataset_id: str, status: bool, session: Session):
|
||||
dataset = DatasetService.get_dataset(dataset_id, session)
|
||||
if dataset is None:
|
||||
raise NotFound("Dataset not found.")
|
||||
dataset.enable_api = status
|
||||
if not current_user or not current_user.id:
|
||||
def update_dataset_api_status(dataset: Dataset, status: bool, actor: Account, session: Session):
|
||||
if not actor.id:
|
||||
raise ValueError("Current user or current user id not found")
|
||||
dataset.updated_by = current_user.id
|
||||
dataset.enable_api = status
|
||||
dataset.updated_by = actor.id
|
||||
dataset.updated_at = naive_utc_now()
|
||||
session.flush()
|
||||
|
||||
@staticmethod
|
||||
def get_dataset_auto_disable_logs(dataset_id: str, session: Session) -> AutoDisableLogsDict:
|
||||
assert isinstance(current_user, Account)
|
||||
assert current_user.current_tenant_id is not None
|
||||
features = FeatureService.get_features(current_user.current_tenant_id, exclude_vector_space=True)
|
||||
def get_dataset_auto_disable_logs(dataset_ref: DatasetRef, session: Session) -> AutoDisableLogsDict:
|
||||
features = FeatureService.get_features(dataset_ref.tenant_id, exclude_vector_space=True)
|
||||
if not features.billing.enabled or features.billing.subscription.plan == CloudPlan.SANDBOX:
|
||||
return {
|
||||
"document_ids": [],
|
||||
@@ -1457,7 +1452,8 @@ class DatasetService:
|
||||
start_date = datetime.datetime.now() - datetime.timedelta(days=30)
|
||||
dataset_auto_disable_logs = session.scalars(
|
||||
select(DatasetAutoDisableLog).where(
|
||||
DatasetAutoDisableLog.dataset_id == dataset_id,
|
||||
DatasetAutoDisableLog.tenant_id == dataset_ref.tenant_id,
|
||||
DatasetAutoDisableLog.dataset_id == dataset_ref.dataset_id,
|
||||
DatasetAutoDisableLog.created_at >= start_date,
|
||||
)
|
||||
).all()
|
||||
@@ -1654,7 +1650,9 @@ class DocumentService:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_documents_by_ids(dataset_id: str, document_ids: Sequence[str], session: Session) -> Sequence[Document]:
|
||||
def get_documents_by_ids(
|
||||
dataset_ref: DatasetRef, document_ids: Sequence[str], session: Session
|
||||
) -> Sequence[Document]:
|
||||
"""Fetch documents for a dataset in a single batch query."""
|
||||
if not document_ids:
|
||||
return []
|
||||
@@ -1662,7 +1660,8 @@ class DocumentService:
|
||||
# Fetch all requested documents in one query to avoid N+1 lookups.
|
||||
documents: Sequence[Document] = session.scalars(
|
||||
select(Document).where(
|
||||
Document.dataset_id == dataset_id,
|
||||
Document.tenant_id == dataset_ref.tenant_id,
|
||||
Document.dataset_id == dataset_ref.dataset_id,
|
||||
Document.id.in_(document_id_list),
|
||||
)
|
||||
).all()
|
||||
@@ -1856,7 +1855,9 @@ class DocumentService:
|
||||
"""
|
||||
document_id_list: list[str] = [str(document_id) for document_id in document_ids]
|
||||
|
||||
documents = DocumentService.get_documents_by_ids(dataset_id, document_id_list, session)
|
||||
documents = DocumentService.get_documents_by_ids(
|
||||
DatasetRef(tenant_id=tenant_id, dataset_id=dataset_id), document_id_list, session
|
||||
)
|
||||
documents_by_id: dict[str, Document] = {str(document.id): document for document in documents}
|
||||
|
||||
missing_document_ids: set[str] = set(document_id_list) - set(documents_by_id.keys())
|
||||
@@ -1866,9 +1867,6 @@ class DocumentService:
|
||||
upload_file_ids: list[str] = []
|
||||
upload_file_ids_by_document_id: dict[str, str] = {}
|
||||
for document_id, document in documents_by_id.items():
|
||||
if document.tenant_id != tenant_id:
|
||||
raise Forbidden("No permission.")
|
||||
|
||||
upload_file_id = DocumentService._get_upload_file_id_for_upload_file_document(
|
||||
document,
|
||||
invalid_source_message="Only uploaded-file documents can be downloaded as ZIP.",
|
||||
@@ -1895,9 +1893,13 @@ class DocumentService:
|
||||
return document
|
||||
|
||||
@staticmethod
|
||||
def get_document_by_ids(document_ids: list[str], session: Session) -> Sequence[Document]:
|
||||
def get_document_by_ids(
|
||||
dataset_ref: DatasetRef, document_ids: Sequence[str], session: Session
|
||||
) -> Sequence[Document]:
|
||||
documents = session.scalars(
|
||||
select(Document).where(
|
||||
Document.tenant_id == dataset_ref.tenant_id,
|
||||
Document.dataset_id == dataset_ref.dataset_id,
|
||||
Document.id.in_(document_ids),
|
||||
Document.enabled == True,
|
||||
Document.indexing_status == IndexingStatus.COMPLETED,
|
||||
@@ -1931,10 +1933,11 @@ class DocumentService:
|
||||
return documents
|
||||
|
||||
@staticmethod
|
||||
def get_error_documents_by_dataset_id(dataset_id: str, session: Session) -> Sequence[Document]:
|
||||
def get_error_documents_by_dataset_ref(dataset_ref: DatasetRef, session: Session) -> Sequence[Document]:
|
||||
documents = session.scalars(
|
||||
select(Document).where(
|
||||
Document.dataset_id == dataset_id,
|
||||
Document.tenant_id == dataset_ref.tenant_id,
|
||||
Document.dataset_id == dataset_ref.dataset_id,
|
||||
Document.indexing_status.in_([IndexingStatus.ERROR, IndexingStatus.PAUSED]),
|
||||
)
|
||||
).all()
|
||||
@@ -2143,7 +2146,11 @@ class DocumentService:
|
||||
retry_document_indexing_task.delay(dataset_id, document_ids, current_user.id)
|
||||
|
||||
@staticmethod
|
||||
def sync_website_document(dataset_id: str, document: Document, session: Session):
|
||||
def sync_website_document(dataset: Dataset, document: Document, session: Session):
|
||||
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
|
||||
if DatasetRefService.create_document_ref(dataset_ref, document) is None:
|
||||
raise ValueError("Document not found.")
|
||||
|
||||
# add sync flag
|
||||
sync_indexing_cache_key = f"document_{document.id}_is_sync"
|
||||
cache_result = redis_client.get(sync_indexing_cache_key)
|
||||
@@ -2160,7 +2167,7 @@ class DocumentService:
|
||||
|
||||
redis_client.setex(sync_indexing_cache_key, 600, 1)
|
||||
|
||||
sync_website_document_indexing_task.delay(dataset_id, document.id)
|
||||
sync_website_document_indexing_task.delay(dataset.id, document.id)
|
||||
|
||||
@staticmethod
|
||||
def get_documents_position(dataset_id, session: Session):
|
||||
|
||||
@@ -6,6 +6,7 @@ from core.rag.entities import Rule
|
||||
from core.rag.entities.metadata_entities import MetadataFilteringCondition
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType
|
||||
from core.rag.retrieval.retrieval_methods import RetrievalMethod
|
||||
from libs.helper import UUIDStr
|
||||
from models.enums import ProcessRuleMode
|
||||
|
||||
DocForm = Annotated[
|
||||
@@ -283,7 +284,7 @@ class MetadataUpdateArgs(BaseModel):
|
||||
|
||||
|
||||
class MetadataDetail(BaseModel):
|
||||
id: str = Field(description="Metadata field ID.")
|
||||
id: UUIDStr = Field(description="Metadata field ID.")
|
||||
name: str = Field(description="Metadata field name.")
|
||||
value: str | int | float | None = Field(
|
||||
default=None,
|
||||
@@ -292,7 +293,7 @@ class MetadataDetail(BaseModel):
|
||||
|
||||
|
||||
class DocumentMetadataOperation(BaseModel):
|
||||
document_id: str = Field(description="Document ID whose metadata should be updated.")
|
||||
document_id: UUIDStr = Field(description="Document ID whose metadata should be updated.")
|
||||
metadata_list: list[MetadataDetail] = Field(description="Metadata fields to update.")
|
||||
partial_update: bool = Field(
|
||||
default=False,
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
class MetadataResourceNotFoundError(Exception):
|
||||
pass
|
||||
@@ -9,13 +9,15 @@ from extensions.ext_redis import redis_client
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.login import resolve_account_fallback
|
||||
from models import Account
|
||||
from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding
|
||||
from models.dataset import Dataset, DatasetMetadata, DatasetMetadataBinding, Document
|
||||
from models.enums import DatasetMetadataType
|
||||
from services.dataset_ref_service import DatasetRefService
|
||||
from services.dataset_service import DocumentService
|
||||
from services.entities.knowledge_entities.knowledge_entities import (
|
||||
MetadataArgs,
|
||||
MetadataOperationData,
|
||||
)
|
||||
from services.errors.metadata import MetadataResourceNotFoundError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -61,11 +63,10 @@ class MetadataService:
|
||||
|
||||
@staticmethod
|
||||
def update_metadata_name(
|
||||
dataset_id: str,
|
||||
dataset: Dataset,
|
||||
metadata_id: str,
|
||||
name: str,
|
||||
current_user: Account | None = None,
|
||||
current_tenant_id: str | None = None, # TODO: the service_api is not migrated yet
|
||||
current_user: Account,
|
||||
*,
|
||||
session: Session,
|
||||
) -> DatasetMetadata | None:
|
||||
@@ -73,14 +74,13 @@ class MetadataService:
|
||||
if len(name) > 255:
|
||||
raise ValueError("Metadata name cannot exceed 255 characters.")
|
||||
|
||||
lock_key = f"dataset_metadata_lock_{dataset_id}"
|
||||
lock_key = f"dataset_metadata_lock_{dataset.id}"
|
||||
# check if metadata name already exists
|
||||
current_user, current_tenant_id = resolve_account_fallback(current_user, current_tenant_id)
|
||||
if session.scalar(
|
||||
select(DatasetMetadata)
|
||||
.where(
|
||||
DatasetMetadata.tenant_id == current_tenant_id,
|
||||
DatasetMetadata.dataset_id == dataset_id,
|
||||
DatasetMetadata.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadata.dataset_id == dataset.id,
|
||||
DatasetMetadata.name == name,
|
||||
)
|
||||
.limit(1)
|
||||
@@ -90,10 +90,14 @@ class MetadataService:
|
||||
if field.value == name:
|
||||
raise ValueError("Metadata name already exists in Built-in fields.")
|
||||
try:
|
||||
MetadataService.knowledge_base_metadata_lock_check(dataset_id, None)
|
||||
MetadataService.knowledge_base_metadata_lock_check(dataset.id, None)
|
||||
metadata = session.scalar(
|
||||
select(DatasetMetadata)
|
||||
.where(DatasetMetadata.id == metadata_id, DatasetMetadata.dataset_id == dataset_id)
|
||||
.where(
|
||||
DatasetMetadata.id == metadata_id,
|
||||
DatasetMetadata.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadata.dataset_id == dataset.id,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
if metadata is None:
|
||||
@@ -105,11 +109,17 @@ class MetadataService:
|
||||
|
||||
# update related documents
|
||||
dataset_metadata_bindings = session.scalars(
|
||||
select(DatasetMetadataBinding).where(DatasetMetadataBinding.metadata_id == metadata_id)
|
||||
select(DatasetMetadataBinding).where(
|
||||
DatasetMetadataBinding.metadata_id == metadata_id,
|
||||
DatasetMetadataBinding.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadataBinding.dataset_id == dataset.id,
|
||||
)
|
||||
).all()
|
||||
if dataset_metadata_bindings:
|
||||
document_ids = [binding.document_id for binding in dataset_metadata_bindings]
|
||||
documents = DocumentService.get_document_by_ids(document_ids, session)
|
||||
documents = DocumentService.get_document_by_ids(
|
||||
DatasetRefService.create_dataset_ref(dataset), document_ids, session
|
||||
)
|
||||
for document in documents:
|
||||
if not document.doc_metadata:
|
||||
doc_metadata = {}
|
||||
@@ -128,13 +138,17 @@ class MetadataService:
|
||||
redis_client.delete(lock_key)
|
||||
|
||||
@staticmethod
|
||||
def delete_metadata(dataset_id: str, metadata_id: str, session: Session):
|
||||
lock_key = f"dataset_metadata_lock_{dataset_id}"
|
||||
def delete_metadata(dataset: Dataset, metadata_id: str, session: Session):
|
||||
lock_key = f"dataset_metadata_lock_{dataset.id}"
|
||||
try:
|
||||
MetadataService.knowledge_base_metadata_lock_check(dataset_id, None)
|
||||
MetadataService.knowledge_base_metadata_lock_check(dataset.id, None)
|
||||
metadata = session.scalar(
|
||||
select(DatasetMetadata)
|
||||
.where(DatasetMetadata.id == metadata_id, DatasetMetadata.dataset_id == dataset_id)
|
||||
.where(
|
||||
DatasetMetadata.id == metadata_id,
|
||||
DatasetMetadata.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadata.dataset_id == dataset.id,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
if metadata is None:
|
||||
@@ -143,11 +157,17 @@ class MetadataService:
|
||||
|
||||
# deal related documents
|
||||
dataset_metadata_bindings = session.scalars(
|
||||
select(DatasetMetadataBinding).where(DatasetMetadataBinding.metadata_id == metadata_id)
|
||||
select(DatasetMetadataBinding).where(
|
||||
DatasetMetadataBinding.metadata_id == metadata_id,
|
||||
DatasetMetadataBinding.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadataBinding.dataset_id == dataset.id,
|
||||
)
|
||||
).all()
|
||||
if dataset_metadata_bindings:
|
||||
document_ids = [binding.document_id for binding in dataset_metadata_bindings]
|
||||
documents = DocumentService.get_document_by_ids(document_ids, session)
|
||||
documents = DocumentService.get_document_by_ids(
|
||||
DatasetRefService.create_dataset_ref(dataset), document_ids, session
|
||||
)
|
||||
for document in documents:
|
||||
if not document.doc_metadata:
|
||||
doc_metadata = {}
|
||||
@@ -237,27 +257,61 @@ class MetadataService:
|
||||
def update_documents_metadata(
|
||||
dataset: Dataset,
|
||||
metadata_args: MetadataOperationData,
|
||||
current_user: Account | None = None, # TODO: the service_api is not migrated yet
|
||||
current_tenant_id: str | None = None,
|
||||
current_user: Account,
|
||||
*,
|
||||
session: Session,
|
||||
):
|
||||
current_user, current_tenant_id = resolve_account_fallback(
|
||||
current_user, current_tenant_id, fallback_tenant_id=dataset.tenant_id
|
||||
metadata_ids = {
|
||||
metadata_value.id
|
||||
for operation in metadata_args.operation_data
|
||||
for metadata_value in operation.metadata_list
|
||||
}
|
||||
metadatas = session.scalars(
|
||||
select(DatasetMetadata).where(
|
||||
DatasetMetadata.id.in_(metadata_ids),
|
||||
DatasetMetadata.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadata.dataset_id == dataset.id,
|
||||
)
|
||||
).all()
|
||||
metadata_by_id = {metadata.id: metadata for metadata in metadatas}
|
||||
if metadata_ids != set(metadata_by_id):
|
||||
raise MetadataResourceNotFoundError("Metadata not found.")
|
||||
|
||||
document_ids = {operation.document_id for operation in metadata_args.operation_data}
|
||||
owned_document_ids = set(
|
||||
session.scalars(
|
||||
select(Document.id).where(
|
||||
Document.id.in_(document_ids),
|
||||
Document.tenant_id == dataset.tenant_id,
|
||||
Document.dataset_id == dataset.id,
|
||||
)
|
||||
).all()
|
||||
)
|
||||
if document_ids != owned_document_ids:
|
||||
raise MetadataResourceNotFoundError("Document not found.")
|
||||
|
||||
for operation in metadata_args.operation_data:
|
||||
lock_key = f"document_metadata_lock_{operation.document_id}"
|
||||
try:
|
||||
MetadataService.knowledge_base_metadata_lock_check(None, operation.document_id)
|
||||
document = DocumentService.get_document(dataset.id, operation.document_id, session=session)
|
||||
document = session.scalar(
|
||||
select(Document)
|
||||
.where(
|
||||
Document.id == operation.document_id,
|
||||
Document.tenant_id == dataset.tenant_id,
|
||||
Document.dataset_id == dataset.id,
|
||||
)
|
||||
.with_for_update()
|
||||
.execution_options(populate_existing=True)
|
||||
)
|
||||
if document is None:
|
||||
raise ValueError("Document not found.")
|
||||
raise MetadataResourceNotFoundError("Document not found.")
|
||||
if operation.partial_update:
|
||||
doc_metadata = copy.deepcopy(document.doc_metadata) if document.doc_metadata else {}
|
||||
else:
|
||||
doc_metadata = {}
|
||||
for metadata_value in operation.metadata_list:
|
||||
doc_metadata[metadata_value.name] = metadata_value.value
|
||||
doc_metadata[metadata_by_id[metadata_value.id].name] = metadata_value.value
|
||||
if dataset.built_in_field_enabled:
|
||||
doc_metadata[BuiltInField.document_name] = document.name
|
||||
doc_metadata[BuiltInField.uploader] = document.get_uploader(session=session)
|
||||
@@ -271,7 +325,9 @@ class MetadataService:
|
||||
if not operation.partial_update:
|
||||
session.execute(
|
||||
delete(DatasetMetadataBinding).where(
|
||||
DatasetMetadataBinding.document_id == operation.document_id
|
||||
DatasetMetadataBinding.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadataBinding.dataset_id == dataset.id,
|
||||
DatasetMetadataBinding.document_id == document.id,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -281,7 +337,9 @@ class MetadataService:
|
||||
existing_binding = session.scalar(
|
||||
select(DatasetMetadataBinding)
|
||||
.where(
|
||||
DatasetMetadataBinding.document_id == operation.document_id,
|
||||
DatasetMetadataBinding.tenant_id == dataset.tenant_id,
|
||||
DatasetMetadataBinding.dataset_id == dataset.id,
|
||||
DatasetMetadataBinding.document_id == document.id,
|
||||
DatasetMetadataBinding.metadata_id == metadata_value.id,
|
||||
)
|
||||
.limit(1)
|
||||
@@ -290,9 +348,9 @@ class MetadataService:
|
||||
continue
|
||||
|
||||
dataset_metadata_binding = DatasetMetadataBinding(
|
||||
tenant_id=current_tenant_id,
|
||||
tenant_id=dataset.tenant_id,
|
||||
dataset_id=dataset.id,
|
||||
document_id=operation.document_id,
|
||||
document_id=document.id,
|
||||
metadata_id=metadata_value.id,
|
||||
created_by=current_user.id,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user