fix(api): bind scoped lookups to nested owner refs (#38470)

This commit is contained in:
WH-2099
2026-07-15 09:31:09 +00:00
committed by GitHub
parent b721e7a32d
commit 6a511da325
28 changed files with 708 additions and 215 deletions
+13 -10
View File
@@ -96,7 +96,7 @@ class AppAnnotationService:
select(MessageAnnotation)
.where(
MessageAnnotation.id == annotation_ref.annotation_id,
MessageAnnotation.app_id == annotation_ref.app_id,
MessageAnnotation.app_id == annotation_ref.app.app_id,
)
.limit(1)
)
@@ -343,15 +343,15 @@ class AppAnnotationService:
session.commit()
# if annotation reply is enabled , add annotation to index
app_annotation_setting = session.scalar(
select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == annotation_ref.app_id).limit(1)
select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == annotation_ref.app.app_id).limit(1)
)
if app_annotation_setting:
update_annotation_to_index_task.delay(
annotation.id,
annotation.question_text,
annotation_ref.tenant_id,
annotation_ref.app_id,
annotation_ref.app.tenant_id,
annotation_ref.app.app_id,
app_annotation_setting.collection_binding_id,
)
@@ -368,7 +368,7 @@ class AppAnnotationService:
annotation_hit_histories = session.scalars(
select(AppAnnotationHitHistory).where(
AppAnnotationHitHistory.app_id == annotation_ref.app_id,
AppAnnotationHitHistory.app_id == annotation_ref.app.app_id,
AppAnnotationHitHistory.annotation_id == annotation_ref.annotation_id,
)
).all()
@@ -379,14 +379,14 @@ class AppAnnotationService:
session.commit()
# if annotation reply is enabled , delete annotation index
app_annotation_setting = session.scalar(
select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == annotation_ref.app_id).limit(1)
select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == annotation_ref.app.app_id).limit(1)
)
if app_annotation_setting:
delete_annotation_index_task.delay(
annotation.id,
annotation_ref.app_id,
annotation_ref.tenant_id,
annotation_ref.app.app_id,
annotation_ref.app.tenant_id,
app_annotation_setting.collection_binding_id,
)
@@ -396,7 +396,10 @@ class AppAnnotationService:
annotations_to_delete = session.execute(
select(MessageAnnotation, AppAnnotationSetting)
.outerjoin(AppAnnotationSetting, MessageAnnotation.app_id == AppAnnotationSetting.app_id)
.where(MessageAnnotation.id.in_(annotation_ids), MessageAnnotation.app_id == app_ref.app_id)
.where(
MessageAnnotation.id.in_(annotation_ids),
MessageAnnotation.app_id == app_ref.app_id,
)
).all()
if not annotations_to_delete:
@@ -576,7 +579,7 @@ class AppAnnotationService:
stmt = (
select(AppAnnotationHitHistory)
.where(
AppAnnotationHitHistory.app_id == annotation_ref.app_id,
AppAnnotationHitHistory.app_id == annotation_ref.app.app_id,
AppAnnotationHitHistory.annotation_id == annotation_ref.annotation_id,
)
.order_by(AppAnnotationHitHistory.created_at.desc())
+8 -10
View File
@@ -15,8 +15,7 @@ class AppRef(NamedTuple):
class MessageRef(NamedTuple):
"""Message identifiers used to scope downstream resource lookups."""
tenant_id: str
app_id: str
app: AppRef
message_id: str
end_user_id: str | None = None
account_id: str | None = None
@@ -25,16 +24,14 @@ class MessageRef(NamedTuple):
class AnnotationRef(NamedTuple):
"""Annotation identifiers used to scope downstream resource lookups."""
tenant_id: str
app_id: str
app: AppRef
annotation_id: str
class AppMCPServerRef(NamedTuple):
"""MCP server identifiers used to scope downstream resource lookups."""
tenant_id: str
app_id: str
app: AppRef
server_id: str
@@ -54,8 +51,7 @@ class AppRefService:
account_id: str | None = None,
) -> MessageRef:
return MessageRef(
tenant_id=app_ref.tenant_id,
app_id=app_ref.app_id,
app=app_ref,
message_id=message_id,
end_user_id=end_user_id,
account_id=account_id,
@@ -63,8 +59,10 @@ class AppRefService:
@staticmethod
def create_annotation_ref(app_ref: AppRef, annotation_id: str) -> AnnotationRef:
return AnnotationRef(tenant_id=app_ref.tenant_id, app_id=app_ref.app_id, annotation_id=annotation_id)
"""Bind a candidate annotation ID; ownership is enforced when the ref is consumed."""
return AnnotationRef(app=app_ref, annotation_id=annotation_id)
@staticmethod
def create_mcp_server_ref(app_ref: AppRef, server_id: str) -> AppMCPServerRef:
return AppMCPServerRef(tenant_id=app_ref.tenant_id, app_id=app_ref.app_id, server_id=server_id)
"""Bind a candidate MCP server ID; ownership is enforced when the ref is consumed."""
return AppMCPServerRef(app=app_ref, server_id=server_id)
+4 -1
View File
@@ -37,7 +37,10 @@ logger = logging.getLogger(__name__)
class AudioService:
@staticmethod
def _get_message_by_ref(session: Session, message_ref: MessageRef) -> Message | None:
stmt = select(Message).where(Message.id == message_ref.message_id, Message.app_id == message_ref.app_id)
stmt = select(Message).where(
Message.id == message_ref.message_id,
Message.app_id == message_ref.app.app_id,
)
if message_ref.end_user_id is not None:
stmt = stmt.where(Message.from_end_user_id == message_ref.end_user_id)
if message_ref.account_id is not None:
+26 -17
View File
@@ -2,6 +2,9 @@
from typing import NamedTuple
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.dataset import Dataset, Document
@@ -13,44 +16,50 @@ class DatasetRef(NamedTuple):
class DocumentRef(NamedTuple):
"""Document identifiers used to scope downstream resource lookups."""
"""Owner-bound lookup coordinates, not proof that a document is authorized or exists."""
tenant_id: str
dataset_id: str
dataset: DatasetRef
document_id: str
class SegmentRef(NamedTuple):
"""Segment identifiers used to scope downstream resource lookups."""
tenant_id: str
dataset_id: str
document_id: str
document: DocumentRef
segment_id: str
class DatasetRefService:
"""Factory helpers for dataset, document, and segment refs."""
"""Build child locators from validated dataset roots and resolve them with owner predicates."""
@staticmethod
def create_dataset_ref(dataset: Dataset) -> DatasetRef:
"""Create a root ref from a dataset already validated by the caller."""
return DatasetRef(tenant_id=dataset.tenant_id, dataset_id=dataset.id)
@staticmethod
def create_document_ref(dataset_ref: DatasetRef, document: Document) -> DocumentRef | None:
if document.tenant_id != dataset_ref.tenant_id or document.dataset_id != dataset_ref.dataset_id:
return None
return DocumentRef(
tenant_id=dataset_ref.tenant_id,
dataset_id=dataset_ref.dataset_id,
document_id=document.id,
)
return DatasetRefService.create_document_ref_from_id(dataset_ref, document.id)
@staticmethod
def create_document_ref_from_id(dataset_ref: DatasetRef, document_id: str) -> DocumentRef:
"""Bind a candidate document ID; ownership is enforced when the ref is consumed."""
return DocumentRef(dataset=dataset_ref, document_id=document_id)
@staticmethod
def create_segment_ref(document_ref: DocumentRef, segment_id: str) -> SegmentRef:
return SegmentRef(
tenant_id=document_ref.tenant_id,
dataset_id=document_ref.dataset_id,
document_id=document_ref.document_id,
segment_id=segment_id,
"""Bind a candidate segment ID; ownership is enforced when the ref is consumed."""
return SegmentRef(document=document_ref, segment_id=segment_id)
@staticmethod
def get_document_by_ref(document_ref: DocumentRef, *, session: Session) -> Document | None:
"""Resolve a document through its complete tenant and dataset ownership chain."""
return session.scalar(
select(Document).where(
Document.id == document_ref.document_id,
Document.dataset_id == document_ref.dataset.dataset_id,
Document.tenant_id == document_ref.dataset.tenant_id,
)
)
+6 -6
View File
@@ -4163,9 +4163,9 @@ class SegmentService:
select(ChildChunk)
.where(
ChildChunk.id == child_chunk_id,
ChildChunk.tenant_id == segment_ref.tenant_id,
ChildChunk.dataset_id == segment_ref.dataset_id,
ChildChunk.document_id == segment_ref.document_id,
ChildChunk.tenant_id == segment_ref.document.dataset.tenant_id,
ChildChunk.dataset_id == segment_ref.document.dataset.dataset_id,
ChildChunk.document_id == segment_ref.document.document_id,
ChildChunk.segment_id == segment_ref.segment_id,
)
.limit(1)
@@ -4219,9 +4219,9 @@ class SegmentService:
select(DocumentSegment)
.where(
DocumentSegment.id == segment_ref.segment_id,
DocumentSegment.tenant_id == segment_ref.tenant_id,
DocumentSegment.dataset_id == segment_ref.dataset_id,
DocumentSegment.document_id == segment_ref.document_id,
DocumentSegment.tenant_id == segment_ref.document.dataset.tenant_id,
DocumentSegment.dataset_id == segment_ref.document.dataset.dataset_id,
DocumentSegment.document_id == segment_ref.document.document_id,
)
.limit(1)
)
@@ -6,10 +6,11 @@ from sqlalchemy.orm import Session
from configs import dify_config
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
from core.app.entities.app_invoke_entities import InvokeFrom
from models.dataset import Document, Pipeline
from models.dataset import Pipeline
from models.enums import IndexingStatus
from models.model import Account, App, EndUser
from models.workflow import Workflow
from services.dataset_ref_service import DatasetRefService, DocumentRef
from services.rag_pipeline.rag_pipeline import RagPipelineService
@@ -37,8 +38,12 @@ class PipelineGenerateService:
try:
workflow = cls._get_workflow(pipeline, invoke_from, session)
if original_document_id := args.get("original_document_id"):
# update document status to waiting
cls.update_document_status(original_document_id, session=session)
dataset = pipeline.retrieve_dataset(session)
if dataset is None or dataset.tenant_id != pipeline.tenant_id:
raise ValueError("Pipeline dataset is required")
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, original_document_id)
cls.update_document_status(document_ref, session=session)
return PipelineGenerator.convert_to_event_stream(
PipelineGenerator().generate(
session=session,
@@ -123,12 +128,10 @@ class PipelineGenerateService:
return workflow
@classmethod
def update_document_status(cls, document_id: str, *, session: Session):
"""
Update document status to waiting
:param document_id: document id
"""
document = session.get(Document, document_id)
if document:
document.indexing_status = IndexingStatus.WAITING
session.add(document)
def update_document_status(cls, document_ref: DocumentRef, *, session: Session) -> None:
"""Set a document in the owner-bound dataset to waiting."""
document = DatasetRefService.get_document_by_ref(document_ref, session=session)
if document is None:
raise ValueError("Pipeline document not found")
document.indexing_status = IndexingStatus.WAITING
session.add(document)
+13 -11
View File
@@ -73,6 +73,7 @@ from models.workflow import (
WorkflowType,
)
from repositories.factory import DifyAPIRepositoryFactory
from services.dataset_ref_service import DatasetRefService
from services.datasource_provider_service import DatasourceProviderService
from services.entities.knowledge_entities.rag_pipeline_entities import (
KnowledgeConfiguration,
@@ -1013,23 +1014,24 @@ class RagPipelineService:
dataset_id = get_system_segment(variable_pool, SystemVariableKey.DATASET_ID)
pipeline_id = get_system_segment(variable_pool, SystemVariableKey.APP_ID)
if document_id and dataset_id and pipeline_id:
document = self._session.scalar(
select(Document)
.join(Dataset, Dataset.id == Document.dataset_id)
dataset = self._session.scalar(
select(Dataset)
.where(
Document.id == document_id.value,
Document.tenant_id == tenant_id,
Document.dataset_id == dataset_id.value,
Dataset.id == dataset_id.value,
Dataset.tenant_id == tenant_id,
Dataset.pipeline_id == pipeline_id.value,
)
.limit(1)
)
if document:
document.indexing_status = IndexingStatus.ERROR
document.error = error
self._session.add(document)
self._session.commit()
if dataset:
dataset_ref = DatasetRefService.create_dataset_ref(dataset)
document_ref = DatasetRefService.create_document_ref_from_id(dataset_ref, document_id.value)
document = DatasetRefService.get_document_by_ref(document_ref, session=self._session)
if document:
document.indexing_status = IndexingStatus.ERROR
document.error = error
self._session.add(document)
self._session.commit()
return workflow_node_execution