mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 10:38:32 +08:00
test: use sqlite3 session in test_retrieval_attachment_access (#38766)
This commit is contained in:
parent
0cea67e72d
commit
213b28f17b
@ -1,11 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
from datetime import UTC, datetime
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
|
||||
from core.app.file_access import (
|
||||
@ -18,46 +19,16 @@ from core.app.file_access import (
|
||||
is_retriever_segment_access_granted,
|
||||
)
|
||||
from core.rag.datasource.retrieval_service import RetrievalService
|
||||
from core.rag.embedding.retrieval import RetrievalSegments
|
||||
from core.rag.models.document import Document as RagDocument
|
||||
from core.rag.retrieval import dataset_retrieval as dataset_retrieval_module
|
||||
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
|
||||
from core.workflow.nodes.knowledge_retrieval.retrieval import KnowledgeRetrievalRequest
|
||||
from extensions.storage.storage_type import StorageType
|
||||
from models import UploadFile
|
||||
|
||||
|
||||
class _ScalarResult:
|
||||
def __init__(self, values):
|
||||
self._values = values
|
||||
|
||||
def all(self):
|
||||
return self._values
|
||||
|
||||
|
||||
class _AttachmentSession:
|
||||
def __init__(self, upload_file, binding):
|
||||
self._results = [
|
||||
_ScalarResult([upload_file]),
|
||||
_ScalarResult([binding]),
|
||||
]
|
||||
|
||||
def scalars(self, _stmt):
|
||||
return self._results.pop(0)
|
||||
|
||||
|
||||
class _DatasetRetrievalSession:
|
||||
def __init__(self, datasets, documents):
|
||||
self._results = [
|
||||
_ScalarResult(datasets),
|
||||
_ScalarResult(documents),
|
||||
]
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def scalars(self, _stmt):
|
||||
return self._results.pop(0)
|
||||
from models.dataset import Dataset, DocumentSegment, SegmentAttachmentBinding
|
||||
from models.dataset import Document as DatasetDocument
|
||||
from models.enums import CreatorUserRole, DataSourceType, DocumentCreatedFrom, SegmentStatus
|
||||
|
||||
|
||||
def test_file_access_grants_ignore_empty_inputs_and_missing_scope() -> None:
|
||||
@ -82,28 +53,44 @@ def test_file_access_grants_ignore_empty_inputs_and_missing_scope() -> None:
|
||||
assert current_scope.granted_retriever_segment_ids == frozenset()
|
||||
|
||||
|
||||
def test_segment_attachment_lookup_grants_returned_upload_files_to_current_scope() -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(UploadFile, SegmentAttachmentBinding)], indirect=True)
|
||||
def test_segment_attachment_lookup_grants_returned_upload_files_to_current_scope(sqlite_session: Session) -> None:
|
||||
tenant_id = str(uuid4())
|
||||
upload_file_id = str(uuid4())
|
||||
segment_id = str(uuid4())
|
||||
upload_file = SimpleNamespace(
|
||||
id=upload_file_id,
|
||||
user_id = str(uuid4())
|
||||
upload_file = UploadFile(
|
||||
tenant_id=tenant_id,
|
||||
storage_type=StorageType.LOCAL,
|
||||
key="uploads/chart.png",
|
||||
name="chart.png",
|
||||
size=1024,
|
||||
extension="png",
|
||||
mime_type="image/png",
|
||||
size=1024,
|
||||
created_by_role=CreatorUserRole.END_USER,
|
||||
created_by=user_id,
|
||||
created_at=datetime.now(UTC),
|
||||
used=False,
|
||||
)
|
||||
binding = SimpleNamespace(attachment_id=upload_file_id, segment_id=segment_id)
|
||||
session = _AttachmentSession(upload_file, binding)
|
||||
upload_file.id = upload_file_id
|
||||
binding = SegmentAttachmentBinding(
|
||||
tenant_id=tenant_id,
|
||||
dataset_id=str(uuid4()),
|
||||
document_id=str(uuid4()),
|
||||
segment_id=segment_id,
|
||||
attachment_id=upload_file_id,
|
||||
)
|
||||
sqlite_session.add_all([upload_file, binding])
|
||||
sqlite_session.commit()
|
||||
scope = FileAccessScope(
|
||||
tenant_id=tenant_id,
|
||||
user_id=str(uuid4()),
|
||||
user_id=user_id,
|
||||
user_from=UserFrom.END_USER,
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
)
|
||||
|
||||
with bind_file_access_scope(scope):
|
||||
result = RetrievalService.get_segment_attachment_infos([upload_file_id], session) # type: ignore[arg-type]
|
||||
result = RetrievalService.get_segment_attachment_infos([upload_file_id], sqlite_session)
|
||||
scoped_stmt = DatabaseFileAccessController().apply_upload_file_filters(
|
||||
select(UploadFile).where(UploadFile.id == upload_file_id)
|
||||
)
|
||||
@ -114,32 +101,60 @@ def test_segment_attachment_lookup_grants_returned_upload_files_to_current_scope
|
||||
assert "upload_files.id IN" in whereclause
|
||||
|
||||
|
||||
def test_knowledge_retrieval_grants_returned_segments_to_current_scope(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(Dataset, DatasetDocument, DocumentSegment)], indirect=True)
|
||||
def test_knowledge_retrieval_grants_returned_segments_to_current_scope(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_engine: Engine,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
tenant_id = str(uuid4())
|
||||
dataset_id = str(uuid4())
|
||||
document_id = str(uuid4())
|
||||
segment_id = str(uuid4())
|
||||
segment = SimpleNamespace(
|
||||
id=segment_id,
|
||||
created_by = str(uuid4())
|
||||
dataset = Dataset(tenant_id=tenant_id, name="Dataset", created_by=created_by)
|
||||
dataset.id = dataset_id
|
||||
document = DatasetDocument(
|
||||
tenant_id=tenant_id,
|
||||
dataset_id=dataset_id,
|
||||
position=1,
|
||||
data_source_type=DataSourceType.UPLOAD_FILE,
|
||||
batch="batch-1",
|
||||
name="Document",
|
||||
created_from=DocumentCreatedFrom.API,
|
||||
created_by=created_by,
|
||||
doc_metadata={},
|
||||
)
|
||||
document.id = document_id
|
||||
segment = DocumentSegment(
|
||||
tenant_id=tenant_id,
|
||||
dataset_id=dataset_id,
|
||||
document_id=document_id,
|
||||
hit_count=1,
|
||||
word_count=10,
|
||||
position=1,
|
||||
content="segment content",
|
||||
word_count=10,
|
||||
tokens=2,
|
||||
created_by=created_by,
|
||||
index_node_hash="hash",
|
||||
answer=None,
|
||||
get_sign_content=lambda: "segment content",
|
||||
hit_count=1,
|
||||
)
|
||||
record = SimpleNamespace(segment=segment, score=0.8, child_chunks=None, files=None, summary=None)
|
||||
dataset = SimpleNamespace(id=dataset_id, name="Dataset")
|
||||
document = SimpleNamespace(id=document_id, name="Document", data_source_type="upload_file", doc_metadata={})
|
||||
segment.id = segment_id
|
||||
segment.status = SegmentStatus.COMPLETED
|
||||
sqlite_session.add_all([dataset, document, segment])
|
||||
sqlite_session.commit()
|
||||
|
||||
record = RetrievalSegments(segment=segment, score=0.8)
|
||||
retrieval = DatasetRetrieval()
|
||||
monkeypatch.setattr(retrieval, "_check_knowledge_rate_limit", lambda tenant_id: None)
|
||||
monkeypatch.setattr(retrieval, "_get_available_datasets", lambda tenant_id, dataset_ids: [dataset])
|
||||
monkeypatch.setattr(retrieval, "multiple_retrieve", lambda **kwargs: [SimpleNamespace(provider="dify")])
|
||||
monkeypatch.setattr(
|
||||
retrieval,
|
||||
"multiple_retrieve",
|
||||
lambda **kwargs: [RagDocument(page_content="segment content", provider="dify")],
|
||||
)
|
||||
monkeypatch.setattr(RetrievalService, "format_retrieval_documents", lambda documents: [record])
|
||||
session = _DatasetRetrievalSession([dataset], [document])
|
||||
monkeypatch.setattr(dataset_retrieval_module.session_factory, "create_session", lambda: session)
|
||||
factory = sessionmaker(bind=sqlite_engine, expire_on_commit=False)
|
||||
monkeypatch.setattr(dataset_retrieval_module.session_factory, "create_session", factory)
|
||||
scope = FileAccessScope(
|
||||
tenant_id=tenant_id,
|
||||
user_id=str(uuid4()),
|
||||
@ -149,7 +164,7 @@ def test_knowledge_retrieval_grants_returned_segments_to_current_scope(monkeypat
|
||||
|
||||
with bind_file_access_scope(scope):
|
||||
results = retrieval.knowledge_retrieval(
|
||||
MagicMock(),
|
||||
sqlite_session,
|
||||
KnowledgeRetrievalRequest(
|
||||
tenant_id=tenant_id,
|
||||
user_id=str(uuid4()),
|
||||
|
||||
Loading…
Reference in New Issue
Block a user