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:
Checo 2026-07-04 12:50:01 +08:00 committed by GitHub
parent 070aed81d9
commit 2e1ab194b7
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
22 changed files with 383 additions and 308 deletions

View File

@ -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:

View File

@ -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)

View File

@ -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(

View File

@ -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]

View File

@ -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:

View File

@ -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
View 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,
)

View File

@ -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

View File

@ -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

View File

@ -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:

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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))

View File

@ -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(

View File

@ -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

View File

@ -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)

View File

@ -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

View File

@ -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(

View File

@ -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

View File

@ -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(