Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 24 additions & 23 deletions py/core/main/services/ingestion_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -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] = (
Expand All @@ -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}"
Expand Down
99 changes: 99 additions & 0 deletions py/tests/unit/ingestion/test_ingestion_service.py
Original file line number Diff line number Diff line change
@@ -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"
]
Loading