mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 18:58:35 +08:00
refactor: replace db.paginate with plain SQLAlchemy pagination (#38280)
Co-authored-by: Asuka Minato <i@asukaminato.eu.org> Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
070aed81d9
commit
2e1ab194b7
@ -13,6 +13,7 @@ from core.rag.index_processor.constant.built_in_field import BuiltInField
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
|
||||
from core.rag.models.document import ChildDocument, Document
|
||||
from extensions.ext_database import db
|
||||
from libs.pagination import paginate_query
|
||||
from models.dataset import Dataset, DatasetCollectionBinding, DatasetMetadata, DatasetMetadataBinding, DocumentSegment
|
||||
from models.dataset import Document as DatasetDocument
|
||||
from models.enums import DatasetMetadataType, IndexingStatus, SegmentStatus
|
||||
@ -183,7 +184,7 @@ def migrate_knowledge_vector_database():
|
||||
.order_by(Dataset.created_at.desc())
|
||||
)
|
||||
|
||||
datasets = db.paginate(select=stmt, page=page, per_page=50, max_per_page=50, error_out=False)
|
||||
datasets = paginate_query(stmt, page=page, per_page=50, max_per_page=50)
|
||||
if not datasets.items:
|
||||
break
|
||||
except SQLAlchemyError:
|
||||
@ -409,7 +410,7 @@ def old_metadata_migration():
|
||||
.where(DatasetDocument.doc_metadata.is_not(None))
|
||||
.order_by(DatasetDocument.created_at.desc())
|
||||
)
|
||||
documents = db.paginate(select=stmt, page=page, per_page=50, max_per_page=50, error_out=False)
|
||||
documents = paginate_query(stmt, page=page, per_page=50, max_per_page=50)
|
||||
except SQLAlchemyError:
|
||||
raise
|
||||
if not documents:
|
||||
|
||||
@ -41,6 +41,7 @@ from fields.conversation_fields import (
|
||||
from libs.datetime_utils import naive_utc_now, parse_time_range
|
||||
from libs.helper import dump_response
|
||||
from libs.login import login_required
|
||||
from libs.pagination import paginate_query
|
||||
from models import Conversation, EndUser, Message, MessageAnnotation
|
||||
from models.account import Account
|
||||
from models.model import App, AppMode
|
||||
@ -156,7 +157,7 @@ class CompletionConversationApi(Resource):
|
||||
|
||||
query = query.order_by(Conversation.created_at.desc())
|
||||
|
||||
conversations = db.paginate(query, page=args.page, per_page=args.limit, error_out=False)
|
||||
conversations = paginate_query(query, page=args.page, per_page=args.limit)
|
||||
|
||||
return dump_response(ConversationPaginationResponse, conversations)
|
||||
|
||||
@ -310,7 +311,7 @@ class ChatConversationApi(Resource):
|
||||
case _:
|
||||
query = query.order_by(Conversation.created_at.desc())
|
||||
|
||||
conversations = db.paginate(query, page=args.page, per_page=args.limit, error_out=False)
|
||||
conversations = paginate_query(query, page=args.page, per_page=args.limit)
|
||||
|
||||
return dump_response(ConversationWithSummaryPaginationResponse, conversations)
|
||||
|
||||
|
||||
@ -46,6 +46,7 @@ from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.helper import dump_response, to_timestamp
|
||||
from libs.login import login_required
|
||||
from libs.pagination import paginate_query
|
||||
from models import Account, DatasetProcessRule, Document, DocumentSegment, UploadFile
|
||||
from models.dataset import DocumentPipelineExecutionLog
|
||||
from models.enums import IndexingStatus, SegmentStatus
|
||||
@ -368,7 +369,7 @@ class DatasetDocumentListApi(Resource):
|
||||
desc(Document.position),
|
||||
)
|
||||
|
||||
paginated_documents = db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False)
|
||||
paginated_documents = paginate_query(query, page=page, per_page=limit, max_per_page=100)
|
||||
documents = paginated_documents.items
|
||||
|
||||
DocumentService.enrich_documents_with_summary_index_status(
|
||||
|
||||
@ -57,6 +57,7 @@ from fields.segment_fields import (
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
from libs.helper import dump_response, escape_like_pattern
|
||||
from libs.login import login_required
|
||||
from libs.pagination import paginate_query
|
||||
from models import Account
|
||||
from models.dataset import Dataset, Document, DocumentSegment
|
||||
from models.model import UploadFile
|
||||
@ -270,7 +271,7 @@ class DatasetDocumentSegmentListApi(Resource):
|
||||
elif args.enabled.lower() == "false":
|
||||
query = query.where(DocumentSegment.enabled == False)
|
||||
|
||||
segments = db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False)
|
||||
segments = paginate_query(query, page=page, per_page=limit, max_per_page=100)
|
||||
|
||||
segment_list = list(segments.items)
|
||||
segment_ids = [segment.id for segment in segment_list]
|
||||
|
||||
@ -33,6 +33,7 @@ from extensions.ext_database import db
|
||||
from fields.base import ResponseModel
|
||||
from libs.helper import OptionalTimestampField, TimestampField, dump_response, to_timestamp
|
||||
from libs.login import login_required
|
||||
from libs.pagination import paginate_query
|
||||
from models.account import Account, Tenant, TenantAccountJoin, TenantCustomConfigDict, TenantStatus
|
||||
from services.account_service import TenantService
|
||||
from services.billing_service import BillingService, SubscriptionPlan
|
||||
@ -294,7 +295,7 @@ class WorkspaceListApi(Resource):
|
||||
args = WorkspaceListQuery.model_validate(payload)
|
||||
|
||||
stmt = select(Tenant).order_by(Tenant.created_at.desc())
|
||||
tenants = db.paginate(select=stmt, page=args.page, per_page=args.limit, error_out=False)
|
||||
tenants = paginate_query(stmt, page=args.page, per_page=args.limit)
|
||||
has_more = False
|
||||
|
||||
if tenants.has_next:
|
||||
|
||||
@ -59,6 +59,7 @@ from fields.document_fields import (
|
||||
)
|
||||
from libs.helper import dump_response
|
||||
from libs.login import current_user
|
||||
from libs.pagination import paginate_query
|
||||
from models.dataset import Dataset, Document, DocumentSegment
|
||||
from models.enums import SegmentStatus
|
||||
from services.dataset_service import DatasetService, DocumentService
|
||||
@ -945,8 +946,8 @@ class DocumentListApi(DatasetApiResource):
|
||||
|
||||
query = query.order_by(desc(Document.created_at), desc(Document.position))
|
||||
|
||||
paginated_documents = db.paginate(
|
||||
select=query, page=query_params.page, per_page=query_params.limit, max_per_page=100, error_out=False
|
||||
paginated_documents = paginate_query(
|
||||
query, page=query_params.page, per_page=query_params.limit, max_per_page=100
|
||||
)
|
||||
documents = paginated_documents.items
|
||||
|
||||
|
||||
86
api/libs/pagination.py
Normal file
86
api/libs/pagination.py
Normal file
@ -0,0 +1,86 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
from sqlalchemy import Select, func, select
|
||||
from sqlalchemy.orm import Session, scoped_session
|
||||
|
||||
|
||||
@dataclass
|
||||
class PaginatedResult[T]:
|
||||
"""Minimal pagination container backed by plain SQLAlchemy queries.
|
||||
|
||||
Drop-in replacement for Flask-SQLAlchemy's ``db.paginate`` return value.
|
||||
Only the attributes actually consumed across the codebase are exposed:
|
||||
``items``, ``total``, ``page``, ``per_page``, ``pages``, ``has_next``.
|
||||
"""
|
||||
|
||||
items: list[T]
|
||||
total: int
|
||||
page: int
|
||||
per_page: int
|
||||
|
||||
@property
|
||||
def pages(self) -> int:
|
||||
if self.per_page == 0:
|
||||
return 0
|
||||
return max(1, math.ceil(self.total / self.per_page))
|
||||
|
||||
@property
|
||||
def has_next(self) -> bool:
|
||||
return self.page < self.pages
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.items)
|
||||
|
||||
|
||||
def paginate_query(
|
||||
stmt: Select,
|
||||
*,
|
||||
page: int = 1,
|
||||
per_page: int = 20,
|
||||
max_per_page: int | None = None,
|
||||
session: Session | scoped_session | None = None,
|
||||
) -> PaginatedResult:
|
||||
"""Execute *stmt* as a paginated query using plain SQLAlchemy.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
stmt:
|
||||
A SQLAlchemy ``select()`` statement.
|
||||
page:
|
||||
1-based page number.
|
||||
per_page:
|
||||
Number of items per page.
|
||||
max_per_page:
|
||||
Hard ceiling for *per_page*; ``None`` means no cap.
|
||||
session:
|
||||
The session to use. Falls back to ``db.session`` when omitted.
|
||||
"""
|
||||
if session is None:
|
||||
from extensions.ext_database import db
|
||||
|
||||
session = db.session
|
||||
|
||||
if max_per_page is not None:
|
||||
per_page = min(per_page, max_per_page)
|
||||
|
||||
page = max(1, page)
|
||||
per_page = max(1, per_page)
|
||||
|
||||
# total count — wrap in a scalar subquery so arbitrary selects work
|
||||
count_stmt = select(func.count()).select_from(stmt.subquery())
|
||||
total: int = session.scalar(count_stmt) or 0 # type: ignore[assignment]
|
||||
|
||||
# fetch the page
|
||||
offset = (page - 1) * per_page
|
||||
page_stmt = stmt.limit(per_page).offset(offset)
|
||||
items = list(session.scalars(page_stmt).all())
|
||||
|
||||
return PaginatedResult(
|
||||
items=items,
|
||||
total=total,
|
||||
page=page,
|
||||
per_page=per_page,
|
||||
)
|
||||
@ -12,6 +12,7 @@ from core.rag.index_processor.index_processor_factory import IndexProcessorFacto
|
||||
from enums.cloud_plan import CloudPlan
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from libs.pagination import paginate_query
|
||||
from models.dataset import Dataset, DatasetAutoDisableLog, DatasetQuery, Document
|
||||
from services.feature_service import FeatureService
|
||||
|
||||
@ -88,7 +89,7 @@ def clean_unused_datasets_task():
|
||||
.order_by(Dataset.created_at.desc())
|
||||
)
|
||||
|
||||
datasets = db.paginate(stmt, page=page, per_page=50, error_out=False)
|
||||
datasets = paginate_query(stmt, page=page, per_page=50)
|
||||
|
||||
except SQLAlchemyError:
|
||||
raise
|
||||
|
||||
@ -13,6 +13,7 @@ from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.login import current_account_with_tenant
|
||||
from libs.pagination import paginate_query
|
||||
from models.model import App, AppAnnotationHitHistory, AppAnnotationSetting, Message, MessageAnnotation
|
||||
from services.app_ref_service import AnnotationRef, AppRef
|
||||
from services.feature_service import FeatureService
|
||||
@ -242,7 +243,7 @@ class AppAnnotationService:
|
||||
.where(MessageAnnotation.app_id == app_id)
|
||||
.order_by(MessageAnnotation.created_at.desc(), MessageAnnotation.id.desc())
|
||||
)
|
||||
annotations = db.paginate(select=stmt, page=page, per_page=limit, max_per_page=100, error_out=False)
|
||||
annotations = paginate_query(stmt, page=page, per_page=limit, max_per_page=100)
|
||||
return annotations.items, annotations.total or 0
|
||||
|
||||
@classmethod
|
||||
@ -573,9 +574,7 @@ class AppAnnotationService:
|
||||
)
|
||||
.order_by(AppAnnotationHitHistory.created_at.desc())
|
||||
)
|
||||
annotation_hit_histories = db.paginate(
|
||||
select=stmt, page=page, per_page=limit, max_per_page=100, error_out=False
|
||||
)
|
||||
annotation_hit_histories = paginate_query(stmt, page=page, per_page=limit, max_per_page=100)
|
||||
return annotation_hit_histories.items, annotation_hit_histories.total or 0
|
||||
|
||||
@classmethod
|
||||
|
||||
@ -5,7 +5,6 @@ from datetime import datetime
|
||||
from typing import Any, Literal, NotRequired, TypedDict, cast, override
|
||||
|
||||
import sqlalchemy as sa
|
||||
from flask_sqlalchemy.pagination import Pagination
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy import ColumnElement, select
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
@ -24,6 +23,7 @@ from graphon.model_runtime.entities.model_entities import ModelPropertyKey, Mode
|
||||
from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.login import current_user
|
||||
from libs.pagination import PaginatedResult, paginate_query
|
||||
from models import Account, AppStar
|
||||
from models.agent import Agent, AgentIconType, AgentScope, AgentSource, AgentStatus
|
||||
from models.model import App, AppMode, AppModelConfig, IconType, Site
|
||||
@ -220,7 +220,7 @@ class AppService:
|
||||
|
||||
def get_paginate_apps(
|
||||
self, user_id: str, tenant_id: str, params: AppListParams, session: scoped_session
|
||||
) -> Pagination | None:
|
||||
) -> PaginatedResult | None:
|
||||
"""
|
||||
Get app list with pagination, filters, and explicit sort order.
|
||||
:param user_id: user id
|
||||
@ -234,11 +234,10 @@ class AppService:
|
||||
|
||||
order_by = self._build_app_list_order_by(params.sort_by)
|
||||
|
||||
app_models = db.paginate(
|
||||
app_models = paginate_query(
|
||||
sa.select(App).where(*filters).order_by(order_by),
|
||||
page=params.page,
|
||||
per_page=params.limit,
|
||||
error_out=False,
|
||||
)
|
||||
|
||||
app_ids = [str(app.id) for app in app_models.items]
|
||||
@ -255,7 +254,7 @@ class AppService:
|
||||
|
||||
def get_paginate_starred_apps(
|
||||
self, user_id: str, tenant_id: str, params: StarredAppListParams, session: scoped_session
|
||||
) -> Pagination | None:
|
||||
) -> PaginatedResult | None:
|
||||
"""
|
||||
Get apps starred by the current account with pagination, filters, and explicit sort order.
|
||||
"""
|
||||
@ -264,7 +263,7 @@ class AppService:
|
||||
return None
|
||||
|
||||
order_by = self._build_app_list_order_by(params.sort_by)
|
||||
app_models = db.paginate(
|
||||
app_models = paginate_query(
|
||||
sa.select(App)
|
||||
.join(
|
||||
AppStar,
|
||||
@ -278,7 +277,6 @@ class AppService:
|
||||
.order_by(order_by),
|
||||
page=params.page,
|
||||
per_page=params.limit,
|
||||
error_out=False,
|
||||
)
|
||||
|
||||
for app in app_models.items:
|
||||
|
||||
@ -34,6 +34,7 @@ from graphon.model_runtime.model_providers.base.text_embedding_model import Text
|
||||
from libs import helper
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.login import current_user
|
||||
from libs.pagination import paginate_query
|
||||
from models import Account, TenantAccountRole
|
||||
from models.dataset import (
|
||||
AppDatasetJoin,
|
||||
@ -355,7 +356,7 @@ class DatasetService:
|
||||
else:
|
||||
return [], 0
|
||||
|
||||
datasets = db.paginate(select=query, page=page, per_page=per_page, max_per_page=100, error_out=False)
|
||||
datasets = paginate_query(query, page=page, per_page=per_page, max_per_page=100)
|
||||
|
||||
return datasets.items, datasets.total
|
||||
|
||||
@ -399,7 +400,7 @@ class DatasetService:
|
||||
accessible_filter = sa.or_(Dataset.maintainer == user.id, accessible_filter)
|
||||
stmt = stmt.where(accessible_filter)
|
||||
|
||||
datasets = db.paginate(select=stmt, page=1, per_page=len(ids), max_per_page=len(ids), error_out=False)
|
||||
datasets = paginate_query(stmt, page=1, per_page=len(ids), max_per_page=len(ids))
|
||||
|
||||
return datasets.items, datasets.total
|
||||
|
||||
@ -1403,7 +1404,7 @@ class DatasetService:
|
||||
def get_dataset_queries(dataset_id: str, page: int, per_page: int):
|
||||
stmt = select(DatasetQuery).filter_by(dataset_id=dataset_id).order_by(db.desc(DatasetQuery.created_at))
|
||||
|
||||
dataset_queries = db.paginate(select=stmt, page=page, per_page=per_page, max_per_page=100, error_out=False)
|
||||
dataset_queries = paginate_query(stmt, page=page, per_page=per_page, max_per_page=100)
|
||||
|
||||
return dataset_queries.items, dataset_queries.total
|
||||
|
||||
@ -4120,7 +4121,7 @@ class SegmentService:
|
||||
if keyword:
|
||||
escaped_keyword = helper.escape_like_pattern(keyword)
|
||||
query = query.where(ChildChunk.content.ilike(f"%{escaped_keyword}%", escape="\\"))
|
||||
return db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False)
|
||||
return paginate_query(query, page=page, per_page=limit, max_per_page=100)
|
||||
|
||||
@classmethod
|
||||
def get_child_chunk_by_id(
|
||||
@ -4172,7 +4173,7 @@ class SegmentService:
|
||||
query = query.where(DocumentSegment.content.ilike(f"%{escaped_keyword}%", escape="\\"))
|
||||
|
||||
query = query.order_by(DocumentSegment.position.asc(), DocumentSegment.id.asc())
|
||||
paginated_segments = db.paginate(select=query, page=page, per_page=limit, max_per_page=100, error_out=False)
|
||||
paginated_segments = paginate_query(query, page=page, per_page=limit, max_per_page=100)
|
||||
|
||||
return paginated_segments.items, paginated_segments.total
|
||||
|
||||
|
||||
@ -10,9 +10,9 @@ from sqlalchemy.orm import Session
|
||||
from constants import HIDDEN_VALUE
|
||||
from core.helper import ssrf_proxy
|
||||
from core.rag.entities import MetadataFilteringCondition
|
||||
from extensions.ext_database import db
|
||||
from graphon.nodes.http_request.exc import InvalidHttpMethodError
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from libs.pagination import paginate_query
|
||||
from models.dataset import (
|
||||
Dataset,
|
||||
ExternalKnowledgeApis,
|
||||
@ -42,9 +42,7 @@ class ExternalDatasetService:
|
||||
escaped_search = escape_like_pattern(search)
|
||||
query = query.where(ExternalKnowledgeApis.name.ilike(f"%{escaped_search}%", escape="\\"))
|
||||
|
||||
external_knowledge_apis = db.paginate(
|
||||
select=query, page=page, per_page=per_page, max_per_page=100, error_out=False
|
||||
)
|
||||
external_knowledge_apis = paginate_query(query, page=page, per_page=per_page, max_per_page=100)
|
||||
|
||||
return external_knowledge_apis.items, external_knowledge_apis.total
|
||||
|
||||
|
||||
@ -25,6 +25,7 @@ from core.plugin.impl.plugin import PluginInstaller
|
||||
from core.plugin.plugin_service import PluginService
|
||||
from core.tools.entities.tool_entities import ToolProviderType
|
||||
from extensions.ext_database import db
|
||||
from libs.pagination import paginate_query
|
||||
from models.account import Tenant
|
||||
from models.model import App, AppMode, AppModelConfig
|
||||
from models.provider_ids import ModelProviderID, ToolProviderID
|
||||
@ -499,7 +500,7 @@ class PluginMigration:
|
||||
total_failed_tenant = 0
|
||||
while True:
|
||||
# paginate
|
||||
tenants = db.paginate(sa.select(Tenant).order_by(Tenant.created_at.desc()), page=page, per_page=100)
|
||||
tenants = paginate_query(sa.select(Tenant).order_by(Tenant.created_at.desc()), page=page, per_page=100)
|
||||
if tenants.items is None or len(tenants.items) == 0:
|
||||
break
|
||||
|
||||
|
||||
@ -30,7 +30,7 @@ def test_completion_conversation_list_returns_paginated_result(app: Flask, monke
|
||||
paginate_result.total = 0
|
||||
paginate_result.has_next = False
|
||||
paginate_result.items = []
|
||||
monkeypatch.setattr(conversation_module.db, "paginate", lambda *_args, **_kwargs: paginate_result)
|
||||
monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/completion-conversations", method="GET"):
|
||||
response = method(api, account, app_model=SimpleNamespace(id="app-1"))
|
||||
@ -71,7 +71,7 @@ def test_chat_conversation_list_advanced_chat_calls_paginate(app: Flask, monkeyp
|
||||
paginate_result.total = 0
|
||||
paginate_result.has_next = False
|
||||
paginate_result.items = []
|
||||
monkeypatch.setattr(conversation_module.db, "paginate", lambda *_args, **_kwargs: paginate_result)
|
||||
monkeypatch.setattr(conversation_module, "paginate_query", lambda *_args, **_kwargs: paginate_result)
|
||||
|
||||
with app.test_request_context("/console/api/apps/app-1/chat-conversations", method="GET"):
|
||||
response = method(api, account, app_model=SimpleNamespace(id="app-1", mode=AppMode.ADVANCED_CHAT))
|
||||
|
||||
@ -207,7 +207,7 @@ class TestDatasetDocumentListApi:
|
||||
with (
|
||||
app.test_request_context("/?fetch=true"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.db.paginate",
|
||||
"controllers.console.datasets.datasets_document.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -237,7 +237,7 @@ class TestDatasetDocumentListApi:
|
||||
with (
|
||||
app.test_request_context("/?keyword=test&status=enabled&sort=created_at"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.db.paginate",
|
||||
"controllers.console.datasets.datasets_document.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -263,7 +263,7 @@ class TestDatasetDocumentListApi:
|
||||
with (
|
||||
app.test_request_context("/"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.db.paginate",
|
||||
"controllers.console.datasets.datasets_document.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -341,7 +341,7 @@ class TestDatasetDocumentListApi:
|
||||
with (
|
||||
app.test_request_context("/?fetch=maybe"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.db.paginate",
|
||||
"controllers.console.datasets.datasets_document.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -363,7 +363,7 @@ class TestDatasetDocumentListApi:
|
||||
with (
|
||||
app.test_request_context("/?sort=hit_count"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.db.paginate",
|
||||
"controllers.console.datasets.datasets_document.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -1537,7 +1537,7 @@ class TestDocumentListAdvancedCases:
|
||||
with (
|
||||
app.test_request_context("/?sort=updated_at"),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_document.db.paginate",
|
||||
"controllers.console.datasets.datasets_document.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
|
||||
@ -161,7 +161,7 @@ class TestDatasetDocumentSegmentListApi:
|
||||
return_value=document,
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_segments.db.paginate",
|
||||
"controllers.console.datasets.datasets_segments.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -1207,7 +1207,7 @@ class TestSegmentListAdvancedCases:
|
||||
return_value=document,
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_segments.db.paginate",
|
||||
"controllers.console.datasets.datasets_segments.paginate_query",
|
||||
return_value=pagination,
|
||||
),
|
||||
patch(
|
||||
@ -1255,7 +1255,7 @@ class TestSegmentListAdvancedCases:
|
||||
SimpleNamespace(SQLALCHEMY_DATABASE_URI_SCHEME="postgresql"),
|
||||
),
|
||||
patch(
|
||||
"controllers.console.datasets.datasets_segments.db.paginate",
|
||||
"controllers.console.datasets.datasets_segments.paginate_query",
|
||||
return_value=pagination,
|
||||
) as paginate_mock,
|
||||
):
|
||||
@ -1267,7 +1267,7 @@ class TestSegmentListAdvancedCases:
|
||||
"33333333-3333-3333-3333-333333333333",
|
||||
)
|
||||
|
||||
query = paginate_mock.call_args.kwargs["select"]
|
||||
query = paginate_mock.call_args.args[0]
|
||||
sql = str(query.compile(compile_kwargs={"literal_binds": True}))
|
||||
assert "jsonb_array_elements_text(CASE" in sql
|
||||
assert "ELSE CAST('[]' AS JSONB)" in sql
|
||||
|
||||
@ -291,7 +291,7 @@ class TestWorkspaceListApi:
|
||||
|
||||
with (
|
||||
app.test_request_context("/all-workspaces", query_string={"page": 1, "limit": 20}),
|
||||
patch("controllers.console.workspace.workspace.db.paginate", return_value=paginate_result),
|
||||
patch("controllers.console.workspace.workspace.paginate_query", return_value=paginate_result),
|
||||
):
|
||||
result, status = method(api)
|
||||
|
||||
@ -308,7 +308,7 @@ class TestWorkspaceListApi:
|
||||
|
||||
with (
|
||||
app.test_request_context("/all-workspaces", query_string={"page": 1, "limit": 1}),
|
||||
patch("controllers.console.workspace.workspace.db.paginate", return_value=paginate_result),
|
||||
patch("controllers.console.workspace.workspace.paginate_query", return_value=paginate_result),
|
||||
):
|
||||
result, status = method(api)
|
||||
|
||||
|
||||
@ -851,9 +851,10 @@ class TestDocumentApiDelete:
|
||||
class TestDocumentListApi:
|
||||
"""Test suite for DocumentListApi endpoint."""
|
||||
|
||||
@patch("controllers.service_api.dataset.document.paginate_query")
|
||||
@patch("controllers.service_api.dataset.document.DocumentService")
|
||||
@patch("controllers.service_api.dataset.document.db")
|
||||
def test_list_documents_success(self, mock_db, mock_doc_svc, app: Flask, mock_tenant, mock_dataset):
|
||||
def test_list_documents_success(self, mock_db, mock_doc_svc, mock_paginate, app: Flask, mock_tenant, mock_dataset):
|
||||
"""Test successful document list retrieval."""
|
||||
# Arrange
|
||||
mock_db.session.scalar.return_value = mock_dataset
|
||||
@ -868,7 +869,7 @@ class TestDocumentListApi:
|
||||
make_serializable_document(id="doc-2", name="Document 2"),
|
||||
]
|
||||
mock_pagination.total = 2
|
||||
mock_db.paginate.return_value = mock_pagination
|
||||
mock_paginate.return_value = mock_pagination
|
||||
|
||||
mock_doc_svc.enrich_documents_with_summary_index_status.return_value = None
|
||||
|
||||
|
||||
@ -395,10 +395,11 @@ class TestAppAnnotationServiceListAndExport:
|
||||
with (
|
||||
patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)),
|
||||
patch("services.annotation_service.db") as mock_db,
|
||||
patch("services.annotation_service.paginate_query") as mock_paginate,
|
||||
patch("libs.helper.escape_like_pattern", return_value="safe"),
|
||||
):
|
||||
mock_db.session.scalar.return_value = app
|
||||
mock_db.paginate.return_value = pagination
|
||||
mock_paginate.return_value = pagination
|
||||
|
||||
# Act
|
||||
items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "keyword")
|
||||
@ -417,9 +418,10 @@ class TestAppAnnotationServiceListAndExport:
|
||||
with (
|
||||
patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)),
|
||||
patch("services.annotation_service.db") as mock_db,
|
||||
patch("services.annotation_service.paginate_query") as mock_paginate,
|
||||
):
|
||||
mock_db.session.scalar.return_value = app
|
||||
mock_db.paginate.return_value = pagination
|
||||
mock_paginate.return_value = pagination
|
||||
|
||||
# Act
|
||||
items, total = AppAnnotationService.get_annotation_list_by_app_id(app.id, 1, 10, "")
|
||||
@ -1101,9 +1103,11 @@ class TestAppAnnotationServiceHitHistoryAndSettings:
|
||||
with (
|
||||
patch("services.annotation_service.current_account_with_tenant", return_value=(_make_user(), tenant_id)),
|
||||
patch("services.annotation_service.db") as mock_db,
|
||||
patch("services.annotation_service.paginate_query") as mock_paginate,
|
||||
):
|
||||
mock_db.session.scalar.return_value = annotation
|
||||
mock_db.paginate.return_value = pagination
|
||||
mock_db.session.scalar.return_value = app
|
||||
mock_db.session.get.return_value = annotation
|
||||
mock_paginate.return_value = pagination
|
||||
|
||||
# Act
|
||||
items, total = AppAnnotationService.get_annotation_hit_histories(
|
||||
|
||||
@ -174,18 +174,18 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
mock_db = MagicMock()
|
||||
explicit_session = MagicMock()
|
||||
explicit_session.scalars.return_value.all.return_value = []
|
||||
mock_db.paginate.return_value.items = []
|
||||
mock_db.paginate.return_value.total = 0
|
||||
user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL)
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db", mock_db),
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.dify_config.RBAC_ENABLED", True),
|
||||
patch(
|
||||
"services.dataset_service.enterprise_rbac_service.RBACService.MyPermissions.get",
|
||||
return_value=SimpleNamespace(workspace=SimpleNamespace(permission_keys=[])),
|
||||
),
|
||||
):
|
||||
mock_paginate.return_value = SimpleNamespace(items=[], total=0)
|
||||
DatasetService.get_datasets(
|
||||
page=1,
|
||||
per_page=20,
|
||||
@ -198,7 +198,7 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
|
||||
explicit_session.scalars.assert_called_once()
|
||||
mock_db.session.scalars.assert_not_called()
|
||||
select_stmt = mock_db.paginate.call_args.kwargs["select"]
|
||||
select_stmt = mock_paginate.call_args.args[0]
|
||||
visibility_clause = str(select_stmt._where_criteria[1])
|
||||
assert "maintainer" in visibility_clause
|
||||
assert "IN" in visibility_clause
|
||||
@ -206,18 +206,18 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
def test_get_datasets_filters_only_by_rbac_overrides_without_manage_own_permission(self):
|
||||
mock_db = MagicMock()
|
||||
mock_db.session.scalars.return_value.all.return_value = []
|
||||
mock_db.paginate.return_value.items = []
|
||||
mock_db.paginate.return_value.total = 0
|
||||
user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL)
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db", mock_db),
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.dify_config.RBAC_ENABLED", True),
|
||||
patch(
|
||||
"services.dataset_service.enterprise_rbac_service.RBACService.MyPermissions.get",
|
||||
return_value=SimpleNamespace(workspace=SimpleNamespace(permission_keys=[])),
|
||||
),
|
||||
):
|
||||
mock_paginate.return_value = SimpleNamespace(items=[], total=0)
|
||||
DatasetService.get_datasets(
|
||||
page=1,
|
||||
per_page=20,
|
||||
@ -227,21 +227,21 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
accessible_dataset_ids=["dataset-shared"],
|
||||
)
|
||||
|
||||
select_stmt = mock_db.paginate.call_args.kwargs["select"]
|
||||
select_stmt = mock_paginate.call_args.args[0]
|
||||
visibility_clause = str(select_stmt._where_criteria[1])
|
||||
assert "maintainer" not in visibility_clause
|
||||
assert "IN" in visibility_clause
|
||||
|
||||
def test_get_datasets_by_ids_applies_rbac_visibility(self):
|
||||
mock_db = MagicMock()
|
||||
mock_db.paginate.return_value.items = []
|
||||
mock_db.paginate.return_value.total = 0
|
||||
user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL)
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db", mock_db),
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.dify_config.RBAC_ENABLED", True),
|
||||
):
|
||||
mock_paginate.return_value = SimpleNamespace(items=[], total=0)
|
||||
DatasetService.get_datasets_by_ids(
|
||||
["dataset-requested", "dataset-shared"],
|
||||
"tenant-1",
|
||||
@ -250,7 +250,7 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
include_own_datasets=True,
|
||||
)
|
||||
|
||||
select_stmt = mock_db.paginate.call_args.kwargs["select"]
|
||||
select_stmt = mock_paginate.call_args.args[0]
|
||||
visibility_clause = str(select_stmt._where_criteria[-1])
|
||||
assert "maintainer" in visibility_clause
|
||||
assert "IN" in visibility_clause
|
||||
@ -262,20 +262,20 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
def test_get_datasets_rbac_include_all_uses_workspace_permission(self):
|
||||
mock_db = MagicMock()
|
||||
mock_db.session.scalars.return_value.all.return_value = []
|
||||
mock_db.paginate.return_value.items = []
|
||||
mock_db.paginate.return_value.total = 0
|
||||
|
||||
user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.NORMAL)
|
||||
mock_permissions = SimpleNamespace(workspace=SimpleNamespace(permission_keys=["dataset.create_and_management"]))
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db", mock_db),
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.dify_config.RBAC_ENABLED", True),
|
||||
patch(
|
||||
"services.dataset_service.enterprise_rbac_service.RBACService.MyPermissions.get",
|
||||
return_value=mock_permissions,
|
||||
),
|
||||
):
|
||||
mock_paginate.return_value = SimpleNamespace(items=[], total=0)
|
||||
DatasetService.get_datasets(
|
||||
page=1,
|
||||
per_page=20,
|
||||
@ -286,39 +286,39 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
)
|
||||
|
||||
mock_db.session.scalars.assert_called_once()
|
||||
mock_db.paginate.assert_called_once()
|
||||
select_stmt = mock_db.paginate.call_args.kwargs["select"]
|
||||
mock_paginate.assert_called_once()
|
||||
select_stmt = mock_paginate.call_args.args[0]
|
||||
assert len(select_stmt._where_criteria) == 1
|
||||
|
||||
def test_get_datasets_rbac_without_user_returns_empty_result(self):
|
||||
mock_db = MagicMock()
|
||||
mock_db.session.scalars.return_value.all.return_value = []
|
||||
mock_db.paginate.return_value.items = []
|
||||
mock_db.paginate.return_value.total = 0
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db", mock_db),
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.dify_config.RBAC_ENABLED", True),
|
||||
):
|
||||
mock_paginate.return_value = SimpleNamespace(items=[], total=0)
|
||||
DatasetService.get_datasets(page=1, per_page=20, session=mock_db.session, tenant_id="tenant-1", user=None)
|
||||
|
||||
mock_db.session.scalars.assert_not_called()
|
||||
mock_db.paginate.assert_called_once()
|
||||
select_stmt = mock_db.paginate.call_args.kwargs["select"]
|
||||
mock_paginate.assert_called_once()
|
||||
select_stmt = mock_paginate.call_args.args[0]
|
||||
assert len(select_stmt._where_criteria) == 2
|
||||
|
||||
def test_get_datasets_legacy_owner_include_all_keeps_full_access(self):
|
||||
mock_db = MagicMock()
|
||||
mock_db.session.scalars.return_value.all.return_value = []
|
||||
mock_db.paginate.return_value.items = []
|
||||
mock_db.paginate.return_value.total = 0
|
||||
|
||||
user = DatasetServiceUnitDataFactory.create_user_mock(role=TenantAccountRole.OWNER)
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db", mock_db),
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.dify_config.RBAC_ENABLED", False),
|
||||
):
|
||||
mock_paginate.return_value = SimpleNamespace(items=[], total=0)
|
||||
DatasetService.get_datasets(
|
||||
page=1,
|
||||
per_page=20,
|
||||
@ -329,8 +329,8 @@ class TestDatasetServiceRetrievalPermissions:
|
||||
)
|
||||
|
||||
mock_db.session.scalars.assert_called_once()
|
||||
mock_db.paginate.assert_called_once()
|
||||
select_stmt = mock_db.paginate.call_args.kwargs["select"]
|
||||
mock_paginate.assert_called_once()
|
||||
select_stmt = mock_paginate.call_args.args[0]
|
||||
assert len(select_stmt._where_criteria) == 1
|
||||
|
||||
|
||||
|
||||
@ -270,9 +270,10 @@ class TestSegmentServiceQueries:
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db") as mock_db,
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.helper.escape_like_pattern", return_value="escaped") as escape_like,
|
||||
):
|
||||
mock_db.paginate.return_value = paginated
|
||||
mock_paginate.return_value = paginated
|
||||
|
||||
result = SegmentService.get_child_chunks(
|
||||
segment_id="segment-1",
|
||||
@ -285,7 +286,7 @@ class TestSegmentServiceQueries:
|
||||
|
||||
assert result is paginated
|
||||
escape_like.assert_called_once_with("needle")
|
||||
mock_db.paginate.assert_called_once()
|
||||
mock_paginate.assert_called_once()
|
||||
|
||||
def test_get_child_chunk_by_id_returns_only_child_chunk_instances(self):
|
||||
child_chunk = _make_child_chunk()
|
||||
@ -324,9 +325,10 @@ class TestSegmentServiceQueries:
|
||||
|
||||
with (
|
||||
patch("services.dataset_service.db") as mock_db,
|
||||
patch("services.dataset_service.paginate_query") as mock_paginate,
|
||||
patch("services.dataset_service.helper.escape_like_pattern", return_value="escaped") as escape_like,
|
||||
):
|
||||
mock_db.paginate.return_value = paginated
|
||||
mock_paginate.return_value = paginated
|
||||
|
||||
items, total = SegmentService.get_segments(
|
||||
document_id="doc-1",
|
||||
@ -340,7 +342,7 @@ class TestSegmentServiceQueries:
|
||||
assert items == ["segment"]
|
||||
assert total == 1
|
||||
escape_like.assert_called_once_with("needle")
|
||||
mock_db.paginate.assert_called_once()
|
||||
mock_paginate.assert_called_once()
|
||||
|
||||
def test_get_segment_by_id_returns_only_document_segment_instances(self):
|
||||
segment = DocumentSegment(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Loading…
Reference in New Issue
Block a user