From 457f65eff933101fd38437e069d098d751583462 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 2 Jun 2026 17:45:28 -0700 Subject: [PATCH] test(proxy/utils): pin PrismaClient and spend-update behavior (#29488) * test(proxy/utils): pin PrismaClient and spend-update behavior PR2 of the litellm/proxy/utils.py behavior-pinning plan (https://www.notion.so/37343b8acdab81f68f39f66915f62bcf). Adds tests/test_litellm/proxy/utils/prisma_and_spend/, with happy + error pins for every symbol in the PR2 list: the config-param cache, PrismaClient lifecycle/data ops/engine watcher/reconnect/health clusters, the user-row cache and SMTP helper, password/token helpers, ProxyUpdateSpend, and the module-level spend functions. Tests run against fully-mocked Prisma stacks (patched ``Prisma`` / ``PrismaWrapper`` at fixture setup), with a fake SMTP transport and a clock-driven asyncio.sleep for the monitor loop, so unit runs need no DB or network. ``_pin_check.py`` enforces happy + error coverage for every symbol; ``_coverage_check.py`` filters branch + line coverage to the PR2 source range (lines 2,668-5,541) and prints PASS / FAIL with no numbers. Workflow shard ``tests/test_litellm/proxy/utils`` is added to the existing proxy-endpoints job. * test(proxy/utils): commit pin list and drop dead exclusion line Addresses Greptile review feedback on PR #29488: - Check in ``.pin_list.txt`` (force-added, overriding the repo-wide ``.gitignore`` rule) so reviewers can reproduce the ``_pin_check.py`` PASS shown in the PR description without first regenerating the file from Notion. - Remove the unreachable ``_harness_smoke_test.py`` continue in ``_pin_check.py``: the surrounding ``test_*.py`` glob already excludes underscore-prefixed files; rephrase the docstring instead. * test(proxy/utils): shift PR2 coverage line range by +1 after merge ``litellm_internal_staging`` added one line in ``ProxyLogging`` at ``utils.py:645`` (PR1 territory, before the PR2 region). Bump the ``_PR2_LINE_START`` / ``_PR2_LINE_END`` constants accordingly so the coverage gate keeps scoring the same source region after the merge. * test(proxy/utils): drop committed pin-list and gate scripts ``_pin_check.py``, ``_coverage_check.py``, and ``.pin_list.txt`` are local-only stopping signals: no workflow or pytest collection invokes them, so committing them adds rot risk (line-range drift in the coverage check, pin-list staleness) without any enforcement upside. The pin-list contract lives in the Notion plan; the tests themselves are the durable artifact. --------- Co-authored-by: Claude --- .../proxy/utils/prisma_and_spend/__init__.py | 0 .../prisma_and_spend/_harness_smoke_test.py | 84 +++ .../proxy/utils/prisma_and_spend/conftest.py | 387 +++++++++++++ .../prisma_and_spend/test_cache_user_row.py | 81 +++ .../test_config_param_cache.py | 267 +++++++++ .../prisma_and_spend/test_password_helpers.py | 223 ++++++++ .../test_prisma_client_engine_watcher.py | 521 ++++++++++++++++++ .../test_prisma_client_get_data.py | 400 ++++++++++++++ .../test_prisma_client_health.py | 292 ++++++++++ .../test_prisma_client_lifecycle.py | 207 +++++++ .../test_prisma_client_reconnect.py | 371 +++++++++++++ .../test_prisma_client_writes.py | 260 +++++++++ .../test_proxy_update_spend.py | 275 +++++++++ .../utils/prisma_and_spend/test_send_email.py | 105 ++++ .../prisma_and_spend/test_spend_functions.py | 360 ++++++++++++ 15 files changed, 3833 insertions(+) create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/__init__.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py create mode 100644 tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/__init__.py b/tests/test_litellm/proxy/utils/prisma_and_spend/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py b/tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py new file mode 100644 index 00000000000..2243d46ae7f --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/_harness_smoke_test.py @@ -0,0 +1,84 @@ +"""Self-tests for the prisma_and_spend test harness fixtures. + +Verifies the fixtures themselves do what their docstrings claim. +""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +def test_normalize_scrubs_volatile_keys() -> None: + from tests.test_litellm.proxy.utils.prisma_and_spend.conftest import normalize + + out = normalize({"id": 1, "spend": 2.0, "team_id": "t1"}) + assert out == {"id": "", "spend": "", "team_id": "t1"} + + +def test_normalize_recurses_into_lists() -> None: + from tests.test_litellm.proxy.utils.prisma_and_spend.conftest import normalize + + out = normalize([{"id": "x"}, {"team_id": "t"}]) + assert out == [{"id": ""}, {"team_id": "t"}] + + +def test_mock_prisma_client_has_common_tables(mock_prisma_client: Any) -> None: + for table in ( + "litellm_verificationtoken", + "litellm_teamtable", + "litellm_usertable", + "litellm_spendlogs", + "litellm_config", + "litellm_healthchecktable", + ): + assert hasattr(mock_prisma_client.db, table) + + +@pytest.mark.asyncio +async def test_mock_dual_cache_round_trip(mock_dual_cache: Any) -> None: + await mock_dual_cache.async_set_cache("k", "v") + assert await mock_dual_cache.async_get_cache("k") == "v" + await mock_dual_cache.async_delete_cache("k") + assert await mock_dual_cache.async_get_cache("k") is None + + +def test_prisma_client_fixture_is_a_real_prismaclient( + prisma_client: PrismaClient, +) -> None: + assert isinstance(prisma_client, PrismaClient) + assert callable(prisma_client.hash_token) + + +@pytest.mark.asyncio +async def test_fake_clock_advances(fake_clock: Any) -> None: + start = fake_clock.now + await asyncio.sleep(2.5) + assert fake_clock.now == start + 2.5 + assert fake_clock.sleep_calls == [2.5] + + +def test_make_spend_log_row_factory(make_spend_log_row: Any) -> None: + row = make_spend_log_row(request_id="abc", spend=0.5) + assert row["request_id"] == "abc" + assert row["spend"] == 0.5 + + +@pytest.mark.asyncio +async def test_in_memory_smtp_captures(in_memory_smtp: Any) -> None: + factory = in_memory_smtp.server_factory() + conn = factory("smtp.invalid", 25) + with conn: + conn.starttls() + from email.message import EmailMessage + + m = EmailMessage() + m["Subject"] = "S" + m.set_content("

x

", subtype="html") + conn.send_message(m, from_addr="a@b", to_addrs="c@d") + assert len(in_memory_smtp.sent) == 1 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py new file mode 100644 index 00000000000..2305a88b6dd --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py @@ -0,0 +1,387 @@ +"""Shared fixtures for tests/test_litellm/proxy/utils/prisma_and_spend/. + +All fixtures used by PR2 test files live here. Do NOT add fixtures inside +individual test files; if a fixture is missing, add it here and update the +Notion plan. + +The PrismaClient is exercised against a fully-mocked Prisma stack: the +``prisma.Prisma`` constructor and the writer/reader wrappers are patched +before PrismaClient.__init__ runs so the init code paths execute without +needing a generated Prisma client or a real database. +""" + +from __future__ import annotations + +import asyncio +import sys +from dataclasses import dataclass, field +from email.message import EmailMessage +from pathlib import Path +from typing import Any, Callable, Dict, Iterator, List, Optional +from unittest.mock import AsyncMock, MagicMock + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[5])) + + +VOLATILE_KEYS = frozenset( + { + "created_at", + "updated_at", + "checked_at", + "started_at", + "request_id", + "id", + "token", + "expires", + "expires_at", + "litellm_call_id", + "created", + "spend", + "last_refreshed_at", + "startTime", + "endTime", + "salt", + } +) + + +def normalize(data: Any, volatile: frozenset = VOLATILE_KEYS) -> Any: + """Recursively replace values for volatile keys with ''.""" + if isinstance(data, dict): + return { + k: ("" if k in volatile else normalize(v, volatile)) + for k, v in data.items() + } + if isinstance(data, list): + return [normalize(v, volatile) for v in data] + return data + + +_PRISMA_TABLES: List[str] = [ + "litellm_verificationtoken", + "litellm_teamtable", + "litellm_usertable", + "litellm_endusertable", + "litellm_organizationtable", + "litellm_proxymodeltable", + "litellm_modeltable", + "litellm_budgettable", + "litellm_spendlogs", + "litellm_config", + "litellm_usernotifications", + "litellm_healthchecktable", + "litellm_dailyuserspend", + "litellm_dailyteamspend", + "litellm_dailytagspend", + "litellm_managed_object_table", + "litellm_credentialstable", + "litellm_mcpservertable", + "litellm_audit_log", + "litellm_invitationlink", + "litellm_session_token_table", + "litellm_passthrough_endpoint_table", + "litellm_cron_job", + "litellm_passthrough_logs", + "litellm_promptstable", + "litellm_guardrailstable", + "litellm_managed_files", + "litellm_mcpusercredentials", + "litellm_objectpermissiontable", + "litellm_organizationmembership", +] + + +def _make_table_mock() -> MagicMock: + table = MagicMock() + table.find_unique = AsyncMock(return_value=None) + table.find_many = AsyncMock(return_value=[]) + table.find_first = AsyncMock(return_value=None) + table.create = AsyncMock() + table.create_many = AsyncMock() + table.update = AsyncMock() + table.update_many = AsyncMock() + table.upsert = AsyncMock() + table.delete = AsyncMock() + table.delete_many = AsyncMock() + table.count = AsyncMock(return_value=0) + table.group_by = AsyncMock(return_value=[]) + table.aggregate = AsyncMock(return_value={}) + return table + + +@pytest.fixture +def mock_prisma_client() -> MagicMock: + """Bare ``db`` mock with all common LiteLLM_* tables stubbed. + + Override individual return values in a test:: + + mock_prisma_client.db.litellm_usertable.find_unique.return_value = user + """ + client = MagicMock(name="MockPrismaClient") + client.db = MagicMock(name="MockPrismaDB") + client.connect = AsyncMock() + client.disconnect = AsyncMock() + client.health_check = AsyncMock(return_value=[{"?column?": 1}]) + client.proxy_logging_obj = MagicMock() + client.proxy_logging_obj.failure_handler = AsyncMock() + client.spend_log_transactions = [] + client._spend_log_transactions_lock = asyncio.Lock() + client.jsonify_object = lambda data: dict(data) + client.db.is_connected = MagicMock(return_value=False) + client.db.connect = AsyncMock() + client.db.disconnect = AsyncMock() + client.db.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + client.db.execute_raw = AsyncMock() + client.db.tx = MagicMock() + client.db.batch_ = MagicMock() + for table_name in _PRISMA_TABLES: + setattr(client.db, table_name, _make_table_mock()) + return client + + +@pytest.fixture +def mock_dual_cache() -> MagicMock: + """In-memory DualCache stand-in. + + Sync and async get/set wired against a private dict. Override or read + ``cache._store`` directly in a test for assertion convenience. + """ + cache = MagicMock(name="MockDualCache") + cache._store: Dict[str, Any] = {} + + def _sync_get(key: str, **_: Any) -> Any: + return cache._store.get(key) + + def _sync_set(key: str, value: Any, **_: Any) -> None: + cache._store[key] = value + + async def _async_get(key: str, **_: Any) -> Any: + return cache._store.get(key) + + async def _async_set(key: str, value: Any, **_: Any) -> None: + cache._store[key] = value + + async def _async_delete(key: str, **_: Any) -> None: + cache._store.pop(key, None) + + cache.get_cache = MagicMock(side_effect=_sync_get) + cache.set_cache = MagicMock(side_effect=_sync_set) + cache.async_get_cache = AsyncMock(side_effect=_async_get) + cache.async_set_cache = AsyncMock(side_effect=_async_set) + cache.async_delete_cache = AsyncMock(side_effect=_async_delete) + return cache + + +@pytest.fixture +def patched_prisma_import(monkeypatch: pytest.MonkeyPatch) -> Iterator[MagicMock]: + """Replace ``prisma.Prisma`` and ``PrismaWrapper`` so PrismaClient.__init__ + runs without a generated client. Yields the fake Prisma instance. + + ``prisma`` raises RuntimeError (not AttributeError) for the missing + ``Prisma`` attribute, so ``monkeypatch.setattr`` can't probe it; assign + directly and restore in teardown. + """ + import prisma as _prisma_pkg + import litellm.proxy.utils as _utils_mod + + fake_prisma = MagicMock(name="FakePrisma") + fake_prisma.is_connected = MagicMock(return_value=False) + fake_prisma.connect = AsyncMock() + fake_prisma.disconnect = AsyncMock() + + fake_prisma_factory = MagicMock(name="FakePrismaFactory", return_value=fake_prisma) + had_prisma_attr = "Prisma" in _prisma_pkg.__dict__ + previous_prisma_attr = _prisma_pkg.__dict__.get("Prisma") + _prisma_pkg.Prisma = fake_prisma_factory # type: ignore[attr-defined] + + fake_wrapper = MagicMock(name="FakePrismaWrapper") + fake_wrapper.is_connected = MagicMock(return_value=False) + fake_wrapper.connect = AsyncMock() + fake_wrapper.disconnect = AsyncMock() + fake_wrapper.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + + def _fake_wrapper_ctor(*args: Any, **kwargs: Any) -> MagicMock: + return fake_wrapper + + monkeypatch.setattr(_utils_mod, "PrismaWrapper", _fake_wrapper_ctor) + fake_prisma.__wrapper__ = fake_wrapper + try: + yield fake_prisma + finally: + if had_prisma_attr: + _prisma_pkg.Prisma = previous_prisma_attr # type: ignore[attr-defined] + else: + try: + del _prisma_pkg.Prisma # type: ignore[attr-defined] + except AttributeError: + pass + + +@pytest.fixture +def prisma_client( + patched_prisma_import: MagicMock, + mock_prisma_client: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> Any: + """Wired ``PrismaClient`` whose ``db`` attribute is the table mock. + + The init runs through the real code path (testing the constructor's + config-attribute setup) and is then snapped to the easier-to-assert + table mock for downstream behavior pinning. + """ + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + from litellm.proxy.utils import PrismaClient + + proxy_logging_obj = MagicMock(name="MockProxyLogging") + proxy_logging_obj.failure_handler = AsyncMock() + pc = PrismaClient( + database_url="postgresql://test:test@localhost:5432/test", + proxy_logging_obj=proxy_logging_obj, + ) + pc.db = mock_prisma_client.db + return pc + + +@dataclass +class FakeClock: + """Monotonic-time controller for the spend monitor loop. + + Tests advance time via ``clock.advance(seconds)`` while asyncio.sleep + is replaced with a clock-driven no-op. + """ + + now: float = 0.0 + sleep_calls: List[float] = field(default_factory=list) + + def advance(self, seconds: float) -> None: + self.now += seconds + + def time(self) -> float: + return self.now + + async def sleep(self, seconds: float) -> None: + self.sleep_calls.append(seconds) + self.now += seconds + + +@pytest.fixture +def fake_clock(monkeypatch: pytest.MonkeyPatch) -> FakeClock: + """Install a controllable clock + asyncio.sleep replacement.""" + clock = FakeClock() + monkeypatch.setattr("time.time", clock.time) + monkeypatch.setattr("time.monotonic", clock.time) + + async def _fast_sleep(seconds: float, *_: Any, **__: Any) -> None: + clock.sleep_calls.append(seconds) + clock.now += seconds + + monkeypatch.setattr("asyncio.sleep", _fast_sleep) + return clock + + +@pytest.fixture +def make_spend_log_row() -> Callable[..., Dict[str, Any]]: + """Factory for fake LiteLLM_SpendLogs rows.""" + + def _make( + request_id: str = "req-1", + spend: float = 0.01, + model: str = "gpt-4o-mini", + **overrides: Any, + ) -> Dict[str, Any]: + row = { + "request_id": request_id, + "spend": spend, + "model": model, + "user": "user-1", + "team_id": "team-1", + "api_key": "hashed-key", + "startTime": "2026-06-02T00:00:00Z", + "endTime": "2026-06-02T00:00:01Z", + "metadata": {}, + } + row.update(overrides) + return row + + return _make + + +@dataclass +class _SentMessage: + from_addr: Optional[str] + to_addrs: Any + subject: Optional[str] + body: Optional[str] + starttls_called: bool + login_args: Optional[tuple] + + +@dataclass +class InMemorySMTP: + """Captures outbound SMTP traffic for ``send_email`` tests.""" + + sent: List[_SentMessage] = field(default_factory=list) + raise_on_send: Optional[Exception] = None + + def server_factory(self) -> Callable[..., Any]: + outer = self + + class _Conn: + def __init__(self) -> None: + self._starttls_called = False + self._login_args: Optional[tuple] = None + + def __enter__(self) -> "_Conn": + return self + + def __exit__(self, *exc: Any) -> None: + return None + + def starttls(self) -> None: + self._starttls_called = True + + def login(self, user: str, password: str) -> None: + self._login_args = (user, password) + + def send_message( + self, + msg: EmailMessage, + from_addr: Optional[str] = None, + to_addrs: Any = None, + ) -> None: + if outer.raise_on_send is not None: + raise outer.raise_on_send + body = "" + for part in msg.walk(): + if part.get_content_type() == "text/html": + body = part.get_payload(decode=False) or "" + break + outer.sent.append( + _SentMessage( + from_addr=from_addr, + to_addrs=to_addrs, + subject=msg["Subject"], + body=body, + starttls_called=self._starttls_called, + login_args=self._login_args, + ) + ) + + def _factory(*args: Any, **kwargs: Any) -> _Conn: + return _Conn() + + return _factory + + +@pytest.fixture +def in_memory_smtp(monkeypatch: pytest.MonkeyPatch) -> InMemorySMTP: + """Patch ``smtplib.SMTP`` to capture sends in memory. + + Override ``smtp.raise_on_send`` to test the SMTP error path. + """ + smtp = InMemorySMTP() + monkeypatch.setattr("smtplib.SMTP", smtp.server_factory()) + return smtp diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py new file mode 100644 index 00000000000..d1270b60b19 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_cache_user_row.py @@ -0,0 +1,81 @@ +"""Pin ``_cache_user_row``. + +Symbols pinned here: + - ``_cache_user_row`` +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import _cache_user_row + + +@pytest.mark.asyncio +async def test_cache_user_row_caches_on_miss( + mock_dual_cache: Any, +) -> None: + user_row = SimpleNamespace( + user_id="u1", spend=2.5, max_budget=10.0, name="Alice" + ) + user_row.model_dump_json = MagicMock( + return_value='{"user_id":"u1","spend":2.5,"max_budget":10.0,"name":"Alice"}' + ) + db = MagicMock() + db.get_data = AsyncMock(return_value=user_row) + + result = await _cache_user_row("u1", mock_dual_cache, db) + cache_key = "u1_user_api_key_user_id" + pinned = { + "result": result, + "cache_value": mock_dual_cache._store[cache_key], + "get_calls": mock_dual_cache.get_cache.call_count, + "set_calls": mock_dual_cache.set_cache.call_count, + "db_called": db.get_data.await_count, + } + assert pinned == { + "result": None, + "cache_value": '{"user_id":"u1","spend":2.5,"max_budget":10.0,"name":"Alice"}', + "get_calls": 1, + "set_calls": 1, + "db_called": 1, + } + + +@pytest.mark.asyncio +async def test_cache_user_row_skips_db_on_cache_hit( + mock_dual_cache: Any, +) -> None: + cache_key = "u-hit_user_api_key_user_id" + mock_dual_cache._store[cache_key] = "cached-blob" + db = MagicMock() + db.get_data = AsyncMock(return_value=None) + result = await _cache_user_row("u-hit", mock_dual_cache, db) + assert result is None + assert db.get_data.await_count == 0 + + +@pytest.mark.asyncio +async def test_cache_user_row_skips_set_when_user_row_lacks_model_dump_json( + mock_dual_cache: Any, +) -> None: + user_row = SimpleNamespace(user_id="u2", spend=1.0) + db = MagicMock() + db.get_data = AsyncMock(return_value=user_row) + await _cache_user_row("u2", mock_dual_cache, db) + assert mock_dual_cache._store == {} + assert mock_dual_cache.set_cache.call_count == 0 + + +@pytest.mark.asyncio +async def test_cache_user_row_propagates_db_error( + mock_dual_cache: Any, +) -> None: + db = MagicMock() + db.get_data = AsyncMock(side_effect=RuntimeError("db down")) + with pytest.raises(RuntimeError, match="db down"): + await _cache_user_row("u3", mock_dual_cache, db) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py new file mode 100644 index 00000000000..761835078f4 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py @@ -0,0 +1,267 @@ +"""Pin the LiteLLM_Config cached-read layer. + +Symbols pinned here: + - ``_ConfigRow`` + - ``_config_cache_key`` + - ``_pack_config_row`` + - ``_unpack_config_row`` + - ``get_config_param`` + - ``invalidate_config_param`` + - ``prefetch_config_params`` +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm.proxy.utils as utils_mod +from litellm.proxy.utils import ( + _config_cache_key, + _ConfigRow, + _pack_config_row, + _unpack_config_row, + get_config_param, + invalidate_config_param, + prefetch_config_params, +) + + +@pytest.fixture(autouse=True) +def _swap_config_cache( + monkeypatch: pytest.MonkeyPatch, mock_dual_cache: Any +) -> Any: + """Replace the module-level cache so tests see a clean store per run.""" + monkeypatch.setattr(utils_mod, "litellm_config_cache", mock_dual_cache) + return mock_dual_cache + + +def test_config_cache_key_uses_documented_prefix() -> None: + actual = { + "key": _config_cache_key("max_budget"), + "another": _config_cache_key("disable_spend_updates"), + "prefix": _config_cache_key("x").split(":")[0], + } + assert actual == { + "key": "litellm_config:param:max_budget", + "another": "litellm_config:param:disable_spend_updates", + "prefix": "litellm_config", + } + + +def test_config_cache_key_error_propagates_from_bad_format() -> None: + class _Boom: + def __format__(self, _spec: str) -> str: + raise ValueError("format failure") + + with pytest.raises(ValueError, match="format failure"): + _config_cache_key(_Boom()) # type: ignore[arg-type] + + +def test_config_row_dataclass_shape() -> None: + row = _ConfigRow(param_name="alpha", param_value={"k": 1}) + assert { + "param_name": row.param_name, + "param_value": row.param_value, + "slots": _ConfigRow.__slots__, + } == { + "param_name": "alpha", + "param_value": {"k": 1}, + "slots": ("param_name", "param_value"), + } + + +def test_config_row_rejects_unknown_attribute() -> None: + row = _ConfigRow("a", 1) + with pytest.raises(AttributeError): + row.something_else = 2 # type: ignore[attr-defined] + + +def test_pack_config_row_returns_dict_for_caching() -> None: + row = SimpleNamespace(param_name="zeta", param_value=[1, 2, 3]) + actual = _pack_config_row(row) + expanded = {**actual, "is_dict": isinstance(actual, dict)} + assert expanded == { + "param_name": "zeta", + "param_value": [1, 2, 3], + "is_dict": True, + } + + +def test_pack_config_row_error_on_missing_attribute() -> None: + bad = SimpleNamespace(param_name="only_name") + with pytest.raises(AttributeError): + _pack_config_row(bad) + + +def test_unpack_config_row_round_trips_dict() -> None: + packed = {"param_name": "alpha", "param_value": "abc"} + unpacked = _unpack_config_row(packed) + assert isinstance(unpacked, _ConfigRow) + actual = { + "param_name": unpacked.param_name, + "param_value": unpacked.param_value, + "from_none": _unpack_config_row(None), + "from_miss_sentinel": _unpack_config_row(utils_mod._CONFIG_CACHE_MISS), + "from_other_type": _unpack_config_row(123), + } + assert actual == { + "param_name": "alpha", + "param_value": "abc", + "from_none": None, + "from_miss_sentinel": None, + "from_other_type": None, + } + + +def test_unpack_config_row_error_on_malformed_dict() -> None: + with pytest.raises(KeyError): + _unpack_config_row({"only_name": "x"}) + + +@pytest.mark.asyncio +async def test_get_config_param_cache_hit_returns_unpacked_row( + _swap_config_cache: Any, +) -> None: + cache_key = _config_cache_key("p1") + await _swap_config_cache.async_set_cache( + cache_key, {"param_name": "p1", "param_value": {"x": 1}} + ) + prisma = MagicMock() + prisma.get_generic_data = AsyncMock() + + row = await get_config_param(prisma, "p1") + actual = { + "type": type(row).__name__, + "param_name": row.param_name, + "param_value": row.param_value, + "db_not_touched": prisma.get_generic_data.await_count == 0, + } + assert actual == { + "type": "_ConfigRow", + "param_name": "p1", + "param_value": {"x": 1}, + "db_not_touched": True, + } + + +@pytest.mark.asyncio +async def test_get_config_param_cache_miss_fetches_from_db_and_caches( + _swap_config_cache: Any, +) -> None: + db_row = SimpleNamespace(param_name="p2", param_value={"y": 2}) + prisma = MagicMock() + prisma.get_generic_data = AsyncMock(return_value=db_row) + + row = await get_config_param(prisma, "p2") + cached = _swap_config_cache._store[_config_cache_key("p2")] + actual = { + "returned": row, + "cached": cached, + "db_called": prisma.get_generic_data.await_count, + "db_args": prisma.get_generic_data.await_args.kwargs, + } + assert actual == { + "returned": db_row, + "cached": {"param_name": "p2", "param_value": {"y": 2}}, + "db_called": 1, + "db_args": {"key": "param_name", "value": "p2", "table_name": "config"}, + } + + +@pytest.mark.asyncio +async def test_get_config_param_caches_negative_lookup_as_miss_sentinel( + _swap_config_cache: Any, +) -> None: + prisma = MagicMock() + prisma.get_generic_data = AsyncMock(return_value=None) + row = await get_config_param(prisma, "absent") + assert row is None + assert _swap_config_cache._store[_config_cache_key("absent")] == ( + utils_mod._CONFIG_CACHE_MISS + ) + + +@pytest.mark.asyncio +async def test_get_config_param_raises_when_db_raises() -> None: + prisma = MagicMock() + prisma.get_generic_data = AsyncMock(side_effect=RuntimeError("db down")) + with pytest.raises(RuntimeError, match="db down"): + await get_config_param(prisma, "p3") + + +@pytest.mark.asyncio +async def test_invalidate_config_param_evicts_from_cache( + _swap_config_cache: Any, +) -> None: + cache_key = _config_cache_key("p4") + await _swap_config_cache.async_set_cache(cache_key, {"param_name": "p4", "param_value": 1}) + await invalidate_config_param("p4") + actual = { + "store_empty": _swap_config_cache._store == {}, + "delete_calls": _swap_config_cache.async_delete_cache.await_count, + "delete_arg": _swap_config_cache.async_delete_cache.await_args.args[0], + } + assert actual == { + "store_empty": True, + "delete_calls": 1, + "delete_arg": "litellm_config:param:p4", + } + + +@pytest.mark.asyncio +async def test_invalidate_config_param_propagates_cache_error( + _swap_config_cache: Any, +) -> None: + _swap_config_cache.async_delete_cache = AsyncMock( + side_effect=ConnectionError("redis down") + ) + with pytest.raises(ConnectionError): + await invalidate_config_param("p5") + + +@pytest.mark.asyncio +async def test_prefetch_config_params_populates_cache_for_each_name( + _swap_config_cache: Any, +) -> None: + rows: List[SimpleNamespace] = [ + SimpleNamespace(param_name="a", param_value={"av": 1}), + SimpleNamespace(param_name="c", param_value=[3]), + ] + prisma = MagicMock() + prisma.db.litellm_config.find_many = AsyncMock(return_value=rows) + await prefetch_config_params(prisma, ["a", "b", "c"]) + actual = { + "a": _swap_config_cache._store[_config_cache_key("a")], + "b": _swap_config_cache._store[_config_cache_key("b")], + "c": _swap_config_cache._store[_config_cache_key("c")], + } + assert actual == { + "a": {"param_name": "a", "param_value": {"av": 1}}, + "b": utils_mod._CONFIG_CACHE_MISS, + "c": {"param_name": "c", "param_value": [3]}, + } + + +@pytest.mark.asyncio +async def test_prefetch_config_params_empty_list_is_noop( + _swap_config_cache: Any, +) -> None: + prisma = MagicMock() + prisma.db.litellm_config.find_many = AsyncMock(return_value=[]) + await prefetch_config_params(prisma, []) + assert prisma.db.litellm_config.find_many.await_count == 0 + assert _swap_config_cache._store == {} + + +@pytest.mark.asyncio +async def test_prefetch_config_params_swallows_db_error_without_caching( + _swap_config_cache: Any, +) -> None: + prisma = MagicMock() + prisma.db.litellm_config.find_many = AsyncMock(side_effect=RuntimeError("boom")) + await prefetch_config_params(prisma, ["a", "b"]) + assert _swap_config_cache._store == {} diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py new file mode 100644 index 00000000000..3c028473479 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py @@ -0,0 +1,223 @@ +"""Pin password/token helper behavior. + +Symbols pinned here: + - ``hash_token`` + - ``hash_password`` + - ``verify_password`` + - ``migrate_passwords_to_scrypt_async`` + - ``_hash_token_if_needed`` + - ``PrismaClient._is_sha256_hex`` (a nested helper inside + ``migrate_passwords_to_scrypt_async``; the pin list labels it under the + PrismaClient health cluster as a documentation artifact) +""" + +from __future__ import annotations + +import hashlib +from types import SimpleNamespace +from typing import List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import ( + _hash_token_if_needed, + hash_password, + hash_token, + migrate_passwords_to_scrypt_async, + verify_password, +) + + +def test_hash_token_returns_sha256_hex_of_input() -> None: + token = "sk-abcDEF12345" + result = hash_token(token) + expected = hashlib.sha256(token.encode()).hexdigest() + actual = { + "len": len(result), + "hex": all(c in "0123456789abcdef" for c in result), + "hash": result, + "matches_sha256": result == expected, + } + assert actual == { + "len": 64, + "hex": True, + "hash": expected, + "matches_sha256": True, + } + + +def test_hash_token_empty_string_still_hashes() -> None: + result = hash_token("") + assert result == hashlib.sha256(b"").hexdigest() + + +def test_hash_token_raises_for_non_string() -> None: + with pytest.raises(AttributeError): + hash_token(None) # type: ignore[arg-type] + + +def test_hash_password_uses_scrypt_prefix() -> None: + h = hash_password("hunter2") + fields = { + "prefix": h[:7], + "min_length": len(h) > 60, + "verifies_self": verify_password("hunter2", h), + "rejects_other": verify_password("hunter3", h), + } + assert fields == { + "prefix": "scrypt:", + "min_length": True, + "verifies_self": True, + "rejects_other": False, + } + + +def test_hash_password_returns_distinct_hashes_per_call() -> None: + a = hash_password("same-password") + b = hash_password("same-password") + assert a != b + assert verify_password("same-password", a) + assert verify_password("same-password", b) + + +def test_hash_password_error_for_non_string_raises() -> None: + with pytest.raises(AttributeError): + hash_password(None) # type: ignore[arg-type] + + +def test_verify_password_sha256_legacy_path() -> None: + plaintext = "legacy-pass" + sha = hashlib.sha256(plaintext.encode()).hexdigest() + matrix = { + "correct": verify_password(plaintext, sha), + "wrong": verify_password("other", sha), + "non_hex_short": verify_password(plaintext, "not-hex"), + "empty_stored": verify_password(plaintext, ""), + } + assert matrix == { + "correct": True, + "wrong": False, + "non_hex_short": False, + "empty_stored": False, + } + + +def test_verify_password_scrypt_malformed_returns_false() -> None: + assert verify_password("anything", "scrypt:not-base64") is False + + +def test_verify_password_unknown_format_returns_false() -> None: + assert verify_password("x", "plaintext-not-supported") is False + + +def test_hash_token_if_needed_handles_sk_prefix() -> None: + plain = "sk-secret-xyz" + already_hashed = hashlib.sha256(plain.encode()).hexdigest() + not_a_secret = "token-without-sk-prefix" + actual = { + "sk_input_is_hashed": _hash_token_if_needed(plain) == already_hashed, + "non_sk_passthrough": _hash_token_if_needed(not_a_secret) == not_a_secret, + "double_hash_stable": _hash_token_if_needed(already_hashed) == already_hashed, + } + assert actual == { + "sk_input_is_hashed": True, + "non_sk_passthrough": True, + "double_hash_stable": True, + } + + +def test_hash_token_if_needed_error_on_non_string() -> None: + with pytest.raises(AttributeError): + _hash_token_if_needed(None) # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- +# migrate_passwords_to_scrypt_async — pins behavior of the nested +# ``_is_sha256_hex`` helper too: scrypt-prefixed and sha256-hex rows are +# left alone, plaintext rows are upgraded in place. +# --------------------------------------------------------------------------- + + +def _make_user(user_id: str, password) -> SimpleNamespace: + return SimpleNamespace(user_id=user_id, password=password) + + +@pytest.mark.asyncio +async def test_migrate_passwords_skips_when_no_plaintext() -> None: + pc = MagicMock() + pc.db = MagicMock() + sha = hashlib.sha256(b"already-hashed").hexdigest() + pc.db.litellm_usertable.find_many = AsyncMock( + return_value=[ + _make_user("a", "scrypt:abc"), + _make_user("b", sha), + ] + ) + pc.db.litellm_usertable.update = AsyncMock() + + result = await migrate_passwords_to_scrypt_async(pc) + outcome = { + "message": result, + "updates": pc.db.litellm_usertable.update.await_count, + "find_called": pc.db.litellm_usertable.find_many.await_count, + "fetch_filter": pc.db.litellm_usertable.find_many.await_args.kwargs["where"], + } + assert outcome == { + "message": "No plaintext passwords found", + "updates": 0, + "find_called": 1, + "fetch_filter": {"password": {"not": None}}, + } + + +@pytest.mark.asyncio +async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: + pc = MagicMock() + pc.db = MagicMock() + users: List[SimpleNamespace] = [ + _make_user("plaintext-user-1", "plain-1"), + _make_user("plaintext-user-2", "plain-2"), + _make_user("scrypt-user", "scrypt:already"), + _make_user( + "sha-user", + hashlib.sha256(b"alreadyhashed").hexdigest(), + ), + _make_user("null-pw", None), + ] + pc.db.litellm_usertable.find_many = AsyncMock(return_value=users) + pc.db.litellm_usertable.update = AsyncMock() + + result = await migrate_passwords_to_scrypt_async(pc) + + updated_user_ids = sorted( + call.kwargs["where"]["user_id"] + for call in pc.db.litellm_usertable.update.await_args_list + ) + new_password_prefixes = sorted( + call.kwargs["data"]["password"][:7] + for call in pc.db.litellm_usertable.update.await_args_list + ) + outcome = { + "message": result, + "update_count": pc.db.litellm_usertable.update.await_count, + "updated_ids": updated_user_ids, + "all_scrypt_prefixed": new_password_prefixes, + } + assert outcome == { + "message": "Migrated 2 plaintext passwords to scrypt", + "update_count": 2, + "updated_ids": ["plaintext-user-1", "plaintext-user-2"], + "all_scrypt_prefixed": ["scrypt:", "scrypt:"], + } + + +@pytest.mark.asyncio +async def test_migrate_passwords_raises_on_db_failure() -> None: + pc = MagicMock() + pc.db = MagicMock() + pc.db.litellm_usertable.find_many = AsyncMock( + side_effect=RuntimeError("db unavailable") + ) + with pytest.raises(RuntimeError, match="db unavailable"): + await migrate_passwords_to_scrypt_async(pc) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py new file mode 100644 index 00000000000..7b862eecbd4 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_engine_watcher.py @@ -0,0 +1,521 @@ +"""Pin ``PrismaClient`` engine watcher methods. + +Symbols pinned here: + - ``PrismaClient._get_engine_pid`` + - ``PrismaClient._is_engine_alive`` + - ``PrismaClient._reap_all_zombies`` + - ``PrismaClient._try_waitpid_watch`` + - ``PrismaClient._waitpid_thread_func`` + - ``PrismaClient._on_engine_death_from_thread`` + - ``PrismaClient._try_pidfd_watch`` + - ``PrismaClient._on_pidfd_readable`` + - ``PrismaClient._poll_engine_proc`` + - ``PrismaClient._cleanup_engine_watcher`` + - ``PrismaClient._start_engine_watcher`` + - ``PrismaClient._stop_engine_watcher`` + +Linux-only tests are skipped on Windows; the production code uses +``waitpid``/``pidfd_open`` which are Unix-only. +""" + +from __future__ import annotations + +import asyncio +import os +import sys +import threading +from typing import Any, Optional +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +pytestmark = pytest.mark.skipif( + sys.platform == "win32", reason="engine watcher is Unix-only" +) + + +def test_get_engine_pid_extracts_process_pid(prisma_client: PrismaClient) -> None: + fake_engine = MagicMock() + fake_engine.process = MagicMock() + fake_engine.process.pid = 4242 + prisma_client.db._original_prisma = MagicMock() + prisma_client.db._original_prisma._engine = fake_engine + actual = { + "pid": prisma_client._get_engine_pid(), + "engine_attr": prisma_client.db._original_prisma._engine is fake_engine, + "process_pid": fake_engine.process.pid, + } + assert actual == {"pid": 4242, "engine_attr": True, "process_pid": 4242} + + +def test_get_engine_pid_returns_zero_when_engine_attr_missing( + prisma_client: PrismaClient, +) -> None: + prisma_client.db._original_prisma = MagicMock(spec=[]) + assert prisma_client._get_engine_pid() == 0 + + +def test_is_engine_alive_true_when_pid_zero(prisma_client: PrismaClient) -> None: + prisma_client._engine_pid = 0 + pinned = { + "result": prisma_client._is_engine_alive(), + "pid": prisma_client._engine_pid, + "type": type(prisma_client._is_engine_alive()).__name__, + } + assert pinned == {"result": True, "pid": 0, "type": "bool"} + + +def test_is_engine_alive_false_when_process_lookup_fails( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 99999 + monkeypatch.setattr( + "os.kill", MagicMock(side_effect=ProcessLookupError()) + ) + assert prisma_client._is_engine_alive() is False + + +def test_is_engine_alive_true_on_permission_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 1 + monkeypatch.setattr("os.kill", MagicMock(side_effect=PermissionError())) + assert prisma_client._is_engine_alive() is True + + +def test_reap_all_zombies_returns_set_of_reaped_pids( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls = iter([(111, 0), (222, 0), (0, 0)]) + + def fake_waitpid(pid: int, flags: int) -> Any: + return next(calls) + + monkeypatch.setattr("os.waitpid", fake_waitpid) + reaped = PrismaClient._reap_all_zombies() + pinned = { + "type": type(reaped).__name__, + "size": len(reaped), + "contains_111": 111 in reaped, + "contains_222": 222 in reaped, + } + assert pinned == {"type": "set", "size": 2, "contains_111": True, "contains_222": True} + + +def test_reap_all_zombies_handles_no_children_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "os.waitpid", MagicMock(side_effect=ChildProcessError()) + ) + assert PrismaClient._reap_all_zombies() == set() + + +@pytest.mark.asyncio +async def test_try_waitpid_watch_starts_thread_for_live_child( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(0, 0))) + + threads: list[threading.Thread] = [] + + real_thread_cls = threading.Thread + + def _capture_thread(*args: Any, **kwargs: Any) -> threading.Thread: + t = real_thread_cls(*args, **kwargs) + threads.append(t) + # Replace start so we don't actually launch the thread. + t.start = MagicMock() # type: ignore[method-assign] + return t + + monkeypatch.setattr("threading.Thread", _capture_thread) + monkeypatch.setattr(prisma_client, "_waitpid_thread_func", MagicMock()) + + result = prisma_client._try_waitpid_watch(7777) + pinned = { + "returned": result, + "threads_made": len(threads), + "wait_thread_set": prisma_client._engine_wait_thread is threads[0], + "thread_name_prefix": threads[0].name.startswith("prisma-engine-waitpid-"), + } + assert pinned == { + "returned": True, + "threads_made": 1, + "wait_thread_set": True, + "thread_name_prefix": True, + } + + +@pytest.mark.asyncio +async def test_try_waitpid_watch_returns_false_for_non_child( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + "os.waitpid", MagicMock(side_effect=ChildProcessError()) + ) + assert prisma_client._try_waitpid_watch(123) is False + + +@pytest.mark.asyncio +async def test_try_waitpid_watch_handles_already_dead_pid( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """If the engine PID is already dead at watch start, _try_waitpid_watch + returns True and schedules a reconnect. + """ + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(8888, 0))) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + monkeypatch.setattr(prisma_client, "_cleanup_engine_watcher", MagicMock()) + + result = prisma_client._try_waitpid_watch(8888) + # Drain pending tasks so attempt_db_reconnect is awaited and we don't leak. + await asyncio.sleep(0) + pinned = { + "result": result, + "engine_confirmed_dead": prisma_client._engine_confirmed_dead, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + "reconnect_scheduled": prisma_client.attempt_db_reconnect.await_count >= 1, + } + assert pinned == { + "result": True, + "engine_confirmed_dead": True, + "cleanup_called": 1, + "reconnect_scheduled": True, + } + + +def test_waitpid_thread_func_swallows_child_process_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(side_effect=ChildProcessError())) + loop = MagicMock() + loop.call_soon_threadsafe = MagicMock() + prisma_client._waitpid_thread_func(123, loop) + assert loop.call_soon_threadsafe.call_count == 1 + + +def test_waitpid_thread_func_invokes_on_engine_death_on_normal_exit( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(123, 0))) + loop = MagicMock() + received: list[Any] = [] + loop.call_soon_threadsafe = lambda fn, pid: received.append((fn, pid)) + prisma_client._waitpid_thread_func(123, loop) + pinned = { + "callbacks_received": len(received), + "callback_target": received[0][0] == prisma_client._on_engine_death_from_thread, + "pid_arg": received[0][1], + "first_tuple_size": len(received[0]), + } + assert pinned == { + "callbacks_received": 1, + "callback_target": True, + "pid_arg": 123, + "first_tuple_size": 2, + } + + +def test_waitpid_thread_func_swallows_loop_runtime_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.waitpid", MagicMock(return_value=(123, 0))) + loop = MagicMock() + loop.call_soon_threadsafe = MagicMock(side_effect=RuntimeError("loop closed")) + prisma_client._waitpid_thread_func(123, loop) + + +@pytest.mark.asyncio +async def test_on_engine_death_from_thread_schedules_reconnect( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 7777 + prisma_client._engine_confirmed_dead = False + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + monkeypatch.setattr(prisma_client, "_cleanup_engine_watcher", MagicMock()) + + prisma_client._on_engine_death_from_thread(7777) + await asyncio.sleep(0) + pinned = { + "confirmed_dead": prisma_client._engine_confirmed_dead, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + "reconnect_called": prisma_client.attempt_db_reconnect.await_count, + "reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs["reason"], + } + assert pinned == { + "confirmed_dead": True, + "cleanup_called": 1, + "reconnect_called": 1, + "reconnect_reason": "engine_process_death", + } + + +def test_on_engine_death_from_thread_ignores_wrong_pid_or_already_dead( + prisma_client: PrismaClient, +) -> None: + prisma_client._engine_pid = 1111 + prisma_client._engine_confirmed_dead = True + prisma_client._cleanup_engine_watcher = MagicMock() + prisma_client._on_engine_death_from_thread(1111) + assert prisma_client._cleanup_engine_watcher.call_count == 0 + + +def test_on_engine_death_from_thread_wrong_pid_does_nothing( + prisma_client: PrismaClient, +) -> None: + prisma_client._engine_pid = 1111 + prisma_client._engine_confirmed_dead = False + prisma_client._cleanup_engine_watcher = MagicMock() + prisma_client._on_engine_death_from_thread(2222) + assert prisma_client._cleanup_engine_watcher.call_count == 0 + assert prisma_client._engine_confirmed_dead is False + + +@pytest.mark.asyncio +async def test_try_pidfd_watch_returns_false_when_pidfd_open_missing( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delattr("os.pidfd_open", raising=False) + assert prisma_client._try_pidfd_watch(123) is False + + +@pytest.mark.asyncio +async def test_try_pidfd_watch_arms_reader_when_available( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + def fake_pidfd(pid: int, flags: int) -> int: + return 42 + + monkeypatch.setattr("os.pidfd_open", fake_pidfd, raising=False) + loop = asyncio.get_running_loop() + fake_add_reader = MagicMock() + monkeypatch.setattr(loop, "add_reader", fake_add_reader) + + result = prisma_client._try_pidfd_watch(123) + assert result is True + assert prisma_client._engine_pidfd == 42 + assert fake_add_reader.call_args.args[0] == 42 + + +@pytest.mark.asyncio +async def test_try_pidfd_watch_error_returns_false_and_cleans_up( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + def fake_pidfd(pid: int, flags: int) -> int: + raise OSError("ENOSYS") + + monkeypatch.setattr("os.pidfd_open", fake_pidfd, raising=False) + assert prisma_client._try_pidfd_watch(123) is False + assert prisma_client._engine_pidfd == -1 + + +@pytest.mark.asyncio +async def test_on_pidfd_readable_invokes_reconnect_path( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 4321 + prisma_client._engine_confirmed_dead = False + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + cleanup = MagicMock() + prisma_client._cleanup_engine_watcher = cleanup + + prisma_client._on_pidfd_readable() + await asyncio.sleep(0) + pinned = { + "confirmed_dead": prisma_client._engine_confirmed_dead, + "cleanup_called": cleanup.call_count, + "reconnect_called": prisma_client.attempt_db_reconnect.await_count, + "force_kwarg": prisma_client.attempt_db_reconnect.await_args.kwargs["force"], + } + assert pinned == { + "confirmed_dead": True, + "cleanup_called": 1, + "reconnect_called": 1, + "force_kwarg": True, + } + + +@pytest.mark.asyncio +async def test_on_pidfd_readable_noop_when_already_dead_closes_pidfd( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """When _engine_confirmed_dead is already True, the reader handler should + not schedule another reconnect and should release the pidfd resource. + """ + closed: list[int] = [] + monkeypatch.setattr("os.close", lambda fd: closed.append(fd)) + loop = asyncio.get_running_loop() + removed: list[int] = [] + monkeypatch.setattr(loop, "remove_reader", lambda fd: removed.append(fd)) + + prisma_client._engine_confirmed_dead = True + prisma_client._engine_pidfd = 99 + prisma_client.attempt_db_reconnect = AsyncMock() + + prisma_client._on_pidfd_readable() + pinned = { + "engine_pidfd": prisma_client._engine_pidfd, + "closed": closed, + "removed": removed, + "reconnect_call_count": prisma_client.attempt_db_reconnect.await_count, + } + assert pinned == { + "engine_pidfd": -1, + "closed": [99], + "removed": [99], + "reconnect_call_count": 0, + } + + +@pytest.mark.asyncio +async def test_poll_engine_proc_detects_death_and_reconnects( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 555 + prisma_client._watching_engine = True + prisma_client.attempt_db_reconnect = AsyncMock() + monkeypatch.setattr("os.kill", MagicMock(side_effect=ProcessLookupError())) + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + prisma_client._cleanup_engine_watcher = MagicMock() + + await prisma_client._poll_engine_proc() + pinned = { + "reconnect_count": prisma_client.attempt_db_reconnect.await_count, + "cleanup_count": prisma_client._cleanup_engine_watcher.call_count, + "confirmed_dead": prisma_client._engine_confirmed_dead, + "reason": prisma_client.attempt_db_reconnect.await_args.kwargs["reason"], + } + assert pinned == { + "reconnect_count": 1, + "cleanup_count": 1, + "confirmed_dead": True, + "reason": "engine_process_death", + } + + +@pytest.mark.asyncio +async def test_poll_engine_proc_returns_on_permission_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + prisma_client._engine_pid = 555 + prisma_client._watching_engine = True + monkeypatch.setattr("os.kill", MagicMock(side_effect=PermissionError())) + prisma_client._cleanup_engine_watcher = MagicMock() + await prisma_client._poll_engine_proc() + assert prisma_client._cleanup_engine_watcher.call_count == 1 + + +@pytest.mark.asyncio +async def test_cleanup_engine_watcher_resets_state( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + closed: list[int] = [] + monkeypatch.setattr("os.close", lambda fd: closed.append(fd)) + loop = asyncio.get_running_loop() + removed: list[int] = [] + monkeypatch.setattr(loop, "remove_reader", lambda fd: removed.append(fd)) + + prisma_client._engine_pidfd = 42 + prisma_client._engine_pid = 999 + prisma_client._engine_wait_thread = MagicMock() + prisma_client._watching_engine = True + + prisma_client._cleanup_engine_watcher() + pinned = { + "engine_pidfd": prisma_client._engine_pidfd, + "engine_pid": prisma_client._engine_pid, + "wait_thread": prisma_client._engine_wait_thread, + "watching": prisma_client._watching_engine, + "closed": closed, + "removed": removed, + } + assert pinned == { + "engine_pidfd": -1, + "engine_pid": 0, + "wait_thread": None, + "watching": False, + "closed": [42], + "removed": [42], + } + + +@pytest.mark.asyncio +async def test_cleanup_engine_watcher_swallows_close_error( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("os.close", MagicMock(side_effect=OSError("bad fd"))) + loop = asyncio.get_running_loop() + monkeypatch.setattr(loop, "remove_reader", MagicMock(side_effect=Exception("boom"))) + prisma_client._engine_pidfd = 99 + prisma_client._cleanup_engine_watcher() + assert prisma_client._engine_pidfd == -1 + + +@pytest.mark.asyncio +async def test_start_engine_watcher_picks_waitpid_when_available( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=12345)) + monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock(return_value=True)) + pidfd_called = MagicMock(return_value=False) + monkeypatch.setattr(prisma_client, "_try_pidfd_watch", pidfd_called) + await prisma_client._start_engine_watcher() + pinned = { + "engine_pid": prisma_client._engine_pid, + "confirmed_dead_reset": prisma_client._engine_confirmed_dead, + "waitpid_called": prisma_client._try_waitpid_watch.call_count, + "pidfd_skipped": pidfd_called.call_count, + } + assert pinned == { + "engine_pid": 12345, + "confirmed_dead_reset": False, + "waitpid_called": 1, + "pidfd_skipped": 0, + } + + +@pytest.mark.asyncio +async def test_start_engine_watcher_returns_early_when_pid_unknown( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=0)) + monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock()) + await prisma_client._start_engine_watcher() + assert prisma_client._try_waitpid_watch.call_count == 0 + + +@pytest.mark.asyncio +async def test_start_engine_watcher_falls_back_to_polling_when_no_kernel_apis( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(prisma_client, "_get_engine_pid", MagicMock(return_value=4242)) + monkeypatch.setattr(prisma_client, "_try_waitpid_watch", MagicMock(return_value=False)) + monkeypatch.setattr(prisma_client, "_try_pidfd_watch", MagicMock(return_value=False)) + monkeypatch.setattr(prisma_client, "_poll_engine_proc", AsyncMock()) + await prisma_client._start_engine_watcher() + await asyncio.sleep(0) + assert prisma_client._watching_engine is True + + +def test_stop_engine_watcher_clears_dead_flag( + prisma_client: PrismaClient, +) -> None: + prisma_client._engine_confirmed_dead = True + prisma_client._cleanup_engine_watcher = MagicMock() + prisma_client._stop_engine_watcher() + assert prisma_client._cleanup_engine_watcher.call_count == 1 + assert prisma_client._engine_confirmed_dead is False + + +def test_stop_engine_watcher_error_in_cleanup_propagates( + prisma_client: PrismaClient, +) -> None: + prisma_client._cleanup_engine_watcher = MagicMock(side_effect=RuntimeError("cleanup boom")) + with pytest.raises(RuntimeError, match="cleanup boom"): + prisma_client._stop_engine_watcher() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py new file mode 100644 index 00000000000..7e7e98d1360 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py @@ -0,0 +1,400 @@ +"""Pin ``PrismaClient`` read-side data operations. + +Symbols pinned here: + - ``PrismaClient.hash_token`` + - ``PrismaClient.jsonify_object`` + - ``PrismaClient.jsonify_team_object`` + - ``PrismaClient.check_view_exists`` + - ``PrismaClient.get_request_status`` + - ``PrismaClient.get_generic_data`` + - ``PrismaClient._query_first_with_cached_plan_fallback`` + - ``PrismaClient.get_data`` +""" + +from __future__ import annotations + +import hashlib +import json +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.utils import PrismaClient + + +def test_hash_token_method_returns_sha256(prisma_client: PrismaClient) -> None: + token = "sk-token-xyz" + actual = { + "result": prisma_client.hash_token(token), + "len": len(prisma_client.hash_token(token)), + "expected": hashlib.sha256(token.encode()).hexdigest(), + "deterministic": prisma_client.hash_token(token) + == prisma_client.hash_token(token), + } + assert actual == { + "result": hashlib.sha256(token.encode()).hexdigest(), + "len": 64, + "expected": hashlib.sha256(token.encode()).hexdigest(), + "deterministic": True, + } + + +def test_hash_token_method_error_on_non_string(prisma_client: PrismaClient) -> None: + with pytest.raises(AttributeError): + prisma_client.hash_token(None) # type: ignore[arg-type] + + +def test_jsonify_object_serializes_nested_dicts(prisma_client: PrismaClient) -> None: + data = { + "metadata": {"a": 1, "b": [2, 3]}, + "models": ["gpt-4o", "gpt-4o-mini"], + "token": "abc", + "spend": 1.23, + } + result = prisma_client.jsonify_object(data) + parsed_meta = json.loads(result["metadata"]) + assert result == { + "metadata": json.dumps(data["metadata"]), + "models": ["gpt-4o", "gpt-4o-mini"], + "token": "abc", + "spend": 1.23, + } + assert parsed_meta == {"a": 1, "b": [2, 3]} + + +def test_jsonify_object_fallback_for_unserializable_dict( + prisma_client: PrismaClient, +) -> None: + class _Bad: + pass + + data = {"metadata": {"x": _Bad()}, "label": "ok", "n": 1} + result = prisma_client.jsonify_object(data) + assert result == { + "metadata": "failed-to-serialize-json", + "label": "ok", + "n": 1, + } + + +def test_jsonify_object_error_on_non_dict(prisma_client: PrismaClient) -> None: + with pytest.raises(AttributeError): + prisma_client.jsonify_object(None) # type: ignore[arg-type] + + +def test_jsonify_team_object_converts_members_to_json_string( + prisma_client: PrismaClient, +) -> None: + data = { + "team_id": "t1", + "members_with_roles": [{"role": "admin", "user_id": "u1"}], + "metadata": {"foo": "bar"}, + "models": ["gpt-4"], + } + result = prisma_client.jsonify_team_object(data) + assert result == { + "team_id": "t1", + "members_with_roles": json.dumps(data["members_with_roles"]), + "metadata": json.dumps(data["metadata"]), + "models": ["gpt-4"], + } + + +def test_jsonify_team_object_error_on_non_dict(prisma_client: PrismaClient) -> None: + with pytest.raises(AttributeError): + prisma_client.jsonify_team_object(None) # type: ignore[arg-type] + + +@pytest.mark.parametrize( + "metadata,expected", + [ + ({"status": "failure"}, "failure"), + ({"status": "success"}, "success"), + ({}, "success"), + ("not-json", "success"), + (json.dumps({"status": "failure"}), "failure"), + ], +) +def test_get_request_status_pins_status_resolution( + prisma_client: PrismaClient, metadata: Any, expected: str +) -> None: + assert prisma_client.get_request_status({"metadata": metadata}) == expected + + +def test_get_request_status_error_returns_success_default( + prisma_client: PrismaClient, +) -> None: + """``get_request_status`` swallows AttributeError / JSONDecodeError and + defaults to ``success`` to avoid blocking the request pipeline. + """ + + class _Broken: + def get(self, *_: Any, **__: Any) -> Any: + raise AttributeError("broken metadata") + + actual = prisma_client.get_request_status({"metadata": _Broken()}) + assert actual == "success" + + +@pytest.mark.asyncio +async def test_get_generic_data_dispatches_by_table( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace(user_id="u1", spend=0.5, name="Alice") + prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=row) + result = await prisma_client.get_generic_data( + key="user_id", value="u1", table_name="users" + ) + actual = { + "result_is_row": result is row, + "find_first_count": prisma_client.db.litellm_usertable.find_first.await_count, + "where_kwarg": prisma_client.db.litellm_usertable.find_first.await_args.kwargs[ + "where" + ], + "user_attr": result.user_id, + } + assert actual == { + "result_is_row": True, + "find_first_count": 1, + "where_kwarg": {"user_id": "u1"}, + "user_attr": "u1", + } + + +@pytest.mark.asyncio +async def test_get_generic_data_unknown_table_returns_none( + prisma_client: PrismaClient, +) -> None: + result = await prisma_client.get_generic_data( + key="x", value="y", table_name="bogus" # type: ignore[arg-type] + ) + assert result is None + + +@pytest.mark.asyncio +async def test_get_generic_data_logs_failure_handler_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_usertable.find_first = AsyncMock( + side_effect=RuntimeError("db boom") + ) + with pytest.raises(RuntimeError, match="db boom"): + await prisma_client.get_generic_data( + key="user_id", value="x", table_name="users" + ) + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_happy_returns_row( + prisma_client: PrismaClient, +) -> None: + expected = {"token": "abc", "team_spend": 1.0, "team_max_budget": 5.0} + prisma_client.db.query_first = AsyncMock(return_value=expected) + result = await prisma_client._query_first_with_cached_plan_fallback( + "SELECT * FROM x WHERE token = $1", "abc" + ) + actual = { + "result": result, + "call_count": prisma_client.db.query_first.await_count, + "args": prisma_client.db.query_first.await_args.args, + "matches": result == expected, + } + assert actual == { + "result": expected, + "call_count": 1, + "args": ("SELECT * FROM x WHERE token = $1", "abc"), + "matches": True, + } + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_retries_on_cached_plan_error( + prisma_client: PrismaClient, +) -> None: + expected = {"token": "abc", "team_spend": 1.0, "team_max_budget": 5.0} + prisma_client.db.query_first = AsyncMock( + side_effect=[ + RuntimeError("cached plan must not change result type"), + expected, + ] + ) + result = await prisma_client._query_first_with_cached_plan_fallback( + "SELECT * FROM x WHERE token = $1", "abc" + ) + assert result == expected + assert prisma_client.db.query_first.await_count == 2 + second_call_sql = prisma_client.db.query_first.await_args_list[1].args[0] + assert "cache_invalidated_" in second_call_sql + + +@pytest.mark.asyncio +async def test_query_first_with_cached_plan_fallback_reraises_non_plan_errors( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_first = AsyncMock(side_effect=RuntimeError("totally unrelated")) + with pytest.raises(RuntimeError, match="totally unrelated"): + await prisma_client._query_first_with_cached_plan_fallback("SELECT 1") + + +@pytest.mark.asyncio +async def test_check_view_exists_noop_when_all_views_present( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock( + return_value=[ + { + "view_count": 8, + "view_names": [ + "LiteLLM_VerificationTokenView", + "MonthlyGlobalSpend", + "Last30dKeysBySpend", + "Last30dModelsBySpend", + "MonthlyGlobalSpendPerKey", + "MonthlyGlobalSpendPerUserPerKey", + "Last30dTopEndUsersSpend", + "DailyTagSpend", + ], + } + ] + ) + prisma_client.db.execute_raw = AsyncMock() + result = await prisma_client.check_view_exists() + actual = { + "result": result, + "query_raw_calls": prisma_client.db.query_raw.await_count, + "execute_raw_calls": prisma_client.db.execute_raw.await_count, + "view_query_contains_token_view": "LiteLLM_VerificationTokenView" + in prisma_client.db.query_raw.await_args.args[0], + } + assert actual == { + "result": None, + "query_raw_calls": 1, + "execute_raw_calls": 0, + "view_query_contains_token_view": True, + } + + +@pytest.mark.asyncio +async def test_check_view_exists_creates_token_view_when_missing( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock( + return_value=[ + { + "view_count": 1, + "view_names": ["DailyTagSpend"], + } + ] + ) + prisma_client.db.execute_raw = AsyncMock() + prisma_client.health_check = AsyncMock(return_value=[{"?column?": 1}]) + result = await prisma_client.check_view_exists() + actual = { + "result": result, + "create_called": prisma_client.db.execute_raw.await_count, + "create_sql_starts_with_create_view": prisma_client.db.execute_raw.await_args.args[ + 0 + ] + .strip() + .startswith('CREATE VIEW "LiteLLM_VerificationTokenView"'), + } + assert actual == { + "result": None, + "create_called": 1, + "create_sql_starts_with_create_view": True, + } + + +@pytest.mark.asyncio +async def test_check_view_exists_raises_when_query_raw_fails( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("db down")) + with pytest.raises(RuntimeError, match="db down"): + await prisma_client.check_view_exists() + + +@pytest.mark.asyncio +async def test_get_data_token_find_unique_returns_record( + prisma_client: PrismaClient, +) -> None: + token = "sk-key-1" + hashed = hashlib.sha256(token.encode()).hexdigest() + record = SimpleNamespace(token=hashed, user_id="u1", expires=None, spend=0.5) + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=record + ) + + result = await prisma_client.get_data(token=token, table_name="key") + actual = { + "result_is_record": result is record, + "where_arg": prisma_client.db.litellm_verificationtoken.find_unique.await_args.kwargs[ + "where" + ], + "include_arg": prisma_client.db.litellm_verificationtoken.find_unique.await_args.kwargs[ + "include" + ], + "token_field_matches": result.token == hashed, + } + assert actual == { + "result_is_record": True, + "where_arg": {"token": hashed}, + "include_arg": {"litellm_budget_table": True}, + "token_field_matches": True, + } + + +@pytest.mark.asyncio +async def test_get_data_token_find_unique_missing_token_raises_401( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + with pytest.raises(HTTPException) as excinfo: + await prisma_client.get_data(token="sk-missing", table_name="key") + err = excinfo.value + assert "invalid user key" in err.detail + assert err.status_code == 401 + + +@pytest.mark.asyncio +async def test_get_data_user_find_unique_returns_user_row( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace( + user_id="u-7", + spend=1.5, + max_budget=10.0, + organization_memberships=[], + ) + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row) + result = await prisma_client.get_data(user_id="u-7", table_name="user") + actual = { + "result_is_row": result is row, + "where_arg": prisma_client.db.litellm_usertable.find_unique.await_args.kwargs[ + "where" + ], + "include_arg": prisma_client.db.litellm_usertable.find_unique.await_args.kwargs[ + "include" + ], + "spend": row.spend, + } + assert actual == { + "result_is_row": True, + "where_arg": {"user_id": "u-7"}, + "include_arg": {"organization_memberships": True}, + "spend": 1.5, + } + + +@pytest.mark.asyncio +async def test_get_data_logs_and_raises_on_db_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + side_effect=RuntimeError("network split") + ) + with pytest.raises(RuntimeError, match="network split"): + await prisma_client.get_data(token="sk-broken", table_name="key") diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py new file mode 100644 index 00000000000..220fff1a881 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_health.py @@ -0,0 +1,292 @@ +"""Pin ``PrismaClient`` health + spend-logs counter helpers. + +Symbols pinned here: + - ``PrismaClient.health_check`` + - ``PrismaClient._get_spend_logs_row_count`` + - ``PrismaClient._set_spend_logs_row_count_in_proxy_state`` + - ``PrismaClient._validate_response_time`` + - ``PrismaClient._clean_details`` + - ``PrismaClient.save_health_check_result`` + - ``PrismaClient.get_health_check_history`` + - ``PrismaClient.get_all_latest_health_checks`` + - ``PrismaClient._is_sha256_hex`` (a nested helper inside + ``migrate_passwords_to_scrypt_async``; the pin list assigns it to this + cluster as a documentation artifact) +""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_health_check_returns_query_raw_result( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + result = await prisma_client.health_check() + actual = { + "result": result, + "query_raw_called": prisma_client.db.query_raw.await_count, + "query_sql": prisma_client.db.query_raw.await_args.args[0], + "type": type(result).__name__, + } + assert actual == { + "result": [{"?column?": 1}], + "query_raw_called": 1, + "query_sql": "SELECT 1", + "type": "list", + } + + +@pytest.mark.asyncio +async def test_health_check_raises_when_query_raw_fails( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("connection refused")) + with pytest.raises(RuntimeError, match="connection refused"): + await prisma_client.health_check() + + +@pytest.mark.asyncio +async def test_get_spend_logs_row_count_returns_int_from_pg_class( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(return_value=[{"reltuples": 12345}]) + result = await prisma_client._get_spend_logs_row_count() + actual = { + "result": result, + "query_count": prisma_client.db.query_raw.await_count, + "query_kwargs": prisma_client.db.query_raw.await_args.kwargs, + "type": type(result).__name__, + } + assert actual == { + "result": 12345, + "query_count": 1, + "query_kwargs": { + "query": prisma_client.db.query_raw.await_args.kwargs["query"] + }, + "type": "int", + } + + +@pytest.mark.asyncio +async def test_get_spend_logs_row_count_error_falls_back_to_zero( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.query_raw = AsyncMock(side_effect=RuntimeError("perm denied")) + assert await prisma_client._get_spend_logs_row_count() == 0 + + +@pytest.mark.asyncio +async def test_set_spend_logs_row_count_in_proxy_state_writes_to_state( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + fake_state = MagicMock() + fake_state.set_proxy_state_variable = MagicMock() + + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr(proxy_server_mod, "proxy_state", fake_state, raising=False) + + prisma_client._get_spend_logs_row_count = AsyncMock(return_value=99) + await prisma_client._set_spend_logs_row_count_in_proxy_state() + kwargs = fake_state.set_proxy_state_variable.call_args.kwargs + assert kwargs == {"variable_name": "spend_logs_row_count", "value": 99} + + +@pytest.mark.asyncio +async def test_set_spend_logs_row_count_error_raises_through_backoff( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + fake_state = MagicMock() + fake_state.set_proxy_state_variable = MagicMock(side_effect=RuntimeError("boom")) + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr(proxy_server_mod, "proxy_state", fake_state, raising=False) + + prisma_client._get_spend_logs_row_count = AsyncMock(return_value=1) + with pytest.raises(RuntimeError, match="boom"): + await prisma_client._set_spend_logs_row_count_in_proxy_state() + + +def test_validate_response_time_passes_finite_value(prisma_client: PrismaClient) -> None: + inputs = { + "ok": prisma_client._validate_response_time(123.45), + "none": prisma_client._validate_response_time(None), + "inf": prisma_client._validate_response_time(float("inf")), + "neg_inf": prisma_client._validate_response_time(float("-inf")), + "nan": prisma_client._validate_response_time(float("nan")), + } + assert inputs == { + "ok": 123.45, + "none": None, + "inf": None, + "neg_inf": None, + "nan": None, + } + + +def test_validate_response_time_invalid_string_returns_none( + prisma_client: PrismaClient, +) -> None: + """Non-numeric input is logged and returned as None. The name is the + error hint; the input itself is invalid, not a thrown exception.""" + assert prisma_client._validate_response_time("not-a-float") is None + + +def test_clean_details_round_trips_json(prisma_client: PrismaClient) -> None: + details = {"latency": 1.5, "ok": True, "error": None, "model": "gpt-4o"} + cleaned = prisma_client._clean_details(details) + pinned = { + "cleaned": cleaned, + "is_dict": isinstance(cleaned, dict), + "none_for_non_dict": prisma_client._clean_details("oops"), # type: ignore[arg-type] + "none_for_none": prisma_client._clean_details(None), + } + assert pinned == { + "cleaned": details, + "is_dict": True, + "none_for_non_dict": None, + "none_for_none": None, + } + + +def test_clean_details_invalid_payload_returns_none( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """When ``safe_dumps`` itself blows up (e.g. an internal exception), the + error path swallows it and returns None. + """ + import litellm.proxy.utils as utils_mod + + def _explode(_: Any) -> str: + raise RuntimeError("safe_dumps broken") + + monkeypatch.setattr(utils_mod, "safe_dumps", _explode) + assert prisma_client._clean_details({"x": 1}) is None + + +@pytest.mark.asyncio +async def test_save_health_check_result_creates_record( + prisma_client: PrismaClient, +) -> None: + expected = MagicMock(name="HealthCheckRow") + prisma_client.db.litellm_healthchecktable.create = AsyncMock(return_value=expected) + result = await prisma_client.save_health_check_result( + model_name="gpt-4o", + status="healthy", + healthy_count=3, + unhealthy_count=0, + response_time_ms=150.0, + details={"latency": 1, "ok": True}, + checked_by="probe", + model_id="m-1", + ) + data = prisma_client.db.litellm_healthchecktable.create.await_args.kwargs["data"] + pinned = { + "returned": result, + "model_name": data["model_name"], + "status": data["status"], + "healthy_count": data["healthy_count"], + "response_time_ms": data["response_time_ms"], + "details": data["details"], + "checked_by": data["checked_by"], + "model_id": data["model_id"], + } + assert pinned == { + "returned": expected, + "model_name": "gpt-4o", + "status": "healthy", + "healthy_count": 3, + "response_time_ms": 150.0, + "details": {"latency": 1, "ok": True}, + "checked_by": "probe", + "model_id": "m-1", + } + + +@pytest.mark.asyncio +async def test_save_health_check_result_db_failure_returns_none( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_healthchecktable.create = AsyncMock( + side_effect=RuntimeError("db down") + ) + result = await prisma_client.save_health_check_result( + model_name="gpt-4o", status="healthy" + ) + assert result is None + + +@pytest.mark.asyncio +async def test_get_health_check_history_filters_by_model_and_status( + prisma_client: PrismaClient, +) -> None: + rows = [MagicMock(name=f"row-{i}") for i in range(2)] + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=rows) + result = await prisma_client.get_health_check_history( + model_name="gpt-4o", limit=5, offset=10, status_filter="healthy" + ) + kwargs = prisma_client.db.litellm_healthchecktable.find_many.await_args.kwargs + actual = { + "result_len": len(result), + "where": kwargs["where"], + "order": kwargs["order"], + "take": kwargs["take"], + "skip": kwargs["skip"], + } + assert actual == { + "result_len": 2, + "where": {"model_name": "gpt-4o", "status": "healthy"}, + "order": {"checked_at": "desc"}, + "take": 5, + "skip": 10, + } + + +@pytest.mark.asyncio +async def test_get_health_check_history_db_error_returns_empty_list( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock( + side_effect=RuntimeError("network down") + ) + assert await prisma_client.get_health_check_history() == [] + + +@pytest.mark.asyncio +async def test_get_all_latest_health_checks_uses_distinct( + prisma_client: PrismaClient, +) -> None: + rows = [MagicMock(name=f"row-{i}") for i in range(3)] + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=rows) + result = await prisma_client.get_all_latest_health_checks() + kwargs = prisma_client.db.litellm_healthchecktable.find_many.await_args.kwargs + actual = { + "len": len(result), + "distinct": kwargs["distinct"], + "order_len": len(kwargs["order"]), + "first_order": kwargs["order"][0], + } + assert actual == { + "len": 3, + "distinct": ["model_id", "model_name"], + "order_len": 3, + "first_order": {"model_id": "asc"}, + } + + +@pytest.mark.asyncio +async def test_get_all_latest_health_checks_db_error_returns_empty_list( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_healthchecktable.find_many = AsyncMock( + side_effect=RuntimeError("oops") + ) + assert await prisma_client.get_all_latest_health_checks() == [] diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py new file mode 100644 index 00000000000..30fd4a74bb0 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py @@ -0,0 +1,207 @@ +"""Pin ``PrismaClient`` lifecycle methods. + +Symbols pinned here: + - ``PrismaClient.__init__`` + - ``PrismaClient.writer_db`` + - ``PrismaClient.connect`` + - ``PrismaClient.disconnect`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_prismaclient_init_wires_default_config( + patched_prisma_import: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + monkeypatch.delenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", raising=False) + monkeypatch.delenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", raising=False) + monkeypatch.delenv("PRISMA_HEALTH_WATCHDOG_ENABLED", raising=False) + monkeypatch.delenv("PRISMA_RECONNECT_ESCALATION_THRESHOLD", raising=False) + + proxy_logging = MagicMock() + pc = PrismaClient( + database_url="postgres://x:y@h:5432/db", + proxy_logging_obj=proxy_logging, + ) + pinned = { + "iam_token_db_auth": pc.iam_token_db_auth, + "db_reconnect_cooldown_seconds": pc._db_reconnect_cooldown_seconds, + "db_health_watchdog_interval_seconds": pc._db_health_watchdog_interval_seconds, + "db_health_watchdog_enabled": pc._db_health_watchdog_enabled, + "reconnect_escalation_threshold": pc._reconnect_escalation_threshold, + "consecutive_reconnect_failures": pc._consecutive_reconnect_failures, + "engine_pid": pc._engine_pid, + "watching_engine": pc._watching_engine, + "proxy_logging_obj_set": pc.proxy_logging_obj is proxy_logging, + "db_reconnect_lock_is_lock": isinstance(pc._db_reconnect_lock, asyncio.Lock), + } + assert pinned == { + "iam_token_db_auth": None, + "db_reconnect_cooldown_seconds": 15, + "db_health_watchdog_interval_seconds": 30, + "db_health_watchdog_enabled": True, + "reconnect_escalation_threshold": 3, + "consecutive_reconnect_failures": 0, + "engine_pid": 0, + "watching_engine": False, + "proxy_logging_obj_set": True, + "db_reconnect_lock_is_lock": True, + } + + +def test_prismaclient_init_honors_env_overrides( + patched_prisma_import: MagicMock, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", "42") + monkeypatch.setenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", "60") + monkeypatch.setenv("PRISMA_HEALTH_WATCHDOG_ENABLED", "false") + monkeypatch.setenv("PRISMA_RECONNECT_ESCALATION_THRESHOLD", "7") + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + + pc = PrismaClient( + database_url="postgres://x:y@h:5432/db", + proxy_logging_obj=MagicMock(), + ) + pinned = { + "db_reconnect_cooldown_seconds": pc._db_reconnect_cooldown_seconds, + "db_health_watchdog_interval_seconds": pc._db_health_watchdog_interval_seconds, + "db_health_watchdog_enabled": pc._db_health_watchdog_enabled, + "reconnect_escalation_threshold": pc._reconnect_escalation_threshold, + } + assert pinned == { + "db_reconnect_cooldown_seconds": 42, + "db_health_watchdog_interval_seconds": 60, + "db_health_watchdog_enabled": False, + "reconnect_escalation_threshold": 7, + } + + +def test_prismaclient_init_raises_when_prisma_not_generated() -> None: + """If ``from prisma import Prisma`` fails, the init re-raises with the + 'prisma generate' guidance message. + """ + import prisma as _prisma_pkg + + had_prisma_attr = "Prisma" in _prisma_pkg.__dict__ + previous_prisma_attr = _prisma_pkg.__dict__.get("Prisma") + if had_prisma_attr: + del _prisma_pkg.Prisma # type: ignore[attr-defined] + try: + with pytest.raises(Exception, match="prisma generate"): + PrismaClient( + database_url="postgres://x:y@h:5432/db", + proxy_logging_obj=MagicMock(), + ) + finally: + if had_prisma_attr: + _prisma_pkg.Prisma = previous_prisma_attr # type: ignore[attr-defined] + + +def test_writer_db_returns_db_when_no_routing(prisma_client: PrismaClient) -> None: + actual = { + "writer_is_db": prisma_client.writer_db is prisma_client.db, + "type_consistency": type(prisma_client.writer_db) is type(prisma_client.db), + "callable_query_raw": callable(prisma_client.writer_db.query_raw), + } + assert actual == { + "writer_is_db": True, + "type_consistency": True, + "callable_query_raw": True, + } + + +def test_writer_db_unwraps_routing_wrapper(prisma_client: PrismaClient) -> None: + from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper + + inner_writer = MagicMock(name="WriterInsideRouter") + + class _FakeRouting(RoutingPrismaWrapper): # type: ignore[misc] + def __init__(self) -> None: + self._writer = inner_writer + + prisma_client.db = _FakeRouting() + assert prisma_client.writer_db is inner_writer + + +def test_writer_db_error_when_db_attribute_missing(prisma_client: PrismaClient) -> None: + del prisma_client.db + with pytest.raises(AttributeError): + _ = prisma_client.writer_db + + +@pytest.mark.asyncio +async def test_connect_invokes_underlying_when_disconnected( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.is_connected = MagicMock(return_value=False) + prisma_client.db.connect = AsyncMock() + await prisma_client.connect() + actual = { + "connect_called": prisma_client.db.connect.await_count, + "is_connected_called": prisma_client.db.is_connected.call_count, + "no_failure_handler": prisma_client.proxy_logging_obj.failure_handler.await_count, + } + assert actual == { + "connect_called": 1, + "is_connected_called": 1, + "no_failure_handler": 0, + } + + +@pytest.mark.asyncio +async def test_connect_is_noop_when_already_connected( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.is_connected = MagicMock(return_value=True) + prisma_client.db.connect = AsyncMock() + await prisma_client.connect() + assert prisma_client.db.connect.await_count == 0 + + +@pytest.mark.asyncio +async def test_connect_invokes_failure_handler_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.is_connected = MagicMock(return_value=False) + prisma_client.db.connect = AsyncMock(side_effect=RuntimeError("network down")) + with pytest.raises(RuntimeError, match="network down"): + await prisma_client.connect() + + +@pytest.mark.asyncio +async def test_disconnect_calls_underlying(prisma_client: PrismaClient) -> None: + prisma_client.db.disconnect = AsyncMock() + await prisma_client.disconnect() + actual = { + "disconnect_called": prisma_client.db.disconnect.await_count, + "failure_handler_called": prisma_client.proxy_logging_obj.failure_handler.await_count, + "type": type(prisma_client.db.disconnect).__name__, + } + assert actual == { + "disconnect_called": 1, + "failure_handler_called": 0, + "type": "AsyncMock", + } + + +@pytest.mark.asyncio +async def test_disconnect_raises_when_underlying_fails( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.disconnect = AsyncMock(side_effect=RuntimeError("disconnect boom")) + with pytest.raises(RuntimeError, match="disconnect boom"): + await prisma_client.disconnect() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py new file mode 100644 index 00000000000..f669e6be88d --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_reconnect.py @@ -0,0 +1,371 @@ +"""Pin ``PrismaClient`` reconnect + watchdog symbols. + +Symbols pinned here: + - ``PrismaClient._run_reconnect_cycle`` + - ``PrismaClient._attempt_reconnect_inside_lock`` + - ``PrismaClient.attempt_db_reconnect`` + - ``PrismaClient.start_db_health_watchdog_task`` + - ``PrismaClient.stop_db_health_watchdog_task`` + - ``PrismaClient._db_health_watchdog_loop`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_run_reconnect_cycle_direct_path_when_engine_alive( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("DATABASE_URL", "postgres://x:y@h:5432/db") + prisma_client._engine_confirmed_dead = False + prisma_client._engine_pid = 0 + prisma_client.db.recreate_prisma_client = AsyncMock() + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._cleanup_engine_watcher = MagicMock() + + writer = MagicMock() + writer.query_raw = AsyncMock(return_value=[{"?column?": 1}]) + monkeypatch.setattr( + PrismaClient, + "writer_db", + property(lambda self: writer), + ) + + await prisma_client._run_reconnect_cycle(timeout_seconds=5) + pinned = { + "recreate_called": prisma_client.db.recreate_prisma_client.await_count, + "start_watcher_called": prisma_client._start_engine_watcher.await_count, + "writer_smoke_test_called": writer.query_raw.await_count, + "engine_confirmed_dead": prisma_client._engine_confirmed_dead, + } + assert pinned == { + "recreate_called": 1, + "start_watcher_called": 1, + "writer_smoke_test_called": 1, + "engine_confirmed_dead": False, + } + + +@pytest.mark.asyncio +async def test_run_reconnect_cycle_heavy_path_when_engine_dead( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("DATABASE_URL", "postgres://x:y@h:5432/db") + prisma_client._engine_confirmed_dead = True + prisma_client._engine_pid = 1234 + prisma_client.db.recreate_prisma_client = AsyncMock() + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._cleanup_engine_watcher = MagicMock() + monkeypatch.setattr(PrismaClient, "_reap_all_zombies", staticmethod(lambda: set())) + + await prisma_client._run_reconnect_cycle(timeout_seconds=5) + pinned = { + "recreate_called": prisma_client.db.recreate_prisma_client.await_count, + "start_watcher_called": prisma_client._start_engine_watcher.await_count, + "cleanup_called": prisma_client._cleanup_engine_watcher.call_count, + "dead_flag_cleared": prisma_client._engine_confirmed_dead, + } + assert pinned == { + "recreate_called": 1, + "start_watcher_called": 1, + "cleanup_called": 1, + "dead_flag_cleared": False, + } + + +@pytest.mark.asyncio +async def test_run_reconnect_cycle_raises_when_database_url_missing( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.delenv("DATABASE_URL", raising=False) + with pytest.raises(RuntimeError, match="DATABASE_URL not set"): + await prisma_client._run_reconnect_cycle(timeout_seconds=1) + + +@pytest.mark.asyncio +async def test_attempt_reconnect_inside_lock_runs_cycle_and_resets_counter( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._consecutive_reconnect_failures = 2 + prisma_client._run_reconnect_cycle = AsyncMock() + + ok = await prisma_client._attempt_reconnect_inside_lock( + force=True, reason="test", timeout_seconds=1 + ) + pinned = { + "returned": ok, + "cycle_called": prisma_client._run_reconnect_cycle.await_count, + "failures_reset": prisma_client._consecutive_reconnect_failures, + } + assert pinned == { + "returned": True, + "cycle_called": 1, + "failures_reset": 0, + } + + +@pytest.mark.asyncio +async def test_attempt_reconnect_inside_lock_skips_when_in_cooldown( + prisma_client: PrismaClient, +) -> None: + import time + + prisma_client._db_reconnect_cooldown_seconds = 60 + prisma_client._db_last_reconnect_attempt_ts = time.time() + prisma_client._run_reconnect_cycle = AsyncMock() + + ok = await prisma_client._attempt_reconnect_inside_lock( + force=False, reason="test", timeout_seconds=1 + ) + assert ok is False + assert prisma_client._run_reconnect_cycle.await_count == 0 + + +@pytest.mark.asyncio +async def test_attempt_reconnect_inside_lock_increments_failure_counter_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._consecutive_reconnect_failures = 0 + prisma_client._run_reconnect_cycle = AsyncMock(side_effect=RuntimeError("boom")) + + ok = await prisma_client._attempt_reconnect_inside_lock( + force=True, reason="failing_test", timeout_seconds=1 + ) + assert ok is False + assert prisma_client._consecutive_reconnect_failures == 1 + + +@pytest.mark.asyncio +async def test_attempt_db_reconnect_force_runs_under_lock( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._attempt_reconnect_inside_lock = AsyncMock(return_value=True) + + result = await prisma_client.attempt_db_reconnect(reason="explicit", force=True) + args = prisma_client._attempt_reconnect_inside_lock.await_args + pinned = { + "returned": result, + "calls": prisma_client._attempt_reconnect_inside_lock.await_count, + "passed_force": args.args[0], + "passed_reason": args.args[1], + "passed_timeout": args.args[2], + } + assert pinned == { + "returned": True, + "calls": 1, + "passed_force": True, + "passed_reason": "explicit", + "passed_timeout": None, + } + + +@pytest.mark.asyncio +async def test_attempt_db_reconnect_lock_timeout_returns_false( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """A reconnect attempt that can't acquire the lock within + ``lock_timeout_seconds`` returns False without running the cycle. + + The production code creates an inner task, races it against the + timeout via ``asyncio.wait``, then cancels and awaits the loser. + Under coverage instrumentation on Python 3.11 the CancelledError from + a freshly-cancelled task can outrun the surrounding ``except`` block, + so this test pre-completes the inner task (no cancellation happens) + by replacing ``asyncio.wait`` with a callable that returns the loser + task as still-pending after it's already been completed elsewhere. + """ + completed_task: asyncio.Task[bool] = asyncio.get_running_loop().create_task( + _no_op_returning_true() + ) + # Ensure the inner task has finished before attempt_db_reconnect sees it. + await completed_task + + async def _wait_returns_loser(_tasks: Any, **kwargs: Any) -> Any: + return set(), {completed_task} + + monkeypatch.setattr("asyncio.wait", _wait_returns_loser) + monkeypatch.setattr( + asyncio, + "create_task", + lambda coro, *a, **kw: (coro.close() or completed_task), + ) + + prisma_client._db_last_reconnect_attempt_ts = 0.0 + prisma_client._attempt_reconnect_inside_lock = AsyncMock() + + ok = await prisma_client.attempt_db_reconnect( + reason="lock_busy", + lock_timeout_seconds=0.0, + ) + assert ok is False + assert prisma_client._attempt_reconnect_inside_lock.await_count == 0 + + +async def _no_op_returning_true() -> bool: + return True + + +@pytest.mark.asyncio +async def test_attempt_db_reconnect_skips_in_cooldown_returns_false( + prisma_client: PrismaClient, +) -> None: + import time + + prisma_client._db_reconnect_cooldown_seconds = 60 + prisma_client._db_last_reconnect_attempt_ts = time.time() + ok = await prisma_client.attempt_db_reconnect(reason="cooled_down") + assert ok is False + + +@pytest.mark.asyncio +async def test_start_db_health_watchdog_task_creates_loop_task( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_enabled = True + prisma_client._db_health_watchdog_task = None + prisma_client._start_engine_watcher = AsyncMock() + prisma_client._db_health_watchdog_loop = AsyncMock(return_value=None) + + await prisma_client.start_db_health_watchdog_task() + task = prisma_client._db_health_watchdog_task + # Yield control so the just-scheduled task actually invokes the loop mock. + await asyncio.sleep(0) + pinned = { + "task_type": type(task).__name__, + "watcher_started": prisma_client._start_engine_watcher.await_count, + "loop_invoked": prisma_client._db_health_watchdog_loop.await_count, + } + if task is not None: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + assert pinned == { + "task_type": "Task", + "watcher_started": 1, + "loop_invoked": 1, + } + + +@pytest.mark.asyncio +async def test_start_db_health_watchdog_task_disabled_short_circuits( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_enabled = False + prisma_client._start_engine_watcher = AsyncMock() + await prisma_client.start_db_health_watchdog_task() + assert prisma_client._db_health_watchdog_task is None + assert prisma_client._start_engine_watcher.await_count == 0 + + +@pytest.mark.asyncio +async def test_stop_db_health_watchdog_task_cancels_and_clears( + prisma_client: PrismaClient, +) -> None: + prisma_client._stop_engine_watcher = MagicMock() + + cancel_called = {"n": 0} + + class _FakeTask: + def cancel(self) -> None: + cancel_called["n"] += 1 + + def __await__(self): + return iter([]) + + prisma_client._db_health_watchdog_task = _FakeTask() # type: ignore[assignment] + + await prisma_client.stop_db_health_watchdog_task() + pinned = { + "task_cleared": prisma_client._db_health_watchdog_task, + "engine_stop_called": prisma_client._stop_engine_watcher.call_count, + "cancel_called": cancel_called["n"], + "no_failure": True, + } + assert pinned == { + "task_cleared": None, + "engine_stop_called": 1, + "cancel_called": 1, + "no_failure": True, + } + + +@pytest.mark.asyncio +async def test_stop_db_health_watchdog_task_noop_when_no_task( + prisma_client: PrismaClient, +) -> None: + prisma_client._db_health_watchdog_task = None + prisma_client._stop_engine_watcher = MagicMock(side_effect=RuntimeError("err")) + with pytest.raises(RuntimeError, match="err"): + await prisma_client.stop_db_health_watchdog_task() + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_triggers_reconnect_on_timeout( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """The watchdog loop reconnects when ``wait_for`` raises TimeoutError + or a recognized DB connection error. + """ + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + call_count = {"n": 0} + + async def _timeout_then_cancel(*args: Any, **kwargs: Any) -> None: + call_count["n"] += 1 + if call_count["n"] >= 2: + raise asyncio.CancelledError() + raise asyncio.TimeoutError() + + monkeypatch.setattr("asyncio.wait_for", _timeout_then_cancel) + await prisma_client._db_health_watchdog_loop() + pinned = { + "reconnect_called": prisma_client.attempt_db_reconnect.await_count, + "reconnect_reason": prisma_client.attempt_db_reconnect.await_args.kwargs[ + "reason" + ], + "wait_for_calls": call_count["n"], + "loop_exited_clean": True, + } + assert pinned == { + "reconnect_called": 1, + "reconnect_reason": "db_health_watchdog_connection_error", + "wait_for_calls": 2, + "loop_exited_clean": True, + } + + +@pytest.mark.asyncio +async def test_db_health_watchdog_loop_swallows_non_db_errors( + prisma_client: PrismaClient, monkeypatch: pytest.MonkeyPatch +) -> None: + """A non-DB error during the probe should NOT trigger reconnect; the + loop logs and continues until cancellation. + """ + prisma_client._db_health_watchdog_interval_seconds = 0 + prisma_client.attempt_db_reconnect = AsyncMock() + + call_count = {"n": 0} + + async def _raise_then_cancel(*args: Any, **kwargs: Any) -> None: + call_count["n"] += 1 + if call_count["n"] >= 2: + raise asyncio.CancelledError() + raise ValueError("not a db error") + + monkeypatch.setattr("asyncio.wait_for", _raise_then_cancel) + await prisma_client._db_health_watchdog_loop() + assert prisma_client.attempt_db_reconnect.await_count == 0 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py new file mode 100644 index 00000000000..4e547b81acc --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_writes.py @@ -0,0 +1,260 @@ +"""Pin ``PrismaClient`` write-side data operations. + +Symbols pinned here: + - ``PrismaClient.insert_data`` + - ``PrismaClient.update_data`` + - ``PrismaClient.delete_data`` +""" + +from __future__ import annotations + +import hashlib +import json +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.utils import PrismaClient + + +@pytest.mark.asyncio +async def test_insert_data_hashes_token_and_upserts(prisma_client: PrismaClient) -> None: + token = "sk-secret-1" + response = SimpleNamespace(token=hashlib.sha256(token.encode()).hexdigest(), + key_alias="alias", user_id="u1") + prisma_client.db.litellm_verificationtoken.upsert = AsyncMock(return_value=response) + data = { + "token": token, + "user_id": "u1", + "team_id": "t1", + "metadata": {"a": 1}, + } + result = await prisma_client.insert_data(data=data, table_name="key") + upsert_kwargs = prisma_client.db.litellm_verificationtoken.upsert.await_args.kwargs + actual = { + "returned": result, + "where": upsert_kwargs["where"], + "include": upsert_kwargs["include"], + "create_token": upsert_kwargs["data"]["create"]["token"], + "create_metadata_serialized": isinstance( + upsert_kwargs["data"]["create"]["metadata"], str + ), + "update_empty": upsert_kwargs["data"]["update"], + } + expected_hash = hashlib.sha256(token.encode()).hexdigest() + assert actual == { + "returned": response, + "where": {"token": expected_hash}, + "include": {"litellm_budget_table": True}, + "create_token": expected_hash, + "create_metadata_serialized": True, + "update_empty": {}, + } + + +@pytest.mark.asyncio +async def test_insert_data_strips_null_budget_limits(prisma_client: PrismaClient) -> None: + prisma_client.db.litellm_verificationtoken.upsert = AsyncMock(return_value=None) + await prisma_client.insert_data( + data={"token": "sk-1", "budget_limits": None}, table_name="key" + ) + create_payload = prisma_client.db.litellm_verificationtoken.upsert.await_args.kwargs[ + "data" + ]["create"] + assert "budget_limits" not in create_payload + + +@pytest.mark.asyncio +async def test_insert_data_team_serializes_members(prisma_client: PrismaClient) -> None: + prisma_client.db.litellm_teamtable.upsert = AsyncMock( + return_value=SimpleNamespace(team_id="t1", team_alias="x", spend=0) + ) + data = { + "team_id": "t1", + "team_alias": "x", + "members_with_roles": [{"role": "admin", "user_id": "u1"}], + } + result = await prisma_client.insert_data(data=data, table_name="team") + create_payload = prisma_client.db.litellm_teamtable.upsert.await_args.kwargs["data"][ + "create" + ] + assert result.team_id == "t1" + assert create_payload["members_with_roles"] == json.dumps(data["members_with_roles"]) + assert create_payload["team_id"] == "t1" + + +@pytest.mark.asyncio +async def test_insert_data_user_organization_fk_raises_400( + prisma_client: PrismaClient, +) -> None: + err = RuntimeError( + "Foreign key constraint failed on the field: `LiteLLM_UserTable_organization_id_fkey (index)`" + ) + prisma_client.db.litellm_usertable.upsert = AsyncMock(side_effect=err) + with pytest.raises(HTTPException) as excinfo: + await prisma_client.insert_data( + data={"user_id": "u1", "organization_id": "org-bad"}, table_name="user" + ) + raised = excinfo.value + assert "Foreign Key Constraint failed" in raised.detail["error"] + assert raised.status_code == 400 + + +@pytest.mark.asyncio +async def test_insert_data_logs_and_raises_generic_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.upsert = AsyncMock( + side_effect=RuntimeError("write boom") + ) + with pytest.raises(RuntimeError, match="write boom"): + await prisma_client.insert_data(data={"token": "sk-1"}, table_name="key") + + +@pytest.mark.asyncio +async def test_update_data_token_hashes_and_updates( + prisma_client: PrismaClient, +) -> None: + token = "sk-update-1" + response = SimpleNamespace( + token=hashlib.sha256(token.encode()).hexdigest(), + model_dump=lambda: { + "token": hashlib.sha256(token.encode()).hexdigest(), + "spend": 1.0, + "user_id": "u1", + }, + ) + prisma_client.db.litellm_verificationtoken.update = AsyncMock(return_value=response) + result = await prisma_client.update_data( + token=token, + data={"spend": 1.0}, + ) + update_kwargs = prisma_client.db.litellm_verificationtoken.update.await_args.kwargs + hashed = hashlib.sha256(token.encode()).hexdigest() + actual = { + "result": result, + "where": update_kwargs["where"], + "data_token": update_kwargs["data"]["token"], + "data_spend": update_kwargs["data"]["spend"], + } + assert actual == { + "result": { + "token": hashed, + "data": {"token": hashed, "spend": 1.0, "user_id": "u1"}, + }, + "where": {"token": hashed}, + "data_token": hashed, + "data_spend": 1.0, + } + + +@pytest.mark.asyncio +async def test_update_data_user_upsert_returns_user_envelope( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace(user_id="u2", spend=2.0) + prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=row) + result = await prisma_client.update_data( + data={"user_id": "u2", "spend": 2.0}, + table_name="user", + ) + assert result == {"user_id": "u2", "data": row} + + +@pytest.mark.asyncio +async def test_update_data_team_serializes_members_when_list( + prisma_client: PrismaClient, +) -> None: + row = SimpleNamespace(team_id="t9", team_alias="x") + prisma_client.db.litellm_teamtable.upsert = AsyncMock(return_value=row) + members = [{"role": "admin", "user_id": "u1"}] + result = await prisma_client.update_data( + data={"team_id": "t9", "members_with_roles": members}, + update_key_values={"members_with_roles": members}, + table_name="team", + ) + upsert_kwargs = prisma_client.db.litellm_teamtable.upsert.await_args.kwargs + actual = { + "result_team_id": result["team_id"], + "result_data": result["data"], + "create_members": upsert_kwargs["data"]["create"]["members_with_roles"], + "update_members": upsert_kwargs["data"]["update"]["members_with_roles"], + } + assert actual == { + "result_team_id": "t9", + "result_data": row, + "create_members": json.dumps(members), + "update_members": json.dumps(members), + } + + +@pytest.mark.asyncio +async def test_update_data_logs_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.update = AsyncMock( + side_effect=RuntimeError("update fail") + ) + with pytest.raises(RuntimeError, match="update fail"): + await prisma_client.update_data(token="sk-x", data={"spend": 1.0}) + + +@pytest.mark.asyncio +async def test_delete_data_hashes_sk_tokens_and_calls_delete_many( + prisma_client: PrismaClient, +) -> None: + deleted = SimpleNamespace(count=2) + prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock( + return_value=deleted + ) + tokens = ["sk-one", "sk-two", "raw-hashed-token"] + result = await prisma_client.delete_data(tokens=tokens) + where = prisma_client.db.litellm_verificationtoken.delete_many.await_args.kwargs[ + "where" + ] + expected_hashes = sorted( + [ + hashlib.sha256(b"sk-one").hexdigest(), + hashlib.sha256(b"sk-two").hexdigest(), + "raw-hashed-token", + ] + ) + actual = { + "deleted_keys_attr": result["deleted_keys"], + "where_keys": list(where.keys()), + "filter_in_sorted": sorted(where["token"]["in"]), + "delete_call_count": prisma_client.db.litellm_verificationtoken.delete_many.await_count, + } + assert actual == { + "deleted_keys_attr": deleted, + "where_keys": ["token"], + "filter_in_sorted": expected_hashes, + "delete_call_count": 1, + } + + +@pytest.mark.asyncio +async def test_delete_data_team_calls_team_delete_many( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_teamtable.delete_many = AsyncMock() + result = await prisma_client.delete_data( + team_id_list=["t1", "t2"], table_name="team" + ) + where = prisma_client.db.litellm_teamtable.delete_many.await_args.kwargs["where"] + assert result == {"deleted_teams": ["t1", "t2"]} + assert where == {"team_id": {"in": ["t1", "t2"]}} + + +@pytest.mark.asyncio +async def test_delete_data_logs_and_raises_on_error( + prisma_client: PrismaClient, +) -> None: + prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock( + side_effect=RuntimeError("delete fail") + ) + with pytest.raises(RuntimeError, match="delete fail"): + await prisma_client.delete_data(tokens=["sk-x"]) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py new file mode 100644 index 00000000000..6a4fd516c9b --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py @@ -0,0 +1,275 @@ +"""Pin ``ProxyUpdateSpend`` behavior. + +Symbols pinned here: + - ``ProxyUpdateSpend.update_end_user_spend`` + - ``ProxyUpdateSpend.update_spend_logs`` + - ``ProxyUpdateSpend.disable_spend_updates`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import ProxyUpdateSpend + + +class _AsyncCM: + def __init__(self, target: Any) -> None: + self.target = target + + async def __aenter__(self) -> Any: + return self.target + + async def __aexit__(self, *exc: Any) -> None: + return None + + +@pytest.mark.asyncio +async def test_update_end_user_spend_upserts_each_end_user( + mock_prisma_client: Any, +) -> None: + batcher = MagicMock() + batcher.litellm_endusertable.upsert = MagicMock() + transaction = MagicMock() + transaction.batch_ = lambda: _AsyncCM(batcher) + mock_prisma_client.db.tx = lambda timeout: _AsyncCM(transaction) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + end_user_costs: Dict[str, float] = {"u_b": 1.0, "u_a": 0.5} + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=0, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions=end_user_costs, + ) + calls = batcher.litellm_endusertable.upsert.call_args_list + ordered_ids = [c.kwargs["where"]["user_id"] for c in calls] + creates = [c.kwargs["data"]["create"] for c in calls] + pinned = { + "upsert_count": len(calls), + "ordered_ids": ordered_ids, + "first_create_keys": sorted(creates[0].keys()), + "first_create_user_id": creates[0]["user_id"], + "first_create_spend": creates[0]["spend"], + } + assert pinned == { + "upsert_count": 2, + "ordered_ids": ["u_a", "u_b"], + "first_create_keys": sorted(["user_id", "spend", "blocked"]), + "first_create_user_id": "u_a", + "first_create_spend": 0.5, + } + + +@pytest.mark.asyncio +async def test_update_end_user_spend_retries_on_connection_error( + mock_prisma_client: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + """``DB_CONNECTION_ERROR_TYPES`` failures should be retried with backoff; + once retries are exhausted, ``_raise_failed_update_spend_exception`` is + invoked and the original exception bubbles up. + """ + import httpx + import litellm.proxy.utils as utils_mod + + sleeps: list[float] = [] + + async def _fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + + err = httpx.ReadError("conn reset") + mock_prisma_client.db.tx = MagicMock(side_effect=err) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + with pytest.raises(httpx.ReadError): + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=1, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions={"u": 1.0}, + ) + assert sleeps == [1.0] + + +@pytest.mark.asyncio +async def test_update_end_user_spend_non_connection_error_raises_immediately( + mock_prisma_client: Any, +) -> None: + mock_prisma_client.db.tx = MagicMock(side_effect=RuntimeError("unknown")) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + with pytest.raises(RuntimeError, match="unknown"): + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=3, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions={"u": 1.0}, + ) + + +@pytest.mark.asyncio +async def test_update_spend_logs_writes_batches_via_create_many( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + logs = [make_spend_log_row(request_id=f"r{i}", spend=float(i)) for i in range(3)] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + logs_to_process=logs, + ) + kwargs = mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs + pinned = { + "calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + "data_len": len(kwargs["data"]), + "skip_duplicates": kwargs["skip_duplicates"], + "first_request_id": kwargs["data"][0]["request_id"], + } + assert pinned == { + "calls": 1, + "data_len": 3, + "skip_duplicates": True, + "first_request_id": "r0", + } + + +@pytest.mark.asyncio +async def test_update_spend_logs_uses_spend_logs_url_when_set( + mock_prisma_client: Any, + make_spend_log_row: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("SPEND_LOGS_URL", "http://writer.invalid") + writer = MagicMock() + writer.post = AsyncMock(return_value=MagicMock(status_code=200)) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + logs = [make_spend_log_row(request_id="r1")] + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=writer, + proxy_logging_obj=proxy_logging, + logs_to_process=logs, + ) + pinned = { + "post_calls": writer.post.await_count, + "url": writer.post.await_args.kwargs["url"], + "headers": writer.post.await_args.kwargs["headers"], + "create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + } + assert pinned == { + "post_calls": 1, + "url": "http://writer.invalid/spend/update", + "headers": {"Content-Type": "application/json"}, + "create_many_calls": 0, + } + + +@pytest.mark.asyncio +async def test_update_spend_logs_pops_logs_when_logs_to_process_is_none( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + mock_prisma_client.spend_log_transactions = [ + make_spend_log_row(request_id="a"), + make_spend_log_row(request_id="b"), + ] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert mock_prisma_client.spend_log_transactions == [] + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 1 + + +@pytest.mark.asyncio +async def test_update_spend_logs_failure_raises_after_retries( + mock_prisma_client: Any, + make_spend_log_row: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """When all retries exhaust the underlying DB error, the helper raises + via ``_raise_failed_update_spend_exception``. + """ + import httpx + import litellm.proxy.utils as utils_mod + + async def _fake_sleep(_: float) -> None: + return None + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( + side_effect=httpx.ReadError("network blip") + ) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + with pytest.raises(httpx.ReadError): + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=1, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + logs_to_process=[make_spend_log_row(request_id="r1")], + ) + + +def test_disable_spend_updates_reflects_general_settings( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The static method delegates to ``general_settings['disable_spend_updates']``; + flipping that value toggles the helper's return. + """ + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr( + proxy_server_mod, "general_settings", {"disable_spend_updates": True} + ) + pinned = { + "with_flag_true": ProxyUpdateSpend.disable_spend_updates(), + "type_is_bool": isinstance(ProxyUpdateSpend.disable_spend_updates(), bool), + "method_is_static": isinstance( + ProxyUpdateSpend.__dict__["disable_spend_updates"], staticmethod + ), + } + assert pinned == { + "with_flag_true": True, + "type_is_bool": True, + "method_is_static": True, + } + + +def test_disable_spend_updates_default_false_without_flag( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.setattr(proxy_server_mod, "general_settings", {}) + assert ProxyUpdateSpend.disable_spend_updates() is False + + +def test_disable_spend_updates_error_when_general_settings_unavailable( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy.proxy_server as proxy_server_mod + + monkeypatch.delattr(proxy_server_mod, "general_settings", raising=False) + with pytest.raises(ImportError): + ProxyUpdateSpend.disable_spend_updates() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py new file mode 100644 index 00000000000..5028b65705f --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_send_email.py @@ -0,0 +1,105 @@ +"""Pin ``send_email``. + +Symbols pinned here: + - ``send_email`` +""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from litellm.proxy.utils import send_email + + +@pytest.fixture(autouse=True) +def _smtp_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("SMTP_HOST", "smtp.invalid") + monkeypatch.setenv("SMTP_PORT", "2525") + monkeypatch.setenv("SMTP_USERNAME", "u") + monkeypatch.setenv("SMTP_PASSWORD", "p") + monkeypatch.setenv("SMTP_SENDER_EMAIL", "from@invalid") + monkeypatch.setenv("SMTP_TLS", "True") + + +@pytest.mark.asyncio +async def test_send_email_dispatches_via_smtp(in_memory_smtp: Any) -> None: + await send_email( + receiver_email="to@invalid", + subject="Hello", + html="

body

", + ) + assert len(in_memory_smtp.sent) == 1 + sent = in_memory_smtp.sent[0] + pinned = { + "from_addr": sent.from_addr, + "to_addrs": sent.to_addrs, + "subject": sent.subject, + "starttls": sent.starttls_called, + "login": sent.login_args, + } + assert pinned == { + "from_addr": "from@invalid", + "to_addrs": "to@invalid", + "subject": "Hello", + "starttls": True, + "login": ("u", "p"), + } + assert "

body

" in sent.body + + +@pytest.mark.asyncio +async def test_send_email_skips_starttls_when_disabled( + in_memory_smtp: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("SMTP_TLS", "False") + await send_email( + receiver_email="to@invalid", + subject="Hi", + html="

x

", + ) + assert in_memory_smtp.sent[0].starttls_called is False + + +@pytest.mark.asyncio +async def test_send_email_error_missing_sender_email( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("SMTP_SENDER_EMAIL", raising=False) + with pytest.raises(ValueError, match="SMTP_SENDER_EMAIL"): + await send_email( + receiver_email="x@y", subject="s", html="

h

" + ) + + +@pytest.mark.asyncio +async def test_send_email_error_missing_receiver() -> None: + with pytest.raises(ValueError, match="receiver email"): + await send_email(receiver_email=None, subject="s", html="

h

") + + +@pytest.mark.asyncio +async def test_send_email_error_missing_subject() -> None: + with pytest.raises(ValueError, match="subject"): + await send_email(receiver_email="x@y", subject=None, html="

h

") + + +@pytest.mark.asyncio +async def test_send_email_error_missing_html() -> None: + with pytest.raises(ValueError, match="HTML"): + await send_email(receiver_email="x@y", subject="s", html=None) + + +@pytest.mark.asyncio +async def test_send_email_smtp_failure_is_swallowed( + in_memory_smtp: Any, +) -> None: + """SMTP send_message errors are caught and logged; ``send_email`` itself + does not raise so a failing email never blocks the proxy. + """ + in_memory_smtp.raise_on_send = RuntimeError("smtp boom") + await send_email( + receiver_email="to@invalid", subject="Hi", html="

x

" + ) + assert in_memory_smtp.sent == [] diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py new file mode 100644 index 00000000000..a0b3af54750 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -0,0 +1,360 @@ +"""Pin module-level spend functions. + +Symbols pinned here: + - ``update_spend`` + - ``update_daily_tag_spend`` + - ``update_spend_logs_job`` + - ``_monitor_spend_logs_queue`` + - ``_raise_failed_update_spend_exception`` +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Dict, List +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy.utils import ( + _monitor_spend_logs_queue, + _raise_failed_update_spend_exception, + update_daily_tag_spend, + update_spend, + update_spend_logs_job, +) + + +@pytest.mark.asyncio +async def test_update_spend_invokes_writer_and_skips_empty_queue( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [] + + await update_spend( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + handler = proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler + pinned = { + "handler_called": handler.await_count, + "handler_kwargs": handler.await_args.kwargs, + "queue_empty": mock_prisma_client.spend_log_transactions, + } + assert pinned == { + "handler_called": 1, + "handler_kwargs": { + "prisma_client": mock_prisma_client, + "n_retry_times": 3, + "proxy_logging_obj": proxy_logging, + }, + "queue_empty": [], + } + + +@pytest.mark.asyncio +async def test_update_spend_processes_logs_when_queue_nonempty( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + import litellm.proxy.utils as utils_mod + + job_mock = AsyncMock() + monkeypatch.setattr(utils_mod, "update_spend_logs_job", job_mock) + + await update_spend( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert job_mock.await_count == 1 + + +@pytest.mark.asyncio +async def test_update_spend_handler_failure_propagates( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.db_update_spend_transaction_handler = AsyncMock( + side_effect=RuntimeError("handler down") + ) + with pytest.raises(RuntimeError, match="handler down"): + await update_spend( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + +@pytest.mark.asyncio +async def test_update_daily_tag_spend_redis_path_when_buffered( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + writer = MagicMock() + proxy_logging.db_spend_update_writer = writer + writer.redis_update_buffer = MagicMock() + writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + return_value=True + ) + writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer._commit_daily_tag_spend_to_db = AsyncMock() + + await update_daily_tag_spend( + prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging + ) + redis_kwargs = writer._commit_daily_tag_spend_to_db_with_redis.await_args.kwargs + pinned = { + "redis_calls": writer._commit_daily_tag_spend_to_db_with_redis.await_count, + "direct_calls": writer._commit_daily_tag_spend_to_db.await_count, + "redis_kwargs_keys": sorted(redis_kwargs.keys()), + "redis_n_retries": redis_kwargs["n_retry_times"], + } + assert pinned == { + "redis_calls": 1, + "direct_calls": 0, + "redis_kwargs_keys": sorted( + ["prisma_client", "n_retry_times", "proxy_logging_obj"] + ), + "redis_n_retries": 3, + } + + +@pytest.mark.asyncio +async def test_update_daily_tag_spend_direct_path_when_no_redis( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + writer = MagicMock() + proxy_logging.db_spend_update_writer = writer + writer.redis_update_buffer = MagicMock() + writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + return_value=False + ) + writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock() + writer._commit_daily_tag_spend_to_db = AsyncMock() + + await update_daily_tag_spend( + prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging + ) + assert writer._commit_daily_tag_spend_to_db.await_count == 1 + assert writer._commit_daily_tag_spend_to_db_with_redis.await_count == 0 + + +@pytest.mark.asyncio +async def test_update_daily_tag_spend_logs_and_swallows_errors( + mock_prisma_client: Any, +) -> None: + """A failure in the commit path is logged but not re-raised; this matches + the historical behavior of this site (see plain ``logger.error`` rather + than ``spend_log_error``). + """ + proxy_logging = MagicMock() + proxy_logging.db_spend_update_writer = MagicMock() + proxy_logging.db_spend_update_writer.redis_update_buffer = MagicMock() + proxy_logging.db_spend_update_writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock( + return_value=False + ) + proxy_logging.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock( + side_effect=RuntimeError("commit boom") + ) + await update_daily_tag_spend( + prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging + ) + + +@pytest.mark.asyncio +async def test_update_spend_logs_job_skips_when_queue_empty( + mock_prisma_client: Any, +) -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 0 + + +@pytest.mark.asyncio +async def test_update_spend_logs_job_processes_and_clears_queue( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [ + make_spend_log_row(request_id="r1"), + make_spend_log_row(request_id="r2"), + ] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + + # Stub auxiliary imports so the test focuses on the spend-logs write path. + import litellm.proxy.guardrails.usage_tracking as guard_mod + import litellm.proxy.db.spend_log_tool_index as tool_mod + + monkeypatch.setattr( + guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False + ) + monkeypatch.setattr( + tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False + ) + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + pinned = { + "create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count, + "queue_after": mock_prisma_client.spend_log_transactions, + "first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[ + "data" + ][0]["request_id"], + "skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[ + "skip_duplicates" + ], + } + assert pinned == { + "create_many_calls": 1, + "queue_after": [], + "first_data_request_id": "r1", + "skip_duplicates_set": True, + } + + +@pytest.mark.asyncio +async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty( + mock_prisma_client: Any, + make_spend_log_row: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm.proxy.utils as utils_mod + import litellm.constants as constants_mod + + monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False) + monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_SIZE_THRESHOLD", 1, raising=False) + proxy_logging = MagicMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + cancel_after = {"n": 0} + + async def _fake_job(*args: Any, **kwargs: Any) -> None: + cancel_after["n"] += 1 + if cancel_after["n"] >= 1: + raise asyncio.CancelledError() + + monkeypatch.setattr(utils_mod, "update_spend_logs_job", _fake_job) + + with pytest.raises(asyncio.CancelledError): + await _monitor_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert cancel_after["n"] == 1 + + +@pytest.mark.asyncio +async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off( + mock_prisma_client: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """An exception inside the loop is logged with backoff and the loop + continues running rather than crashing the monitor task. + """ + import litellm.proxy.utils as utils_mod + import litellm.constants as constants_mod + + monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False) + + sleep_count = {"n": 0} + + async def _short_sleep(_: float, *args: Any, **kwargs: Any) -> None: + sleep_count["n"] += 1 + if sleep_count["n"] >= 3: + raise asyncio.CancelledError() + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _short_sleep) + proxy_logging = MagicMock() + + bad_lock = MagicMock() + bad_lock.__aenter__ = AsyncMock(side_effect=RuntimeError("lock broken")) + bad_lock.__aexit__ = AsyncMock(return_value=False) + mock_prisma_client._spend_log_transactions_lock = bad_lock + + with pytest.raises(asyncio.CancelledError): + await _monitor_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + assert sleep_count["n"] == 3 + + +def test_raise_failed_update_spend_exception_emits_failure_handler() -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + async def _runner() -> Any: + try: + _raise_failed_update_spend_exception( + e=RuntimeError("boom"), + start_time=0.0, + proxy_logging_obj=proxy_logging, + ) + except RuntimeError as e: + return e + return None + + err = asyncio.run(_runner()) + pinned = { + "raised": str(err), + "failure_handler_called": proxy_logging.failure_handler.call_count, + "call_type": ( + proxy_logging.failure_handler.call_args.kwargs.get("call_type") + if proxy_logging.failure_handler.call_args + else None + ), + "non_blocking_in_traceback": ( + "Non-Blocking" + in proxy_logging.failure_handler.call_args.kwargs["traceback_str"] + if proxy_logging.failure_handler.call_args + else False + ), + } + assert pinned == { + "raised": "boom", + "failure_handler_called": 1, + "call_type": "update_spend", + "non_blocking_in_traceback": True, + } + + +def test_raise_failed_update_spend_exception_raises_original_error() -> None: + """Error path: the function always re-raises the original exception so + the caller can observe the failure. + """ + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + async def _runner() -> None: + _raise_failed_update_spend_exception( + e=ValueError("specific"), + start_time=0.0, + proxy_logging_obj=proxy_logging, + ) + + with pytest.raises(ValueError, match="specific"): + asyncio.run(_runner())