fix(spend-tracking): forget an empty spend-log lookup after thirty seconds

The daily spend rows and the spend logs of one batch are written a moment apart, so a usage read landing between them used to remember the session as nameless for ten minutes on that worker. Found identities keep the ten minute entry
This commit is contained in:
mateo-berri 2026-09-08 13:28:13 -07:00
parent f4aef5a1db
commit 0c2d0f4777
3 changed files with 38 additions and 2 deletions

View file

@ -1764,6 +1764,7 @@ SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600))
SPEND_LOG_KEY_METADATA_CACHE_TTL: Final = 600
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL: Final = 30
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS: Final = 10000
# Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated
# callers from forcing a DB query per request for unknown names, while bounding

View file

@ -9,7 +9,11 @@ from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS, SPEND_LOG_KEY_METADATA_CACHE_TTL
from litellm.constants import (
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
SPEND_LOG_KEY_METADATA_CACHE_TTL,
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
)
from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
from litellm.proxy.utils import PrismaClient
from litellm.repositories.user_repository import UserRepository
@ -230,6 +234,17 @@ async def _query_spend_log_metadata(
)
def _remember_spend_log_metadata(
cache: InMemoryCache, digest: str, window: tuple[datetime, datetime], meta: KeyMetadataDict | None
) -> None:
if meta is None:
cache.set_cache(
_spend_log_cache_key(digest, window), KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL
)
return
cache.set_cache(_spend_log_cache_key(digest, window), meta)
async def recover_key_metadata_from_spend_logs(
prisma_client: PrismaClient,
missing_keys: AbstractSet[str],
@ -251,7 +266,7 @@ async def recover_key_metadata_from_spend_logs(
fresh: Final = await _query_spend_log_metadata(prisma_client, uncached, window) if uncached else _EMPTY_KEY_METADATA
if fresh is not None:
for digest in uncached:
cache.set_cache(_spend_log_cache_key(digest, window), fresh.get(digest, KeyMetadataDict()))
_remember_spend_log_metadata(cache, digest, window, fresh.get(digest))
return MappingProxyType(
{digest: meta for digest, meta in (*cached.items(), *(fresh or _EMPTY_KEY_METADATA).items()) if meta}
)

View file

@ -1,3 +1,4 @@
import time
from collections.abc import Sequence
from datetime import datetime
from types import SimpleNamespace
@ -7,6 +8,7 @@ import pytest
from prisma.errors import PrismaError
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import SPEND_LOG_KEY_METADATA_CACHE_TTL, SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL
from litellm.proxy.spend_tracking.key_metadata_recovery import (
fill_missing_api_key_aliases,
recover_double_hashed_key_metadata,
@ -395,3 +397,21 @@ async def test_recover_key_metadata_from_spend_logs_does_not_cache_a_failed_quer
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
assert result[digest]["key_alias"] == "back-online"
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_forgets_a_miss_long_before_a_hit():
found = hash_token("cli-session-found")
unknown = hash_token("cli-session-unknown")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(found, "found-alias", None, None)])
started = time.time()
await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache)
hit_expires = next(deadline for key, deadline in cache.ttl_dict.items() if found in key)
miss_expires = next(deadline for key, deadline in cache.ttl_dict.items() if unknown in key)
assert miss_expires - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1
assert hit_expires - started >= SPEND_LOG_KEY_METADATA_CACHE_TTL - 1