fix(spend-tracking): run one spend-log scan at a time and back off repeated misses

Concurrent usage reads on one worker now share a single spend-log query
instead of each scanning the same window, and a digest that comes back
nameless a second time is remembered for the full ten minutes rather than
thirty seconds, so a key that never resolves costs at most two scans per
worker per window per ten minutes. The first miss still expires after
thirty seconds so a read that lands between the daily spend flush and the
spend-log flush recovers on the next read
This commit is contained in:
mateo-berri 2026-09-08 15:00:56 -07:00
parent 0c2d0f4777
commit 907c200b31
2 changed files with 92 additions and 23 deletions

View file

@ -1,3 +1,4 @@
import asyncio
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from datetime import datetime
@ -77,6 +78,7 @@ _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
max_size_in_memory=SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL,
)
_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
_EMPTY_EMAILS: Final[Mapping[str, str]] = MappingProxyType({})
@ -205,11 +207,17 @@ def _spend_log_cache_key(digest: str, window: tuple[datetime, datetime]) -> str:
def _cached_spend_log_metadata(
cache: InMemoryCache,
digest: str,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> KeyMetadataDict | None:
cached: Final[object] = cache.get_cache(_spend_log_cache_key(digest, window))
return None if cached is None else _CACHED_KEY_METADATA.validate_python(cached)
) -> Mapping[str, KeyMetadataDict]:
return MappingProxyType(
{
digest: _CACHED_KEY_METADATA.validate_python(cached)
for digest in digests
for cached in (cache.get_cache(_spend_log_cache_key(digest, window)),)
if cached is not None
}
)
async def _query_spend_log_metadata(
@ -237,12 +245,36 @@ 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
)
key: Final = _spend_log_cache_key(digest, window)
if meta is not None:
cache.set_cache(key, meta)
return
cache.set_cache(_spend_log_cache_key(digest, window), meta)
missed_before: Final = f"{key}:missed-before"
if cache.get_cache(missed_before) is not None:
cache.set_cache(key, KeyMetadataDict())
return
cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL)
cache.set_cache(missed_before, True)
async def _spend_log_metadata_one_query_at_a_time(
prisma_client: PrismaClient,
cache: InMemoryCache,
lock: asyncio.Lock,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict]:
async with lock:
settled: Final = _cached_spend_log_metadata(cache, digests, window)
pending: Final = digests - frozenset(settled)
fresh: Final = (
await _query_spend_log_metadata(prisma_client, pending, window) if pending else _EMPTY_KEY_METADATA
)
if fresh is None:
return settled
for digest in pending:
_remember_spend_log_metadata(cache, digest, window, fresh.get(digest))
return MappingProxyType({**settled, **fresh})
async def recover_key_metadata_from_spend_logs(
@ -250,26 +282,19 @@ async def recover_key_metadata_from_spend_logs(
missing_keys: AbstractSet[str],
window: tuple[datetime, datetime],
cache: InMemoryCache = _SPEND_LOG_METADATA_CACHE,
lock: asyncio.Lock = _SPEND_LOG_QUERY_LOCK,
) -> Mapping[str, KeyMetadataDict]:
digests: Final = frozenset(key for key in missing_keys if _is_spend_log_digest(key))
if not digests:
return _EMPTY_KEY_METADATA
cached: Final = MappingProxyType(
{
digest: meta
for digest in digests
for meta in (_cached_spend_log_metadata(cache, digest, window),)
if meta is not None
}
)
cached: Final = _cached_spend_log_metadata(cache, digests, window)
uncached: Final = digests - frozenset(cached)
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:
_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}
settled: Final = (
await _spend_log_metadata_one_query_at_a_time(prisma_client, cache, lock, uncached, window)
if uncached
else _EMPTY_KEY_METADATA
)
return MappingProxyType({digest: meta for digest, meta in (*cached.items(), *settled.items()) if meta})
def _row_with_recovered_fields(

View file

@ -1,3 +1,4 @@
import asyncio
import time
from collections.abc import Sequence
from datetime import datetime
@ -415,3 +416,46 @@ async def test_recover_key_metadata_from_spend_logs_forgets_a_miss_long_before_a
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
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_runs_one_query_for_concurrent_lookups():
digest = hash_token("cli-session-shared")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
lock = asyncio.Lock()
mock_prisma = MagicMock()
async def slow_query_raw(sql: str, *params: object) -> list[dict[str, str | None]]:
await asyncio.sleep(0.01)
return [_digest_row(digest, "shared-alias", None, None)]
mock_prisma.db.query_raw = AsyncMock(side_effect=slow_query_raw)
results = await asyncio.gather(
*(
recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache, lock=lock)
for _ in range(9)
)
)
assert all(result[digest]["key_alias"] == "shared-alias" for result in results)
assert mock_prisma.db.query_raw.await_count == 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_keeps_a_repeated_miss_as_long_as_a_hit():
unknown = hash_token("cli-session-never-named")
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([])
await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache)
first_miss_key = next(key for key in cache.ttl_dict if unknown in key and not key.endswith(":missed-before"))
cache.ttl_dict[first_miss_key] = time.time() - 1
started = time.time()
await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache)
assert mock_prisma.db.query_raw.await_count == 2
assert cache.ttl_dict[first_miss_key] - started >= SPEND_LOG_KEY_METADATA_CACHE_TTL - 1