From bf04f7dab3952e46ffa450a2d07d8d3833beceb3 Mon Sep 17 00:00:00 2001 From: Alexandru Date: Mon, 3 Aug 2026 19:45:51 +0300 Subject: [PATCH] skip chunk quota checks for superusers --- py/core/main/services/ingestion_service.py | 47 ++++----- .../unit/ingestion/test_ingestion_service.py | 99 +++++++++++++++++++ 2 files changed, 123 insertions(+), 23 deletions(-) create mode 100644 py/tests/unit/ingestion/test_ingestion_service.py diff --git a/py/core/main/services/ingestion_service.py b/py/core/main/services/ingestion_service.py index 189a39cbf..fccc7a33c 100644 --- a/py/core/main/services/ingestion_service.py +++ b/py/core/main/services/ingestion_service.py @@ -445,38 +445,38 @@ async def store_embeddings( vector_batch: list[VectorEntry] = [] document_counts: dict[UUID, int] = {} - # We'll track usage from the first user we see; if your scenario allows - # multiple user owners in a single ingestion, you'd need to refine usage checks. - current_usage = None - user_id_for_usage_check: UUID | None = None - - count = 0 + # We track usage from the first owner because an ingestion is expected + # to contain chunks for a single owner. + owner_id = vector_entries[0].owner_id + owner = await self.providers.database.users_handler.get_user_by_id( + owner_id + ) - for msg in vector_entries: - # If we haven't set usage yet, do so on the first chunk - if current_usage is None: - user_id_for_usage_check = msg.owner_id - usage_data = ( - await self.providers.database.chunks_handler.list_chunks( - limit=1, - offset=0, - filters={"owner_id": msg.owner_id}, - ) + current_usage = None + max_chunks = None + if not owner.is_superuser: + usage_data = ( + await self.providers.database.chunks_handler.list_chunks( + limit=1, + offset=0, + filters={"owner_id": owner_id}, ) - current_usage = usage_data["total_entries"] - - # Figure out the user's limit - user = await self.providers.database.users_handler.get_user_by_id( - msg.owner_id ) + current_usage = usage_data["total_entries"] max_chunks = ( self.providers.database.config.app.default_max_chunks_per_user if self.providers.database.config.app else 1e10 ) - if user.limits_overrides and "max_chunks" in user.limits_overrides: - max_chunks = user.limits_overrides["max_chunks"] + if ( + owner.limits_overrides + and "max_chunks" in owner.limits_overrides + ): + max_chunks = owner.limits_overrides["max_chunks"] + count = 0 + + for msg in vector_entries: # Add to our local batch vector_batch.append(msg) document_counts[msg.document_id] = ( @@ -487,6 +487,7 @@ async def store_embeddings( # Check usage if ( current_usage is not None + and max_chunks is not None and (current_usage + len(vector_batch) + count) > max_chunks ): error_message = f"User {msg.owner_id} has exceeded the maximum number of allowed chunks: {max_chunks}" diff --git a/py/tests/unit/ingestion/test_ingestion_service.py b/py/tests/unit/ingestion/test_ingestion_service.py new file mode 100644 index 000000000..5c9e8e1d5 --- /dev/null +++ b/py/tests/unit/ingestion/test_ingestion_service.py @@ -0,0 +1,99 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock +from uuid import uuid4 + +import pytest + +from core.base import Vector, VectorEntry +from core.base.api.models import User +from core.main.services.ingestion_service import IngestionService + + +def _make_vector_entry(owner_id): + return VectorEntry( + id=uuid4(), + document_id=uuid4(), + owner_id=owner_id, + collection_ids=[], + vector=Vector(data=[0.1]), + text="test chunk", + metadata={}, + ) + + +def _make_service(owner, current_usage=0, max_chunks=10_000): + chunks_handler = SimpleNamespace( + list_chunks=AsyncMock( + return_value={"results": [], "total_entries": current_usage} + ), + upsert_entries=AsyncMock(), + ) + users_handler = SimpleNamespace( + get_user_by_id=AsyncMock(return_value=owner) + ) + database = SimpleNamespace( + chunks_handler=chunks_handler, + users_handler=users_handler, + config=SimpleNamespace( + app=SimpleNamespace( + default_max_chunks_per_user=max_chunks, + ) + ), + ) + providers = SimpleNamespace(database=database) + service = IngestionService(config=Mock(), providers=providers) + return service, chunks_handler, users_handler + + +@pytest.mark.asyncio +async def test_store_embeddings_skips_chunk_quota_for_superuser(): + owner = User( + id=uuid4(), + email="admin@example.com", + is_superuser=True, + ) + vector_entry = _make_vector_entry(owner.id) + service, chunks_handler, users_handler = _make_service( + owner, + current_usage=10_000, + max_chunks=1, + ) + + messages = [ + message async for message in service.store_embeddings([vector_entry]) + ] + + users_handler.get_user_by_id.assert_awaited_once_with(owner.id) + chunks_handler.list_chunks.assert_not_awaited() + chunks_handler.upsert_entries.assert_awaited_once_with([vector_entry]) + assert messages == [ + f"Successful ingestion for document_id: {vector_entry.document_id}, " + "with vector count: 1" + ] + + +@pytest.mark.asyncio +async def test_store_embeddings_keeps_chunk_quota_for_normal_user(): + owner = User( + id=uuid4(), + email="user@example.com", + is_superuser=False, + ) + vector_entry = _make_vector_entry(owner.id) + service, chunks_handler, users_handler = _make_service(owner) + + messages = [ + message async for message in service.store_embeddings([vector_entry]) + ] + + users_handler.get_user_by_id.assert_awaited_once_with(owner.id) + chunks_handler.list_chunks.assert_awaited_once_with( + limit=1, + offset=0, + filters={"owner_id": owner.id}, + ) + chunks_handler.upsert_entries.assert_awaited_once_with([vector_entry]) + assert messages == [ + f"Successful ingestion for document_id: {vector_entry.document_id}, " + "with vector count: 1" + ]