diff --git a/litellm/constants.py b/litellm/constants.py index cfc5b6b86a7..e599a5dd3f9 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1766,6 +1766,7 @@ DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACH 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 +SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS: Final = 5000 # Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated # callers from forcing a DB query per request for unknown names, while bounding # staleness so a transient DB error (which surfaces as an empty list) cannot diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index d555da99397..29688b61b3d 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -1,7 +1,7 @@ import asyncio from collections.abc import Awaitable, Callable, Mapping, Sequence from collections.abc import Set as AbstractSet -from datetime import datetime +from datetime import datetime, timedelta from types import MappingProxyType from typing import Final, TypeVar @@ -14,6 +14,7 @@ from litellm.constants import ( SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS, SPEND_LOG_KEY_METADATA_CACHE_TTL, SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL, + SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS, ) from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash from litellm.proxy.utils import PrismaClient @@ -36,16 +37,15 @@ ORDER BY token, deleted_at DESC """ _SPEND_LOG_ALIAS_SQL: Final = """ -SELECT DISTINCT ON (api_key) - api_key AS digest, - key_alias, - team_id, - user_id, - MIN(user_id) OVER (PARTITION BY api_key) AS first_owner, - MAX(user_id) OVER (PARTITION BY api_key) AS last_owner +SELECT api_key AS digest, + MIN(key_alias) AS first_alias, + MAX(key_alias) AS last_alias, + MIN(team_id) AS first_team, + MAX(team_id) AS last_team, + MIN(user_id) AS first_owner, + MAX(user_id) AS last_owner FROM ( SELECT api_key, - "startTime", NULLIF(metadata->>'user_api_key_alias', '') AS key_alias, COALESCE(NULLIF(team_id, ''), NULLIF(metadata->>'user_api_key_team_id', '')) AS team_id, COALESCE(NULLIF("user", ''), NULLIF(metadata->>'user_api_key_user_id', '')) AS user_id @@ -55,9 +55,12 @@ FROM ( AND "startTime" < $3::timestamp ) named WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL -ORDER BY api_key, "startTime" DESC +GROUP BY api_key """ +_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}" +_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS) + _HASHED_JWT_PREFIX: Final = "hashed-jwt-" @@ -75,12 +78,25 @@ class _TokenDigestRow(BaseModel): user_id: str | None = None -class _SpendLogDigestRow(_TokenDigestRow): +def _unanimous(first: str | None, last: str | None) -> str | None: + return first if first == last else None + + +class _SpendLogDigestRow(BaseModel): + digest: str + first_alias: str | None = None + last_alias: str | None = None + first_team: str | None = None + last_team: str | None = None first_owner: str | None = None last_owner: str | None = None - def unanimous_owner(self) -> str | None: - return self.user_id if self.first_owner == self.last_owner else None + def metadata(self) -> KeyMetadataDict: + return KeyMetadataDict( + key_alias=_unanimous(self.first_alias, self.last_alias), + team_id=_unanimous(self.first_team, self.last_team), + user_id=_unanimous(self.first_owner, self.last_owner), + ) _TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...]) @@ -232,14 +248,24 @@ def _cached_spend_log_metadata( ) +async def _spend_log_rows_within_the_statement_timeout( + prisma_client: PrismaClient, + digests: AbstractSet[str], + window: tuple[datetime, datetime], +) -> Sequence[Mapping[str, object]]: + start, end = window + async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction: + await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL) + return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end) + + async def _query_spend_log_metadata( prisma_client: PrismaClient, digests: AbstractSet[str], window: tuple[datetime, datetime], ) -> Mapping[str, KeyMetadataDict] | None: - start, end = window rows: Final = await _db_or_empty( - lambda: prisma_client.db.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end), + lambda: _spend_log_rows_within_the_statement_timeout(prisma_client, digests, window), "Failed spend-log alias recovery for %d missing keys: %s", len(digests), ) @@ -247,10 +273,10 @@ async def _query_spend_log_metadata( return None return MappingProxyType( { - row.digest: KeyMetadataDict(key_alias=row.key_alias, team_id=row.team_id, user_id=owner) + row.digest: meta for row in _SPEND_LOG_DIGEST_ROWS.validate_python(rows) - for owner in (row.unanimous_owner(),) - if row.digest in digests and (row.key_alias or owner or row.team_id) + for meta in (row.metadata(),) + if row.digest in digests and any(meta.values()) } ) diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index 682b11f5bb7..89be341c87b 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -1,7 +1,7 @@ import asyncio import time from collections.abc import Sequence -from datetime import datetime +from datetime import datetime, timedelta from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock @@ -9,7 +9,11 @@ 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.constants import ( + SPEND_LOG_KEY_METADATA_CACHE_TTL, + SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL, + SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS, +) from litellm.proxy.spend_tracking.key_metadata_recovery import ( fill_missing_api_key_aliases, recover_double_hashed_key_metadata, @@ -31,6 +35,26 @@ def _query_raw_spend_logs(rows: Sequence[dict[str, str | None]]) -> AsyncMock: return AsyncMock(side_effect=query_raw) +def _spend_log_row(digest: str, key_alias: str | None, team_id: str | None, user_id: str | None) -> dict[str, str | None]: + return { + "digest": digest, + "first_alias": key_alias, + "last_alias": key_alias, + "first_team": team_id, + "last_team": team_id, + "first_owner": user_id, + "last_owner": user_id, + } + + +def _spend_log_transaction(mock_prisma: MagicMock, query_raw: AsyncMock) -> AsyncMock: + transaction = MagicMock() + transaction.execute_raw = AsyncMock(return_value=0) + transaction.query_raw = query_raw + mock_prisma.db.tx.return_value.__aenter__.return_value = transaction + return query_raw + + def _query_raw_by_table( active_rows: Sequence[dict[str, str | None]], deleted_rows: Sequence[dict[str, str | None]], @@ -240,15 +264,18 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_ session_digest = hash_token("cli-session-repro-user-6852") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs( - [_digest_row(session_digest, "cli-session-repro-user-6852", None, "repro-user-6852")] + query_raw = _spend_log_transaction( + mock_prisma, + _query_raw_spend_logs( + [_spend_log_row(session_digest, "cli-session-repro-user-6852", None, "repro-user-6852")] + ), ) result = await recover_key_metadata_from_spend_logs(mock_prisma, {session_digest}, window, cache=InMemoryCache()) assert result[session_digest]["key_alias"] == "cli-session-repro-user-6852" assert result[session_digest]["user_id"] == "repro-user-6852" - ((_, digests, start, end),) = [call.args for call in mock_prisma.db.query_raw.call_args_list] + ((_, digests, start, end),) = [call.args for call in query_raw.call_args_list] assert digests == [session_digest] assert (start, end) == window @@ -257,19 +284,19 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_ async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_keys(): window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = AsyncMock(return_value=[]) + query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[])) result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window, cache=InMemoryCache()) assert result == {} - mock_prisma.db.query_raw.assert_not_called() + query_raw.assert_not_called() @pytest.mark.asyncio async def test_recover_key_metadata_from_spend_logs_returns_empty_on_prisma_error(): window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down")) + query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("db down"))) result = await recover_key_metadata_from_spend_logs( mock_prisma, {hash_token("cli-session-x")}, window, cache=InMemoryCache() @@ -285,12 +312,15 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null foreign = hash_token("cli-session-foreign") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs( - [ - _digest_row(wanted, "kept-alias", None, "owner-1"), - _digest_row(all_null, None, None, None), - _digest_row(foreign, "foreign-alias", None, "owner-2"), - ] + query_raw = _spend_log_transaction( + mock_prisma, + _query_raw_spend_logs( + [ + _spend_log_row(wanted, "kept-alias", None, "owner-1"), + _spend_log_row(all_null, None, None, None), + _spend_log_row(foreign, "foreign-alias", None, "owner-2"), + ] + ), ) result = await recover_key_metadata_from_spend_logs(mock_prisma, {wanted, all_null}, window, cache=InMemoryCache()) @@ -304,14 +334,14 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null async def test_recover_key_metadata_from_spend_logs_skips_non_sha256_keys(): window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = AsyncMock(return_value=[]) + query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[])) result = await recover_key_metadata_from_spend_logs( mock_prisma, {"cli-session-raw-1798", "key-hash-short"}, window, cache=InMemoryCache() ) assert result == {} - mock_prisma.db.query_raw.assert_not_called() + query_raw.assert_not_called() @pytest.mark.asyncio @@ -319,13 +349,13 @@ async def test_recover_key_metadata_from_spend_logs_accepts_hashed_jwt_digests() jwt_digest = f"hashed-jwt-{hash_token('jwt-subject-1')}" window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(jwt_digest, None, "team-jwt", "jwt-user")]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(jwt_digest, None, "team-jwt", "jwt-user")])) result = await recover_key_metadata_from_spend_logs(mock_prisma, {jwt_digest}, window, cache=InMemoryCache()) assert result[jwt_digest]["team_id"] == "team-jwt" assert result[jwt_digest]["user_id"] == "jwt-user" - ((_, digests, _, _),) = [call.args for call in mock_prisma.db.query_raw.call_args_list] + ((_, digests, _, _),) = [call.args for call in query_raw.call_args_list] assert digests == [jwt_digest] @@ -336,7 +366,7 @@ async def test_recover_key_metadata_from_spend_logs_serves_repeat_lookups_from_t window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache() mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(found, "found-alias", None, "owner-1")]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, "owner-1")])) first = await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache) second = await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache) @@ -344,7 +374,7 @@ async def test_recover_key_metadata_from_spend_logs_serves_repeat_lookups_from_t assert first == second assert set(first) == {found} assert first[found]["key_alias"] == "found-alias" - assert mock_prisma.db.query_raw.await_count == 1 + assert query_raw.await_count == 1 @pytest.mark.asyncio @@ -354,15 +384,15 @@ async def test_recover_key_metadata_from_spend_logs_only_queries_digests_the_cac window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache() mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(cached_digest, "cached-alias", None, None)]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(cached_digest, "cached-alias", None, None)])) await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest}, window, cache=cache) - mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(new_digest, "new-alias", None, None)]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(new_digest, "new-alias", None, None)])) result = await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest, new_digest}, window, cache=cache) assert result[cached_digest]["key_alias"] == "cached-alias" assert result[new_digest]["key_alias"] == "new-alias" - ((_, digests, _, _),) = [call.args for call in mock_prisma.db.query_raw.call_args_list] + ((_, digests, _, _),) = [call.args for call in query_raw.call_args_list] assert digests == [new_digest] @@ -371,18 +401,18 @@ async def test_recover_key_metadata_from_spend_logs_rescans_when_the_window_chan digest = hash_token("cli-session-windowed") cache = InMemoryCache() mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs([]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([])) await recover_key_metadata_from_spend_logs( mock_prisma, {digest}, (datetime(2026, 9, 1), datetime(2026, 9, 4)), cache=cache ) - mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(digest, "later-alias", None, None)]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, "later-alias", None, None)])) result = await recover_key_metadata_from_spend_logs( mock_prisma, {digest}, (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=cache ) assert result[digest]["key_alias"] == "later-alias" - assert mock_prisma.db.query_raw.await_count == 1 + assert query_raw.await_count == 1 @pytest.mark.asyncio @@ -391,13 +421,13 @@ async def test_recover_key_metadata_from_spend_logs_retries_a_failed_query_only_ 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 = AsyncMock(side_effect=PrismaError("statement timeout")) + query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("statement timeout"))) started = time.time() assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {} - mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(digest, "back-online", None, None)]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, "back-online", None, None)])) assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {} - mock_prisma.db.query_raw.assert_not_awaited() + query_raw.assert_not_awaited() miss_key = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before")) assert cache.ttl_dict[miss_key] - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1 cache.ttl_dict[miss_key] = time.time() - 1 @@ -412,14 +442,11 @@ async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_ shared_ui_digest = hash_token("ui-token") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs( - [ - { - **_digest_row(shared_ui_digest, "ui-token", "litellm-dashboard", "bob"), - "first_owner": "alice", - "last_owner": "bob", - } - ] + query_raw = _spend_log_transaction( + mock_prisma, + _query_raw_spend_logs( + [{**_spend_log_row(shared_ui_digest, "ui-token", "litellm-dashboard", None), "first_owner": "alice", "last_owner": "bob"}] + ), ) result = await recover_key_metadata_from_spend_logs( @@ -434,8 +461,11 @@ async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_when_every_n digest = hash_token("cli-session-one-owner") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() - mock_prisma.db.query_raw = _query_raw_spend_logs( - [{**_digest_row(digest, None, None, "carol"), "first_owner": "carol", "last_owner": "carol"}] + query_raw = _spend_log_transaction( + mock_prisma, + _query_raw_spend_logs( + [_spend_log_row(digest, None, None, "carol")] + ), ) result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache()) @@ -450,7 +480,7 @@ async def test_recover_key_metadata_from_spend_logs_forgets_a_miss_long_before_a 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)]) + query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, None)])) started = time.time() await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache) @@ -471,9 +501,9 @@ async def test_recover_key_metadata_from_spend_logs_runs_one_query_for_concurren 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)] + return [_spend_log_row(digest, "shared-alias", None, None)] - mock_prisma.db.query_raw = AsyncMock(side_effect=slow_query_raw) + query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=slow_query_raw)) results = await asyncio.gather( *( @@ -483,7 +513,7 @@ async def test_recover_key_metadata_from_spend_logs_runs_one_query_for_concurren ) assert all(result[digest]["key_alias"] == "shared-alias" for result in results) - assert mock_prisma.db.query_raw.await_count == 1 + assert query_raw.await_count == 1 @pytest.mark.asyncio @@ -492,7 +522,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_a_repeated_miss_as_lon 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([]) + query_raw = _spend_log_transaction(mock_prisma, _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 @@ -500,5 +530,61 @@ async def test_recover_key_metadata_from_spend_logs_keeps_a_repeated_miss_as_lon await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache) - assert mock_prisma.db.query_raw.await_count == 2 + assert query_raw.await_count == 2 assert cache.ttl_dict[first_miss_key] - started >= SPEND_LOG_KEY_METADATA_CACHE_TTL - 1 + + +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_older_rows_agree_on_when_the_newest_is_nameless(): + digest = hash_token("cli-session-owner-from-older-rows") + window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) + mock_prisma = MagicMock() + _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, None, "team-x", "alice")])) + + result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache()) + + assert result[digest] == {"key_alias": None, "team_id": "team-x", "user_id": "alice"} + + +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_names_nothing_for_a_field_whose_rows_disagree(): + digest = hash_token("cli-session-disagreeing-rows") + window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) + mock_prisma = MagicMock() + _spend_log_transaction( + mock_prisma, + _query_raw_spend_logs( + [ + { + **_spend_log_row(digest, None, None, "carol"), + "first_alias": "old-alias", + "last_alias": "renamed-alias", + "first_team": "team-a", + "last_team": "team-b", + } + ] + ), + ) + + result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache()) + + assert result[digest] == {"key_alias": None, "team_id": None, "user_id": "carol"} + + +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_statement_timeout(): + digest = hash_token("cli-session-bounded-scan") + window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) + mock_prisma = MagicMock() + calls: list[str] = [] + transaction = MagicMock() + transaction.execute_raw = AsyncMock(side_effect=lambda sql: calls.append(sql) or 0) + transaction.query_raw = AsyncMock(side_effect=lambda sql, *args: calls.append("scan") or []) + mock_prisma.db.tx.return_value.__aenter__.return_value = transaction + + await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache()) + + assert calls == [f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}", "scan"] + assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta( + milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS + )