diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index 55bc85483d5..70ce54830c7 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -602,7 +602,7 @@ class DatasetApi(Resource): if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) permissions = enterprise_rbac_service.RBACService.MyPermissions.get( @@ -774,7 +774,7 @@ class DatasetQueryApi(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -915,7 +915,7 @@ class DatasetRelatedAppListApi(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1194,7 +1194,7 @@ class DatasetPermissionUserListApi(Resource): if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/datasets_document.py b/api/controllers/console/datasets/datasets_document.py index 07e150617bf..499558dc4dd 100644 --- a/api/controllers/console/datasets/datasets_document.py +++ b/api/controllers/console/datasets/datasets_document.py @@ -186,7 +186,7 @@ class DocumentResource(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -206,7 +206,7 @@ class DocumentResource(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -247,7 +247,7 @@ class GetProcessRuleApi(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -322,7 +322,7 @@ class DatasetDocumentListApi(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -431,7 +431,7 @@ class DatasetDocumentListApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1173,7 +1173,7 @@ class DocumentStatusApi(DocumentResource): DatasetService.check_dataset_model_setting(dataset) # check user's permission - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) document_ids = request.args.getlist("document_id") @@ -1440,7 +1440,7 @@ class DocumentGenerateSummaryApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -1537,7 +1537,7 @@ class DocumentSummaryStatusApi(DocumentResource): # Check permissions try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/datasets_segments.py b/api/controllers/console/datasets/datasets_segments.py index 4858b5ff6b0..5ba115ff491 100644 --- a/api/controllers/console/datasets/datasets_segments.py +++ b/api/controllers/console/datasets/datasets_segments.py @@ -181,7 +181,7 @@ class DatasetDocumentSegmentListApi(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) @@ -302,7 +302,7 @@ class DatasetDocumentSegmentListApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) SegmentService.delete_segments(segment_ids, document, dataset) @@ -345,7 +345,7 @@ class DatasetDocumentSegmentApi(Resource): raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY: @@ -421,7 +421,7 @@ class DatasetDocumentSegmentAddApi(Resource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args @@ -494,7 +494,7 @@ class DatasetDocumentSegmentUpdateApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args @@ -550,7 +550,7 @@ class DatasetDocumentSegmentUpdateApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) SegmentService.delete_segment(segment, document, dataset) @@ -687,7 +687,7 @@ class ChildChunkAddApi(Resource): except ProviderTokenNotInitError as ex: raise ProviderNotInitializeError(ex.description) try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args @@ -789,7 +789,7 @@ class ChildChunkAddApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args @@ -862,7 +862,7 @@ class ChildChunkUpdateApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) try: @@ -930,7 +930,7 @@ class ChildChunkUpdateApi(Resource): if not current_user.is_dataset_editor: raise Forbidden() try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) # validate args diff --git a/api/controllers/console/datasets/external.py b/api/controllers/console/datasets/external.py index eb7b9aa84f8..99a61807a4d 100644 --- a/api/controllers/console/datasets/external.py +++ b/api/controllers/console/datasets/external.py @@ -382,7 +382,7 @@ class ExternalKnowledgeHitTestingApi(Resource): raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/hit_testing_base.py b/api/controllers/console/datasets/hit_testing_base.py index c343effa9a1..82c30fc7ffb 100644 --- a/api/controllers/console/datasets/hit_testing_base.py +++ b/api/controllers/console/datasets/hit_testing_base.py @@ -90,7 +90,7 @@ class DatasetsHitTestingBase: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/console/datasets/metadata.py b/api/controllers/console/datasets/metadata.py index ebb490cd9e8..90ce263dfe5 100644 --- a/api/controllers/console/datasets/metadata.py +++ b/api/controllers/console/datasets/metadata.py @@ -64,7 +64,7 @@ class DatasetMetadataCreateApi(Resource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) metadata = MetadataService.create_metadata( db.session(), dataset_id_str, metadata_args, current_user, current_tenant_id @@ -108,7 +108,7 @@ class DatasetMetadataApi(Resource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) metadata = MetadataService.update_metadata_name( db.session(), dataset_id_str, metadata_id_str, name, current_user, current_tenant_id @@ -128,7 +128,7 @@ class DatasetMetadataApi(Resource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) MetadataService.delete_metadata(db.session(), dataset_id_str, metadata_id_str) # Frontend callers only await success and invalidate metadata caches; no response body is consumed. @@ -165,7 +165,7 @@ class DatasetMetadataBuiltInFieldActionApi(Resource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) match action: case "enable": @@ -194,7 +194,7 @@ class DocumentMetadataEditApi(Resource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) metadata_args = MetadataOperationData.model_validate(console_ns.payload or {}) diff --git a/api/controllers/service_api/dataset/dataset.py b/api/controllers/service_api/dataset/dataset.py index 292c39f69bc..bfb7a045082 100644 --- a/api/controllers/service_api/dataset/dataset.py +++ b/api/controllers/service_api/dataset/dataset.py @@ -565,7 +565,7 @@ class DatasetApi(DatasetApiResource): if dataset is None: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) data = _dump_service_dataset_detail(dataset) @@ -819,7 +819,7 @@ class DocumentStatusApi(DatasetApiResource): # Check user's permission try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except services.errors.account.NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/controllers/service_api/dataset/metadata.py b/api/controllers/service_api/dataset/metadata.py index 7363e6bdfd4..3bb39f0cd4f 100644 --- a/api/controllers/service_api/dataset/metadata.py +++ b/api/controllers/service_api/dataset/metadata.py @@ -84,7 +84,7 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) metadata = MetadataService.create_metadata(db.session(), dataset_id_str, metadata_args) return dump_response(DatasetMetadataResponse, metadata), 201 @@ -157,7 +157,7 @@ class DatasetMetadataServiceApi(DatasetApiResource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) metadata = MetadataService.update_metadata_name(db.session(), dataset_id_str, metadata_id_str, payload.name) return dump_response(DatasetMetadataResponse, metadata), 200 @@ -192,7 +192,7 @@ class DatasetMetadataServiceApi(DatasetApiResource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) MetadataService.delete_metadata(db.session(), dataset_id_str, metadata_id_str) return "", 204 @@ -260,7 +260,7 @@ class DatasetMetadataBuiltInFieldActionServiceApi(DatasetApiResource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) match action: case "enable": @@ -306,7 +306,7 @@ class DocumentMetadataEditServiceApi(DatasetApiResource): dataset = DatasetService.get_dataset(dataset_id_str) if dataset is None: raise NotFound("Dataset not found.") - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) metadata_args = MetadataOperationData.model_validate(service_api_ns.payload or {}) diff --git a/api/services/dataset_service.py b/api/services/dataset_service.py index a8f341fdd04..2dd8c533828 100644 --- a/api/services/dataset_service.py +++ b/api/services/dataset_service.py @@ -248,7 +248,7 @@ class DatasetService: def get_datasets( page, per_page, - session: scoped_session | Session | None = None, + session: scoped_session | Session, tenant_id=None, user=None, search=None, @@ -257,7 +257,7 @@ class DatasetService: accessible_dataset_ids: list[str] | None = None, include_own_datasets: bool = False, ): - session = session or db.session + """Return visible datasets for a tenant, using the injected session for auxiliary permission lookups.""" query = select(Dataset).where(Dataset.tenant_id == tenant_id).order_by(Dataset.created_at.desc(), Dataset.id) if dify_config.RBAC_ENABLED and accessible_dataset_ids is not None: @@ -268,7 +268,7 @@ class DatasetService: if user: # get permitted dataset ids - dataset_permission = db.session.scalars( + dataset_permission = session.scalars( select(DatasetPermission).where( DatasetPermission.account_id == user.id, DatasetPermission.tenant_id == tenant_id ) @@ -652,7 +652,7 @@ class DatasetService: raise ValueError("Dataset name already exists") # Verify user has permission to update this dataset - DatasetService.check_dataset_permission(dataset, user) + DatasetService.check_dataset_permission(dataset, user, db.session) # Handle external dataset updates if dataset.provider == "external": @@ -1311,7 +1311,7 @@ class DatasetService: if dataset is None: return False - DatasetService.check_dataset_permission(dataset, user) + DatasetService.check_dataset_permission(dataset, user, db.session) dataset_was_deleted.send(dataset) @@ -1325,7 +1325,8 @@ class DatasetService: return db.session.execute(stmt).scalar_one() @staticmethod - def check_dataset_permission(dataset, user): + def check_dataset_permission(dataset, user, session: scoped_session | Session): + """Validate dataset access for a user, using the injected session for partial-member lookups.""" if dataset.tenant_id != user.current_tenant_id: logger.debug("User %s does not have permission to access dataset %s", user.id, dataset.id) raise NoPermissionError("You do not have permission to access this dataset.") @@ -1336,7 +1337,7 @@ class DatasetService: if dataset.permission == DatasetPermissionEnum.PARTIAL_TEAM: # For partial team permission, user needs explicit permission or be the maintainer. if dataset.maintainer != user.id: - user_permission = db.session.scalar( + user_permission = session.scalar( select(DatasetPermission) .where(DatasetPermission.dataset_id == dataset.id, DatasetPermission.account_id == user.id) .limit(1) @@ -1728,7 +1729,7 @@ class DocumentService: if not dataset: raise NotFound("Dataset not found.") try: - DatasetService.check_dataset_permission(dataset, current_user) + DatasetService.check_dataset_permission(dataset, current_user, db.session) except NoPermissionError as e: raise Forbidden(str(e)) diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py b/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py index 1ea1b10a15d..b88204b2a6a 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_permission_service.py @@ -410,7 +410,7 @@ class TestDatasetServiceCheckDatasetPermission: ) with pytest.raises(NoPermissionError): - DatasetService.check_dataset_permission(dataset, other_user) + DatasetService.check_dataset_permission(dataset, other_user, db_session_with_containers) def test_check_dataset_permission_owner_can_access_any_dataset(self, db_session_with_containers: Session): """Test that tenant owners can access any dataset regardless of permission level.""" @@ -423,7 +423,7 @@ class TestDatasetServiceCheckDatasetPermission: tenant.id, creator.id, permission=DatasetPermissionEnum.ONLY_ME ) - DatasetService.check_dataset_permission(dataset, owner) + DatasetService.check_dataset_permission(dataset, owner, db_session_with_containers) def test_check_dataset_permission_only_me_creator_can_access(self, db_session_with_containers: Session): """Test ONLY_ME permission allows only the dataset creator to access.""" @@ -433,7 +433,7 @@ class TestDatasetServiceCheckDatasetPermission: tenant.id, creator.id, permission=DatasetPermissionEnum.ONLY_ME ) - DatasetService.check_dataset_permission(dataset, creator) + DatasetService.check_dataset_permission(dataset, creator, db_session_with_containers) def test_check_dataset_permission_only_me_others_cannot_access(self, db_session_with_containers: Session): """Test ONLY_ME permission denies access to non-creators.""" @@ -447,7 +447,7 @@ class TestDatasetServiceCheckDatasetPermission: ) with pytest.raises(NoPermissionError): - DatasetService.check_dataset_permission(dataset, other) + DatasetService.check_dataset_permission(dataset, other, db_session_with_containers) def test_check_dataset_permission_all_team_allows_access(self, db_session_with_containers: Session): """Test ALL_TEAM permission allows any team member to access the dataset.""" @@ -460,7 +460,7 @@ class TestDatasetServiceCheckDatasetPermission: tenant.id, creator.id, permission=DatasetPermissionEnum.ALL_TEAM ) - DatasetService.check_dataset_permission(dataset, member) + DatasetService.check_dataset_permission(dataset, member, db_session_with_containers) def test_check_dataset_permission_partial_members_with_permission_success( self, db_session_with_containers: Session @@ -483,7 +483,7 @@ class TestDatasetServiceCheckDatasetPermission: DatasetPermissionTestDataFactory.create_dataset_permission(dataset.id, user.id, tenant.id) # Act (should not raise) - DatasetService.check_dataset_permission(dataset, user) + DatasetService.check_dataset_permission(dataset, user, db_session_with_containers) # Assert permissions = DatasetPermissionService.get_dataset_partial_member_list(dataset.id) @@ -510,7 +510,7 @@ class TestDatasetServiceCheckDatasetPermission: # Act & Assert with pytest.raises(NoPermissionError, match="You do not have permission to access this dataset"): - DatasetService.check_dataset_permission(dataset, user) + DatasetService.check_dataset_permission(dataset, user, db_session_with_containers) def test_check_dataset_permission_partial_team_creator_can_access(self, db_session_with_containers: Session): """Test PARTIAL_TEAM permission allows creator to access without explicit permission.""" @@ -520,7 +520,7 @@ class TestDatasetServiceCheckDatasetPermission: tenant.id, creator.id, permission=DatasetPermissionEnum.PARTIAL_TEAM ) - DatasetService.check_dataset_permission(dataset, creator) + DatasetService.check_dataset_permission(dataset, creator, db_session_with_containers) class TestDatasetServiceCheckDatasetOperatorPermission: diff --git a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py index d97c16668fd..ba5883f408d 100644 --- a/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py +++ b/api/tests/test_containers_integration_tests/services/test_dataset_service_permissions.py @@ -237,7 +237,7 @@ class TestDatasetServicePermissionsAndLifecycle: ) with pytest.raises(NoPermissionError, match="do not have permission"): - DatasetService.check_dataset_permission(dataset, outsider) + DatasetService.check_dataset_permission(dataset, outsider, db_session_with_containers) def test_check_dataset_permission_rejects_only_me_dataset_for_non_creator( self, db_session_with_containers: Session @@ -252,7 +252,7 @@ class TestDatasetServicePermissionsAndLifecycle: ) with pytest.raises(NoPermissionError, match="do not have permission"): - DatasetService.check_dataset_permission(dataset, member) + DatasetService.check_dataset_permission(dataset, member, db_session_with_containers) def test_check_dataset_permission_rejects_partial_team_user_without_binding( self, db_session_with_containers: Session @@ -267,7 +267,7 @@ class TestDatasetServicePermissionsAndLifecycle: ) with pytest.raises(NoPermissionError, match="do not have permission"): - DatasetService.check_dataset_permission(dataset, member) + DatasetService.check_dataset_permission(dataset, member, db_session_with_containers) def test_check_dataset_permission_allows_partial_team_creator(self, db_session_with_containers: Session): creator, tenant = DatasetPermissionIntegrationFactory.create_account_with_tenant( @@ -281,7 +281,7 @@ class TestDatasetServicePermissionsAndLifecycle: permission=DatasetPermissionEnum.PARTIAL_TEAM, ) - DatasetService.check_dataset_permission(dataset, creator) + DatasetService.check_dataset_permission(dataset, creator, db_session_with_containers) def test_check_dataset_permission_allows_partial_team_member_with_binding( self, db_session_with_containers: Session @@ -301,7 +301,7 @@ class TestDatasetServicePermissionsAndLifecycle: account_id=member.id, ) - DatasetService.check_dataset_permission(dataset, member) + DatasetService.check_dataset_permission(dataset, member, db_session_with_containers) def test_check_dataset_operator_permission_rejects_only_me_for_non_creator( self, db_session_with_containers: Session diff --git a/api/tests/unit_tests/services/test_dataset_service_dataset.py b/api/tests/unit_tests/services/test_dataset_service_dataset.py index 46f32f93e8a..044e0e5ab40 100644 --- a/api/tests/unit_tests/services/test_dataset_service_dataset.py +++ b/api/tests/unit_tests/services/test_dataset_service_dataset.py @@ -173,7 +173,8 @@ class TestDatasetServiceRetrievalPermissions: def test_get_datasets_filters_by_maintainer_and_rbac_overrides(self): mock_db = MagicMock() - mock_db.session.scalars.return_value.all.return_value = [] + 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) @@ -189,12 +190,15 @@ class TestDatasetServiceRetrievalPermissions: DatasetService.get_datasets( page=1, per_page=20, + session=explicit_session, tenant_id="tenant-1", user=user, accessible_dataset_ids=["dataset-shared"], include_own_datasets=True, ) + explicit_session.scalars.assert_called_once() + mock_db.session.scalars.assert_not_called() select_stmt = mock_db.paginate.call_args.kwargs["select"] visibility_clause = str(select_stmt._where_criteria[1]) assert "maintainer" in visibility_clause @@ -218,6 +222,7 @@ class TestDatasetServiceRetrievalPermissions: DatasetService.get_datasets( page=1, per_page=20, + session=mock_db.session, tenant_id="tenant-1", user=user, accessible_dataset_ids=["dataset-shared"], @@ -272,7 +277,14 @@ class TestDatasetServiceRetrievalPermissions: return_value=mock_permissions, ), ): - DatasetService.get_datasets(page=1, per_page=20, tenant_id="tenant-1", user=user, include_all=True) + DatasetService.get_datasets( + page=1, + per_page=20, + session=mock_db.session, + tenant_id="tenant-1", + user=user, + include_all=True, + ) mock_db.session.scalars.assert_called_once() mock_db.paginate.assert_called_once() @@ -289,7 +301,7 @@ class TestDatasetServiceRetrievalPermissions: patch("services.dataset_service.db", mock_db), patch("services.dataset_service.dify_config.RBAC_ENABLED", True), ): - DatasetService.get_datasets(page=1, per_page=20, tenant_id="tenant-1", user=None) + 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() @@ -308,7 +320,14 @@ class TestDatasetServiceRetrievalPermissions: patch("services.dataset_service.db", mock_db), patch("services.dataset_service.dify_config.RBAC_ENABLED", False), ): - DatasetService.get_datasets(page=1, per_page=20, tenant_id="tenant-1", user=user, include_all=True) + DatasetService.get_datasets( + page=1, + per_page=20, + session=mock_db.session, + tenant_id="tenant-1", + user=user, + include_all=True, + ) mock_db.session.scalars.assert_called_once() mock_db.paginate.assert_called_once() @@ -517,7 +536,9 @@ class TestDatasetServiceCreationAndUpdate: result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user) assert result == "updated" - check_permission.assert_called_once_with(dataset, user) + check_permission.assert_called_once() + assert check_permission.call_args.args[:2] == (dataset, user) + assert len(check_permission.call_args.args) == 3 update_external.assert_called_once_with(dataset, {"name": dataset.name}, user) def test_update_dataset_routes_internal_datasets_to_internal_helper(self): @@ -533,7 +554,9 @@ class TestDatasetServiceCreationAndUpdate: result = DatasetService.update_dataset("dataset-1", {"name": dataset.name}, user) assert result == "updated" - check_permission.assert_called_once_with(dataset, user) + check_permission.assert_called_once() + assert check_permission.call_args.args[:2] == (dataset, user) + assert len(check_permission.call_args.args) == 3 update_internal.assert_called_once_with(dataset, {"name": dataset.name}, user) def test_has_dataset_same_name_returns_true_when_query_matches(self):