diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index cd39aa3931f..89a7ba20d75 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44526 + "limit": 44519 }, "reportUnknownLambdaType": { "limit": 109 @@ -111,7 +111,7 @@ "limit": 19829 }, "reportUnknownVariableType": { - "limit": 30349 + "limit": 30348 }, "reportUnnecessaryCast": { "limit": 117 diff --git a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py index 57cb5e73b64..f593d747558 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py @@ -1,4 +1,5 @@ import asyncio +from collections.abc import Sequence from typing import Final from litellm._logging import verbose_proxy_logger @@ -23,6 +24,22 @@ class SpendUpdateQueue(BaseUpdateQueue): def __init__(self): super().__init__() self.update_queue: asyncio.Queue[SpendUpdateQueueItem] = asyncio.Queue(maxsize=LITELLM_ASYNCIO_QUEUE_MAXSIZE) + self._pending_updates: Final[ + dict[ + tuple[Litellm_EntityType | None, str | None], SpendUpdateQueueItem + ] # mutable-ok: [LIT001] queue-owned index + ] = {} # mutable-ok: [LIT002] index owns only entries still in the bounded queue + + async def flush_all_updates_from_in_memory_queue( + self, + ) -> list[SpendUpdateQueueItem]: # mutable-ok: [LIT001] preserves the base queue's mutable batch contract + updates: Final[list[SpendUpdateQueueItem]] = ( # mutable-ok: [LIT001] batch owned by the base queue + await super().flush_all_updates_from_in_memory_queue() + ) + for key, update in (((item.get("entity_type"), item.get("entity_id")), item) for item in updates): + if self._pending_updates.get(key) is update: + del self._pending_updates[key] + return updates async def flush_and_get_aggregated_db_spend_update_transactions( self, @@ -37,27 +54,28 @@ class SpendUpdateQueue(BaseUpdateQueue): verbose_proxy_logger.debug("Aggregating updates by entity type: %s", updates) return self.get_aggregated_db_spend_update_transactions(updates) - async def add_update(self, update: SpendUpdateQueueItem): + async def add_update(self, update: SpendUpdateQueueItem) -> None: """Enqueue an update to the spend update queue""" verbose_proxy_logger.debug("Adding update to queue: %s", update) - await self.update_queue.put(update) + key: Final = (update.get("entity_type"), update.get("entity_id")) + pending: Final = self._pending_updates.get(key) + if pending is not None: + pending["response_cost"] = (pending.get("response_cost") or 0) + (update.get("response_cost") or 0) + return - # if the queue is full, aggregate the updates - if self.update_queue.qsize() >= self.MAX_SIZE_IN_MEMORY_QUEUE: - verbose_proxy_logger.warning( - "Spend update queue is full. Aggregating all entries in queue to concatenate entries." - ) - await self.aggregate_queue_updates() + queued_update: Final = update.copy() + await self.update_queue.put(queued_update) + self._pending_updates[key] = queued_update async def aggregate_queue_updates(self): """Concatenate all updates in the queue to reduce the size of in-memory queue""" - updates: Final[list[SpendUpdateQueueItem]] = await self.flush_all_updates_from_in_memory_queue() + updates: Final = await self.flush_all_updates_from_in_memory_queue() aggregated_updates: Final = self._get_aggregated_spend_update_queue_item(updates) for update in aggregated_updates: - await self.update_queue.put(update) + await self.add_update(update) def _get_aggregated_spend_update_queue_item( - self, updates: list[SpendUpdateQueueItem] + self, updates: Sequence[SpendUpdateQueueItem] ) -> list[SpendUpdateQueueItem]: """ This is used to reduce the size of the in-memory queue by aggregating updates by entity type + id @@ -126,7 +144,7 @@ class SpendUpdateQueue(BaseUpdateQueue): return aggregated_spend_updates def get_aggregated_db_spend_update_transactions( - self, updates: list[SpendUpdateQueueItem] + self, updates: Sequence[SpendUpdateQueueItem] ) -> DBSpendUpdateTransactions: """Aggregate updates by entity type.""" # Initialize all transaction lists as empty dicts diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 149c44ed083..e6b0d0e074e 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -9,7 +9,7 @@ "limit": 827 }, "ANN201": { - "limit": 2012 + "limit": 2010 }, "ANN202": { "limit": 847 diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py index 43f1a820885..4efd0d514f3 100644 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_update_queue.py @@ -1,10 +1,11 @@ import asyncio import json +from typing import Final +from unittest.mock import AsyncMock import pytest from fastapi.testclient import TestClient -from litellm.constants import MAX_SIZE_IN_MEMORY_QUEUE from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue @@ -149,27 +150,17 @@ async def test_missing_entity_type(spend_queue): @pytest.mark.asyncio -async def test_queue_max_size_triggers_aggregation(monkeypatch, spend_queue): - """Test that reaching MAX_SIZE_IN_MEMORY_QUEUE triggers aggregation""" - # Override MAX_SIZE_IN_MEMORY_QUEUE for testing - monkeypatch.setattr(spend_queue, "MAX_SIZE_IN_MEMORY_QUEUE", 6) - - # Add 6 updates for the same user (exceeding the max size) - for i in range(6): - update: SpendUpdateQueueItem = { +async def test_repeated_updates_share_one_queue_entry(spend_queue: SpendUpdateQueue) -> None: + for _ in range(6): + update: Final[SpendUpdateQueueItem] = { "entity_type": Litellm_EntityType.USER, "entity_id": "user123", "response_cost": 1.0, } await spend_queue.add_update(update) - # Queue should have been aggregated, resulting in a single entry assert spend_queue.update_queue.qsize() == 1 - - # Verify the aggregated cost is correct - aggregated = ( - await spend_queue.flush_and_get_aggregated_db_spend_update_transactions() - ) + aggregated: Final = await spend_queue.flush_and_get_aggregated_db_spend_update_transactions() assert aggregated["user_list_transactions"]["user123"] == 6.0 @@ -254,13 +245,8 @@ def test_get_aggregated_spend_update_queue_item_does_not_mutate_original_updates @pytest.mark.asyncio -async def test_queue_size_reduction_with_large_volume(monkeypatch, spend_queue): - """Test that queue size is actually reduced when dealing with many items""" - # Set a smaller MAX_SIZE for testing - monkeypatch.setattr(spend_queue, "MAX_SIZE_IN_MEMORY_QUEUE", 10) - - # Add 30 updates (200 for user1, 10 for key1) - for i in range(200): +async def test_queue_size_reduction_with_large_volume(spend_queue: SpendUpdateQueue) -> None: + for _ in range(200): await spend_queue.add_update( { "entity_type": Litellm_EntityType.USER, @@ -269,11 +255,9 @@ async def test_queue_size_reduction_with_large_volume(monkeypatch, spend_queue): } ) - # At this point, aggregation should have happened at least once - # Queue size should be much less than 20 - assert spend_queue.update_queue.qsize() <= 10 + assert spend_queue.update_queue.qsize() == 1 - for i in range(300): + for _ in range(300): await spend_queue.add_update( { "entity_type": Litellm_EntityType.KEY, @@ -282,12 +266,108 @@ async def test_queue_size_reduction_with_large_volume(monkeypatch, spend_queue): } ) - # Queue should have at most 2 items after all this activity - assert spend_queue.update_queue.qsize() <= 10 - - # Verify total costs are correct - aggregated = ( - await spend_queue.flush_and_get_aggregated_db_spend_update_transactions() - ) + assert spend_queue.update_queue.qsize() == 2 + aggregated: Final = await spend_queue.flush_and_get_aggregated_db_spend_update_transactions() assert aggregated["user_list_transactions"]["user1"] == 200 * 0.5 assert aggregated["key_list_transactions"]["key1"] == 300 * 1.0 + + +@pytest.mark.asyncio +async def test_high_cardinality_updates_do_not_rescan_the_queue(monkeypatch: pytest.MonkeyPatch) -> None: + queue: Final = SpendUpdateQueue() + queue.MAX_SIZE_IN_MEMORY_QUEUE = 4 + aggregate: Final = AsyncMock(wraps=queue.aggregate_queue_updates) + monkeypatch.setattr(queue, "aggregate_queue_updates", aggregate) + updates: Final = tuple( + SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id=f"key-{index}", response_cost=0.25) + for index in range(6) + ) + + for _ in range(20): + for update in updates: + await queue.add_update(update) + + aggregate.assert_not_awaited() + assert queue.update_queue.qsize() == len(updates) + flushed: Final = await queue.flush_and_get_aggregated_db_spend_update_transactions() + assert flushed["key_list_transactions"] == {f"key-{index}": 5.0 for index in range(6)} + assert all(update["response_cost"] == 0.25 for update in updates) + + +@pytest.mark.asyncio +async def test_full_queue_coalesces_existing_keys_and_cancels_new_keys() -> None: + queue: Final = SpendUpdateQueue() + queue.update_queue = asyncio.Queue(maxsize=1) + first: Final = SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id="key", response_cost=1.0) + await queue.add_update(first) + await asyncio.wait_for(queue.add_update(first), timeout=1) + blocked: Final = asyncio.create_task( + queue.add_update(SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id="cancelled", response_cost=9.0)) + ) + try: + await asyncio.sleep(0) + assert not blocked.done() + blocked.cancel() + with pytest.raises(asyncio.CancelledError): + await blocked + finally: + blocked.cancel() + await asyncio.gather(blocked, return_exceptions=True) + + flushed: Final = await queue.flush_all_updates_from_in_memory_queue() + assert flushed == [SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id="key", response_cost=2.0)] + await queue.add_update(first) + assert flushed[0]["response_cost"] == 2.0 + next_batch: Final = await queue.flush_and_get_aggregated_db_spend_update_transactions() + assert next_batch["key_list_transactions"] == {"key": 1.0} + assert first["response_cost"] == 1.0 + + +@pytest.mark.asyncio +async def test_partial_flush_keeps_pending_costs_separate(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.db.db_transaction_queue import base_update_queue + + monkeypatch.setattr(base_update_queue, "MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1) + queue: Final = SpendUpdateQueue() + first: Final = SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id="first", response_cost=1.0) + second: Final = SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id="second", response_cost=2.0) + await queue.add_update(first) + await queue.add_update(second) + flushed: Final = await queue.flush_all_updates_from_in_memory_queue() + await queue.add_update(first) + await queue.add_update(second) + + assert flushed == [first] + second_batch: Final = await queue.flush_and_get_aggregated_db_spend_update_transactions() + assert second_batch["key_list_transactions"] == {"second": 4.0} + third_batch: Final = await queue.flush_and_get_aggregated_db_spend_update_transactions() + assert third_batch["key_list_transactions"] == {"first": 1.0} + + +@pytest.mark.asyncio +async def test_partial_flush_preserves_same_key_admitted_by_waiting_producers(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.db.db_transaction_queue import base_update_queue + + queue: Final = SpendUpdateQueue() + queue.update_queue = asyncio.Queue(maxsize=2) + for key in ("first", "second"): + await queue.add_update(SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id=key, response_cost=1.0)) + duplicate: Final = SpendUpdateQueueItem(entity_type=Litellm_EntityType.KEY, entity_id="shared", response_cost=1.0) + producers: Final = tuple(asyncio.create_task(queue.add_update(duplicate)) for _ in range(2)) + try: + await asyncio.sleep(0) + assert all(not producer.done() for producer in producers) + await queue.flush_all_updates_from_in_memory_queue() + await asyncio.wait_for(asyncio.gather(*producers), timeout=1) + finally: + for producer in producers: + producer.cancel() + await asyncio.gather(*producers, return_exceptions=True) + + monkeypatch.setattr(base_update_queue, "MAX_IN_MEMORY_QUEUE_FLUSH_COUNT", 1) + first_batch: Final = await queue.flush_and_get_aggregated_db_spend_update_transactions() + assert first_batch["key_list_transactions"] == {"shared": 1.0} + await queue.add_update(duplicate) + second_batch: Final = await queue.flush_and_get_aggregated_db_spend_update_transactions() + assert second_batch["key_list_transactions"] == {"shared": 2.0} + assert queue.update_queue.empty() diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 959b2eada25..89033331424 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22733 + "limit": 22729 }, "LIT002": { "limit": 26860