mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix(api): allow LLM nodes to access retrieved knowledge files (#36175)
This commit is contained in:
@@ -1,6 +1,13 @@
|
||||
from .controller import DatabaseFileAccessController
|
||||
from .protocols import FileAccessControllerProtocol
|
||||
from .scope import FileAccessScope, bind_file_access_scope, get_current_file_access_scope
|
||||
from .scope import (
|
||||
FileAccessScope,
|
||||
bind_file_access_scope,
|
||||
get_current_file_access_scope,
|
||||
grant_retriever_segment_access,
|
||||
grant_upload_file_access,
|
||||
is_retriever_segment_access_granted,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DatabaseFileAccessController",
|
||||
@@ -8,4 +15,7 @@ __all__ = [
|
||||
"FileAccessScope",
|
||||
"bind_file_access_scope",
|
||||
"get_current_file_access_scope",
|
||||
"grant_retriever_segment_access",
|
||||
"grant_upload_file_access",
|
||||
"is_retriever_segment_access_granted",
|
||||
]
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy import and_, or_, select
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy.sql import Select
|
||||
|
||||
@@ -18,7 +18,8 @@ class DatabaseFileAccessController(FileAccessControllerProtocol):
|
||||
|
||||
Tenant scoping remains mandatory. When the current execution belongs to an
|
||||
end user, the lookup is additionally constrained to that end user's file
|
||||
ownership markers.
|
||||
ownership markers, plus upload files explicitly granted by the current
|
||||
execution context.
|
||||
"""
|
||||
|
||||
_scope_getter: Callable[[], FileAccessScope | None]
|
||||
@@ -47,10 +48,19 @@ class DatabaseFileAccessController(FileAccessControllerProtocol):
|
||||
if not resolved_scope.requires_user_ownership:
|
||||
return scoped_stmt
|
||||
|
||||
return scoped_stmt.where(
|
||||
user_owned_filter = and_(
|
||||
UploadFile.created_by_role == CreatorUserRole.END_USER,
|
||||
UploadFile.created_by == resolved_scope.user_id,
|
||||
)
|
||||
if not resolved_scope.granted_upload_file_ids:
|
||||
return scoped_stmt.where(user_owned_filter)
|
||||
|
||||
return scoped_stmt.where(
|
||||
or_(
|
||||
user_owned_filter,
|
||||
UploadFile.id.in_(resolved_scope.granted_upload_file_ids),
|
||||
)
|
||||
)
|
||||
|
||||
def apply_tool_file_filters(
|
||||
self,
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator # Changed from Iterator
|
||||
from collections.abc import Generator, Iterable
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field, replace
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
|
||||
|
||||
@@ -15,12 +15,23 @@ _current_file_access_scope: ContextVar[FileAccessScope | None] = ContextVar(
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FileAccessScope:
|
||||
"""Request-scoped ownership context used by workflow-layer file lookups."""
|
||||
"""Request-scoped ownership context used by workflow-layer file lookups.
|
||||
|
||||
``granted_upload_file_ids`` is execution-local: callers may add upload files
|
||||
that were returned by trusted retrieval paths without changing persistent
|
||||
ownership markers.
|
||||
|
||||
``granted_retriever_segment_ids`` gates lazy attachment loading by segment
|
||||
ID, so user-provided context cannot make a later LLM node load arbitrary
|
||||
same-tenant knowledge attachments.
|
||||
"""
|
||||
|
||||
tenant_id: str
|
||||
user_id: str
|
||||
user_from: UserFrom
|
||||
invoke_from: InvokeFrom
|
||||
granted_upload_file_ids: frozenset[str] = field(default_factory=frozenset)
|
||||
granted_retriever_segment_ids: frozenset[str] = field(default_factory=frozenset)
|
||||
|
||||
@property
|
||||
def requires_user_ownership(self) -> bool:
|
||||
@@ -31,8 +42,49 @@ def get_current_file_access_scope() -> FileAccessScope | None:
|
||||
return _current_file_access_scope.get()
|
||||
|
||||
|
||||
def grant_upload_file_access(upload_file_ids: Iterable[str]) -> None:
|
||||
scope = _current_file_access_scope.get()
|
||||
if scope is None:
|
||||
return
|
||||
|
||||
granted_upload_file_ids = frozenset(str(file_id) for file_id in upload_file_ids if file_id)
|
||||
if not granted_upload_file_ids:
|
||||
return
|
||||
|
||||
_current_file_access_scope.set(
|
||||
replace(
|
||||
scope,
|
||||
granted_upload_file_ids=scope.granted_upload_file_ids | granted_upload_file_ids,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def grant_retriever_segment_access(segment_ids: Iterable[str]) -> None:
|
||||
scope = _current_file_access_scope.get()
|
||||
if scope is None:
|
||||
return
|
||||
|
||||
granted_segment_ids = frozenset(str(segment_id) for segment_id in segment_ids if segment_id)
|
||||
if not granted_segment_ids:
|
||||
return
|
||||
|
||||
_current_file_access_scope.set(
|
||||
replace(
|
||||
scope,
|
||||
granted_retriever_segment_ids=scope.granted_retriever_segment_ids | granted_segment_ids,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def is_retriever_segment_access_granted(segment_id: str) -> bool:
|
||||
scope = _current_file_access_scope.get()
|
||||
if scope is None or not scope.requires_user_ownership:
|
||||
return True
|
||||
return str(segment_id) in scope.granted_retriever_segment_ids
|
||||
|
||||
|
||||
@contextmanager
|
||||
def bind_file_access_scope(scope: FileAccessScope) -> Generator[None, None, None]: # Changed from Iterator[None]
|
||||
def bind_file_access_scope(scope: FileAccessScope) -> Generator[None, None, None]:
|
||||
token = _current_file_access_scope.set(scope)
|
||||
try:
|
||||
yield
|
||||
|
||||
@@ -9,6 +9,7 @@ from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session, load_only
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.file_access import grant_upload_file_access
|
||||
from core.db.session_factory import session_factory
|
||||
from core.model_manager import ModelManager
|
||||
from core.rag.data_post_processor.data_post_processor import DataPostProcessor, RerankingModelDict, WeightsDict
|
||||
@@ -890,6 +891,7 @@ class RetrievalService:
|
||||
.limit(1)
|
||||
)
|
||||
if attachment_binding:
|
||||
grant_upload_file_access([str(upload_file.id)])
|
||||
attachment_info: AttachmentInfoDict = {
|
||||
"id": upload_file.id,
|
||||
"name": upload_file.name,
|
||||
@@ -906,6 +908,7 @@ class RetrievalService:
|
||||
cls, attachment_ids: list[str], session: Session
|
||||
) -> list[SegmentAttachmentInfoResult]:
|
||||
attachment_infos: list[SegmentAttachmentInfoResult] = []
|
||||
granted_upload_file_ids: list[str] = []
|
||||
upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(attachment_ids))).all()
|
||||
if upload_files:
|
||||
upload_file_ids = [upload_file.id for upload_file in upload_files]
|
||||
@@ -926,6 +929,7 @@ class RetrievalService:
|
||||
"size": upload_file.size,
|
||||
}
|
||||
if attachment_binding:
|
||||
granted_upload_file_ids.append(str(upload_file.id))
|
||||
attachment_infos.append(
|
||||
{
|
||||
"attachment_id": attachment_binding.attachment_id,
|
||||
@@ -933,4 +937,5 @@ class RetrievalService:
|
||||
"segment_id": attachment_binding.segment_id,
|
||||
}
|
||||
)
|
||||
grant_upload_file_access(granted_upload_file_ids)
|
||||
return attachment_infos
|
||||
|
||||
@@ -19,6 +19,7 @@ from core.app.app_config.entities import (
|
||||
ModelConfig,
|
||||
)
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity
|
||||
from core.app.file_access import grant_retriever_segment_access, grant_upload_file_access
|
||||
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
|
||||
from core.db.session_factory import session_factory
|
||||
from core.entities.agent_entities import PlanningStrategy
|
||||
@@ -326,6 +327,7 @@ class DatasetRetrieval:
|
||||
if record.summary:
|
||||
source.summary = record.summary
|
||||
|
||||
grant_retriever_segment_access([str(segment.id)])
|
||||
retrieval_resource_list.append(source)
|
||||
|
||||
if retrieval_resource_list:
|
||||
@@ -515,6 +517,9 @@ class DatasetRetrieval:
|
||||
)
|
||||
).all()
|
||||
if attachments_with_bindings:
|
||||
grant_upload_file_access(
|
||||
str(upload_file.id) for _, upload_file in attachments_with_bindings
|
||||
)
|
||||
for _, upload_file in attachments_with_bindings:
|
||||
attachment_info = File(
|
||||
file_id=upload_file.id,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import importlib
|
||||
import pkgutil
|
||||
from collections.abc import Callable, Iterator, Mapping, MutableMapping
|
||||
from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, cast, final, override
|
||||
@@ -56,6 +56,7 @@ from graphon.nodes.http_request import build_http_request_config
|
||||
from graphon.nodes.llm.entities import LLMNodeData
|
||||
from graphon.nodes.parameter_extractor.entities import ParameterExtractorNodeData
|
||||
from graphon.nodes.question_classifier.entities import QuestionClassifierNodeData
|
||||
from graphon.variables.segments import ArrayObjectSegment
|
||||
from models.model import Conversation
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -496,13 +497,47 @@ class DifyNodeFactory(NodeFactory):
|
||||
if include_prompt_message_serializer:
|
||||
node_init_kwargs["prompt_message_serializer"] = self._prompt_message_serializer
|
||||
if include_retriever_attachment_loader:
|
||||
node_init_kwargs["retriever_attachment_loader"] = self._retriever_attachment_loader
|
||||
node_init_kwargs["retriever_attachment_loader"] = self._build_retriever_attachment_loader(
|
||||
cast(LLMNodeData, validated_node_data)
|
||||
)
|
||||
if include_jinja2_template_renderer:
|
||||
node_init_kwargs["jinja2_template_renderer"] = self._jinja2_template_renderer
|
||||
if validated_node_data.type == BuiltinNodeTypes.LLM:
|
||||
node_init_kwargs["default_query_selector"] = system_variable_selector(SystemVariableKey.QUERY)
|
||||
return node_init_kwargs
|
||||
|
||||
def _build_retriever_attachment_loader(self, node_data: LLMNodeData) -> DifyRetrieverAttachmentLoader:
|
||||
return DifyRetrieverAttachmentLoader(
|
||||
file_reference_factory=self._file_reference_factory,
|
||||
segment_access_checker=self._build_retriever_segment_access_checker(
|
||||
node_data.context.variable_selector if node_data.context.enabled else None
|
||||
),
|
||||
)
|
||||
|
||||
def _build_retriever_segment_access_checker(
|
||||
self,
|
||||
context_variable_selector: Sequence[str] | None,
|
||||
) -> Callable[[str], bool]:
|
||||
def checker(segment_id: str) -> bool:
|
||||
if not context_variable_selector:
|
||||
return False
|
||||
|
||||
context_value = self.graph_runtime_state.variable_pool.get(context_variable_selector)
|
||||
if not isinstance(context_value, ArrayObjectSegment):
|
||||
return False
|
||||
|
||||
for item in context_value.value:
|
||||
if not isinstance(item, Mapping):
|
||||
continue
|
||||
metadata = item.get("metadata")
|
||||
if not isinstance(metadata, Mapping):
|
||||
continue
|
||||
if metadata.get("_source") == "knowledge" and str(metadata.get("segment_id")) == str(segment_id):
|
||||
return True
|
||||
return False
|
||||
|
||||
return checker
|
||||
|
||||
def _build_model_instance_for_llm_node(self, node_data: LLMCompatibleNodeData) -> ModelInstance:
|
||||
node_data_model = node_data.model
|
||||
model_instance, _ = fetch_model_config(
|
||||
|
||||
@@ -8,7 +8,11 @@ from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
|
||||
from core.app.file_access import DatabaseFileAccessController
|
||||
from core.app.file_access import (
|
||||
DatabaseFileAccessController,
|
||||
grant_upload_file_access,
|
||||
is_retriever_segment_access_granted,
|
||||
)
|
||||
from core.callback_handler.workflow_tool_callback_handler import DifyWorkflowCallbackHandler
|
||||
from core.helper.trace_id_helper import ParentTraceContext
|
||||
from core.llm_generator.output_parser.errors import OutputParserError
|
||||
@@ -275,10 +279,23 @@ class DifyPromptMessageSerializer(PromptMessageSerializerProtocol):
|
||||
class DifyRetrieverAttachmentLoader(RetrieverAttachmentLoaderProtocol):
|
||||
"""Resolve retriever attachments through Dify persistence and return graph file references."""
|
||||
|
||||
def __init__(self, *, file_reference_factory: FileReferenceFactoryProtocol) -> None:
|
||||
_segment_access_checker: Callable[[str], bool] | None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
file_reference_factory: FileReferenceFactoryProtocol,
|
||||
segment_access_checker: Callable[[str], bool] | None = None,
|
||||
) -> None:
|
||||
self._file_reference_factory = file_reference_factory
|
||||
self._segment_access_checker = segment_access_checker
|
||||
|
||||
def load(self, *, segment_id: str) -> Sequence[File]:
|
||||
if not is_retriever_segment_access_granted(segment_id):
|
||||
return []
|
||||
if self._segment_access_checker is not None and not self._segment_access_checker(segment_id):
|
||||
return []
|
||||
|
||||
with Session(db.engine, expire_on_commit=False) as session:
|
||||
attachments_with_bindings = session.execute(
|
||||
select(SegmentAttachmentBinding, UploadFile)
|
||||
@@ -286,6 +303,7 @@ class DifyRetrieverAttachmentLoader(RetrieverAttachmentLoaderProtocol):
|
||||
.where(SegmentAttachmentBinding.segment_id == segment_id)
|
||||
).all()
|
||||
|
||||
grant_upload_file_access(str(upload_file.id) for _, upload_file in attachments_with_bindings)
|
||||
return [
|
||||
self._file_reference_factory.build_from_mapping(
|
||||
mapping={
|
||||
|
||||
Reference in New Issue
Block a user