Merge pull request #37833 from BerriAI/litellm_deflake_20260821

fix: roll up the open deflake fixes for the MCP logging queue, PTU rollup, license gate, and pricing test isolation
This commit is contained in:
Mateo Wang 2026-08-27 13:42:26 -07:00 committed by GitHub
commit 0441faadca
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 284 additions and 146 deletions

View file

@ -9253,6 +9253,7 @@ class ProxyStartupEvent:
prisma_client,
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
alert=_alert_ptu_rollup_failure,
router=llm_router,
)
scheduler.add_job(

View file

@ -14,7 +14,6 @@ and share the existing unique constraint.
import asyncio
import json
import sys
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from datetime import date, datetime, time, timedelta, timezone
@ -326,16 +325,6 @@ class _LoadedDeployments:
scanned_ids: frozenset[str]
def _running_router() -> object | None:
"""The proxy's router, or None outside a running proxy.
Read out of ``sys.modules`` rather than imported, so a rollup driven from a test or a
script does not pull the whole proxy server in behind it.
"""
proxy_server: Final = sys.modules.get("litellm.proxy.proxy_server")
return getattr(proxy_server, "llm_router", None) if proxy_server is not None else None
def _config_deployments(router: object | None, *, owned_by_db: frozenset[str]) -> tuple[_PTUDeployment, ...]:
"""Deployments the router holds that no ``LiteLLM_ProxyModelTable`` row owns.
@ -356,15 +345,17 @@ def _config_deployments(router: object | None, *, owned_by_db: frozenset[str]) -
)
async def _load_ptu_models(prisma_client: "PrismaClient") -> _LoadedDeployments:
async def _load_ptu_models(prisma_client: "PrismaClient", *, router: object | None) -> _LoadedDeployments:
"""Every deployment carrying valid manual PTU config, and every id the scan saw.
Reserved capacity is billed by the provider whichever file declared it, so a
deployment the proxy only knows from config.yaml accrues alongside the stored ones.
The router is handed in rather than read off the proxy module, so a run prices exactly
the deployments its caller declares and nothing a co-resident process left behind.
"""
rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many()
db_ids: Final = frozenset(model_id for row in rows if (model_id := str(getattr(row, "model_id", "") or "")))
config_records: Final = _config_deployments(_running_router(), owned_by_db=db_ids)
config_records: Final = _config_deployments(router, owned_by_db=db_ids)
models: Final = tuple(
parsed for parsed in (_parse_ptu_model(row) for row in (*rows, *config_records)) if parsed is not None
)
@ -380,6 +371,7 @@ async def run_ptu_flat_cost_rollup(
prisma_client: "PrismaClient",
target_date: date | None = None,
may_prune: bool = True,
router: object | None = None,
) -> RollupResult:
"""Rollup one UTC day of flat PTU cost across all PTU-configured model deployments.
@ -406,7 +398,7 @@ async def run_ptu_flat_cost_rollup(
date_str: Final = day.isoformat()
run_started: Final = datetime.now(timezone.utc)
loaded: Final = await _load_ptu_models(prisma_client)
loaded: Final = await _load_ptu_models(prisma_client, router=router)
ptu_models: Final = loaded.models
charges: Final = _aggregate_charges(ptu_models, day)
@ -527,6 +519,7 @@ async def _existing_sentinel_keys(
async def run_ptu_flat_cost_backfill(
prisma_client: "PrismaClient",
today: date | None = None,
router: object | None = None,
) -> BackfillResult:
"""Price the elapsed days of every PTU window that carry no sentinel row yet.
@ -546,7 +539,7 @@ async def run_ptu_flat_cost_backfill(
verbose_proxy_logger.warning("PTU backfill: prisma_client is None, skipping")
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
ptu_models: Final = (await _load_ptu_models(prisma_client)).models
ptu_models: Final = (await _load_ptu_models(prisma_client, router=router)).models
days: Final = _backfill_window(ptu_models, end)
if not days:
@ -591,6 +584,7 @@ async def run_scheduled_ptu_rollup(
pod_lock_manager: "PodLockManager | None" = None,
target_date: date | None = None,
alert: Callable[[str], Awaitable[None]] | None = None,
router: object | None = None,
) -> RollupResult | None:
"""Run the daily rollup under a cross-pod lock so only one proxy reconciles a day.
@ -615,7 +609,7 @@ async def run_scheduled_ptu_rollup(
return None
if pod_lock_manager is None or pod_lock_manager.redis_cache is None:
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router)
if not await pod_lock_manager.acquire_lock(cronjob_id=PTU_ROLLUP_JOB_ID, ttl=PTU_ROLLUP_LOCK_TTL_SECONDS):
if await _lock_is_held(pod_lock_manager):
@ -629,10 +623,10 @@ async def run_scheduled_ptu_rollup(
"PTU rollup: could not take the rollup lock and no other pod holds it, "
"running unguarded rather than skipping the day"
)
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False)
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=False, router=router)
try:
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True)
return await _run_and_alert(prisma_client, target_date=target_date, alert=alert, may_prune=True, router=router)
finally:
await pod_lock_manager.release_lock(cronjob_id=PTU_ROLLUP_JOB_ID)
@ -657,6 +651,7 @@ async def _run_and_alert(
target_date: date | None,
alert: "Callable[[str], Awaitable[None]] | None",
may_prune: bool = True,
router: object | None = None,
) -> RollupResult:
"""Reconcile the day, catch up any days left unpriced, and alert on charges that did not land.
@ -669,7 +664,9 @@ async def _run_and_alert(
explicit date means reconcile exactly that day, so it stays a single-day operation.
Its failure is contained: the day's own result is returned either way.
"""
result: Final = await run_ptu_flat_cost_rollup(prisma_client, target_date=target_date, may_prune=may_prune)
result: Final = await run_ptu_flat_cost_rollup(
prisma_client, target_date=target_date, may_prune=may_prune, router=router
)
if result.rows_failed:
await _deliver_alert(
alert,
@ -686,7 +683,7 @@ async def _run_and_alert(
"by the provider with nothing attributing it here. Extend the window, or retire the deployment.",
)
if target_date is None:
await _backfill_and_alert(prisma_client, alert=alert)
await _backfill_and_alert(prisma_client, alert=alert, router=router)
return result
@ -694,6 +691,7 @@ async def _backfill_and_alert(
prisma_client: "PrismaClient",
*,
alert: "Callable[[str], Awaitable[None]] | None",
router: object | None = None,
) -> None:
"""Catch up unpriced PTU days, alerting on charges that did not land.
@ -701,7 +699,7 @@ async def _backfill_and_alert(
caller whatever the catch-up pass does.
"""
try:
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client)
backfill: Final = await run_ptu_flat_cost_backfill(prisma_client, router=router)
except Exception as exc: # noqa: BLE001 # the catch-up pass must not fail the day's rollup
verbose_proxy_logger.error("PTU backfill: catch-up pass failed, the day's rollup still stands: %s", exc)
return

View file

@ -5,8 +5,9 @@ import json
from pathlib import Path
import re
import sys
import time
import tomllib
from typing import Dict, List, Optional, Set, Tuple
from typing import Callable, Dict, Final, List, Optional, Protocol, Set, Tuple
from packaging.requirements import Requirement
import requests
@ -37,6 +38,13 @@ DEFAULT_TRANSITIVE_PIN_PACKAGES = (
# of the identifier, not an operator.
_SPDX_OPERATOR_SPLIT = re.compile(r"\s+(?:OR|AND)\s+")
_SPDX_WITH_SUFFIX = re.compile(r"\s+WITH\s+.*", re.DOTALL)
_PYPI_FETCH_ATTEMPTS: Final[int] = 3
_PYPI_FETCH_BACKOFF_SECONDS: Final[float] = 0.5
class _HttpGet(Protocol):
def __call__(self, url: str, *, timeout: float) -> requests.Response:
...
@dataclass
@ -50,7 +58,10 @@ class PackageLicense:
class LicenseChecker:
def __init__(
self, config_file: Path = Path("./tests/code_coverage_tests/liccheck.ini")
self,
config_file: Path = Path("./tests/code_coverage_tests/liccheck.ini"),
http_get: Optional[_HttpGet] = None,
sleep: Optional[Callable[[float], None]] = None,
):
if not config_file.exists():
print(f"Error: Config file {config_file} not found")
@ -79,6 +90,8 @@ class LicenseChecker:
# Track package results
self.package_results: List[PackageLicense] = []
self._http_get = http_get
self._sleep = sleep
@staticmethod
def _normalize_package_name(package_name: str) -> str:
@ -123,21 +136,38 @@ class LicenseChecker:
last resort derives the license from the ``License :: OSI Approved ::
...`` trove classifiers.
"""
try:
url = f"https://pypi.org/pypi/{package_name}/{version}/json"
response = requests.get(url, timeout=10)
response.raise_for_status()
info = response.json().get("info", {}) or {}
return (
info.get("license_expression")
or info.get("license")
or self._license_from_classifiers(info.get("classifiers") or [])
)
except Exception as e:
print(
f"Warning: Failed to fetch license for {package_name} {version}: {str(e)}"
)
return None
url = f"https://pypi.org/pypi/{package_name}/{version}/json"
http_get = self._http_get if self._http_get is not None else requests.get
sleep = self._sleep if self._sleep is not None else time.sleep
for attempt in range(_PYPI_FETCH_ATTEMPTS):
try:
response = http_get(url, timeout=10)
response.raise_for_status()
info = response.json().get("info", {}) or {}
return (
info.get("license_expression")
or info.get("license")
or self._license_from_classifiers(info.get("classifiers") or [])
)
except Exception as error:
if self._is_retryable_pypi_error(error) and attempt < _PYPI_FETCH_ATTEMPTS - 1:
sleep(_PYPI_FETCH_BACKOFF_SECONDS)
continue
print(
f"Warning: Failed to fetch license for {package_name} {version}: {str(error)}"
)
return None
return None
@staticmethod
def _is_retryable_pypi_error(error: Exception) -> bool:
if isinstance(error, (requests.ConnectionError, requests.Timeout)):
return True
if not isinstance(error, requests.HTTPError) or error.response is None:
return False
status_code = error.response.status_code
return status_code == 429 or status_code >= 500
@staticmethod
def _license_from_classifiers(classifiers: List[str]) -> Optional[str]:

View file

@ -45,6 +45,22 @@ def setup_and_teardown():
asyncio.set_event_loop(None) # Remove the reference to the loop
@pytest.fixture(scope="function", autouse=True)
async def drain_logging_worker():
"""
The logging queue is bound to the running loop, so anything left queued when a test's loop
goes away is carried onto the next test's loop and fires against its callbacks.
"""
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
yield
try:
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.clear_queue(), timeout=10)
except asyncio.TimeoutError:
pass
def pytest_collection_modifyitems(config, items):
# Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests
custom_logger_tests = [

View file

@ -375,6 +375,9 @@ def isolate_litellm_state():
litellm.in_memory_llm_clients_cache.flush_cache()
image_handling_module.in_memory_cache.flush_cache()
_reset_module_level_aws_auth_caches()
# litellm.get_model_info() memoizes ModelInfo built from litellm.model_cost, so a
# test that rebinds the cost map leaves later tests pricing against the old map.
litellm_utils_module._invalidate_model_cost_lowercase_map()
# Clear all callback lists to prevent cross-test contamination
if hasattr(litellm, "callbacks"):
@ -418,6 +421,7 @@ def isolate_litellm_state():
litellm_utils_module._runtime_registered_model_cost.clear()
litellm_utils_module._runtime_registered_model_cost.update(original_runtime_registered_model_cost)
litellm_utils_module._invalidate_model_cost_lowercase_map()
for _router in tuple(litellm_router_module._live_routers):
litellm_router_module._live_routers.discard(_router)

View file

@ -1767,16 +1767,20 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts():
@pytest.mark.asyncio
async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan(monkeypatch):
async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan():
"""Staleness alone stops being evidence once two hosts hold different configuration: a
row this run never considered belongs to a deployment another host is pricing from its
own file, and sweeping it drops that charge."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-elsewhere", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
entry = _router_entry(model_id="cfg-here", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
await run_scheduled_ptu_rollup(
_prisma_for([], table),
pod_lock_manager=_pod_lock(acquired=True),
target_date=DAY,
router=_router_holding(entry),
)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-elsewhere") in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-here") in table.rows
@ -1784,7 +1788,7 @@ async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan(monkeypat
@pytest.mark.asyncio
async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged(monkeypatch):
async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged():
"""The accepted cost of bounding the prune, driven through the sequence that produces
it: charge the day while the deployment exists, remove it, run the day again. Nothing
scans it now, so nothing may judge its row, and the amount it was billed stands."""
@ -1793,18 +1797,19 @@ async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged(
live_row = _model_row(model_id="dep-live", model_info=ptu)
doomed_row = _model_row(model_id="dep-doomed", model_info=ptu)
charged_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-doomed")
monkeypatch.setattr(
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="cfg", model_info=dict(ptu)))
)
router = _router_holding(_router_entry(model_id="cfg", model_info=dict(ptu)))
await run_scheduled_ptu_rollup(
_prisma_for([live_row, doomed_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
_prisma_for([live_row, doomed_row], table),
pod_lock_manager=_pod_lock(acquired=True),
target_date=DAY,
router=router,
)
billed = table.rows[charged_key]["ptu_flat_cost"]
table.rows[charged_key]["updated_at"] = datetime(2020, 1, 1, tzinfo=timezone.utc)
await run_scheduled_ptu_rollup(
_prisma_for([live_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
_prisma_for([live_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY, router=router
)
assert table.rows[charged_key]["ptu_flat_cost"] == billed
@ -1843,7 +1848,7 @@ async def test_every_deployment_that_prices_is_inside_the_set_that_bounds_the_pr
table,
)
loaded = await ptu_rollup._load_ptu_models(prisma)
loaded = await ptu_rollup._load_ptu_models(prisma, router=None)
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
assert loaded.scanned_ids == {"dep-a", "dep-b", "dep-unpriced"}
@ -1859,7 +1864,7 @@ async def test_a_priced_deployment_is_in_the_bound_even_with_an_id_the_scan_skip
_FakeSentinelTable(),
)
loaded = await ptu_rollup._load_ptu_models(prisma)
loaded = await ptu_rollup._load_ptu_models(prisma, router=None)
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
@ -1873,13 +1878,13 @@ async def test_the_prune_splits_the_id_set_across_statements(monkeypatch):
table = _FakeSentinelTable()
ptu = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
deployments = [_model_row(model_id=f"dep-{n}", model_info=ptu) for n in range(4)]
monkeypatch.setattr(
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="dep-4", model_info=dict(ptu)))
)
table.seed("t", DAY, "dep-3", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
await run_scheduled_ptu_rollup(
_prisma_for(deployments, table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
_prisma_for(deployments, table),
pod_lock_manager=_pod_lock(acquired=True),
target_date=DAY,
router=_router_holding(_router_entry(model_id="dep-4", model_info=dict(ptu))),
)
chunks = [call["model"]["in"] for call in table.delete_many_calls]
@ -1912,140 +1917,123 @@ def _router_holding(*entries):
@pytest.mark.asyncio
async def test_a_config_declared_deployment_is_priced(monkeypatch):
async def test_a_config_declared_deployment_is_priced():
"""The whole point. A PTU deployment the proxy only knows from config.yaml is not in
LiteLLM_ProxyModelTable, so a DB-only scan bills the provider's reservation to nobody."""
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=_router_holding(entry))
assert [(m.model_id, m.model_name, m.team_id) for m in loaded.models] == [("cfg-1", "gpt-4o-ptu", "t")]
assert "cfg-1" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_a_database_backed_router_entry_is_not_counted_twice(monkeypatch):
async def test_a_database_backed_router_entry_is_not_counted_twice():
"""Every deployment loaded from the table is also in the router, flagged db_model. Pricing
both copies would write two charges for one reservation."""
row = _model_row(model_id="db-1", model_info=dict(_VALID_PTU))
mirrored = _router_entry(model_id="db-1", model_info={**_VALID_PTU, "db_model": True})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(mirrored))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_router_entry_sharing_an_id_with_the_table_is_priced_once(monkeypatch):
"""db_model is data the router carries rather than something this module controls, so the
id anti-join is what actually maps onto the failure: two charges under one id."""
row = _model_row(model_id="both-1", model_info=dict(_VALID_PTU))
unflagged = _router_entry(model_id="both-1", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(unflagged))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["both-1"]
@pytest.mark.asyncio
async def test_a_client_credential_clone_is_not_priced(monkeypatch):
"""Supplying an api_key on a request mints a clone of the deployment under a fresh id,
carrying the source's PTU config. Pricing it bills one reservation per distinct caller key."""
source = _router_entry(model_id="cfg-1", model_info=dict(_VALID_PTU))
clone = _router_entry(model_id="cfg-1-clone", model_info={**_VALID_PTU, "original_model_id": "cfg-1"})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(source, clone))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["cfg-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_without_ptu_config_is_scanned_but_not_priced(monkeypatch):
"""It has to stay in the scanned set or its leftover sentinel rows become unprunable."""
entry = _router_entry(model_id="cfg-plain", model_info={"team_id": "t"})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert loaded.models == ()
assert "cfg-plain" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_no_router_in_the_process_prices_the_database_alone(monkeypatch):
"""The rollup is importable and callable outside a running proxy."""
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: None)
loaded = await ptu_rollup._load_ptu_models(
_prisma_for([_model_row(model_id="db-1", model_info=dict(_VALID_PTU))], _FakeSentinelTable())
_prisma_for([row], _FakeSentinelTable()), router=_router_holding(mirrored)
)
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_is_charged_end_to_end(monkeypatch):
async def test_a_router_entry_sharing_an_id_with_the_table_is_priced_once():
"""db_model is data the router carries rather than something this module controls, so the
id anti-join is what actually maps onto the failure: two charges under one id."""
row = _model_row(model_id="both-1", model_info=dict(_VALID_PTU))
unflagged = _router_entry(model_id="both-1", model_info=dict(_VALID_PTU))
loaded = await ptu_rollup._load_ptu_models(
_prisma_for([row], _FakeSentinelTable()), router=_router_holding(unflagged)
)
assert [m.model_id for m in loaded.models] == ["both-1"]
@pytest.mark.asyncio
async def test_a_client_credential_clone_is_not_priced():
"""Supplying an api_key on a request mints a clone of the deployment under a fresh id,
carrying the source's PTU config. Pricing it bills one reservation per distinct caller key."""
source = _router_entry(model_id="cfg-1", model_info=dict(_VALID_PTU))
clone = _router_entry(model_id="cfg-1-clone", model_info={**_VALID_PTU, "original_model_id": "cfg-1"})
loaded = await ptu_rollup._load_ptu_models(
_prisma_for([], _FakeSentinelTable()), router=_router_holding(source, clone)
)
assert [m.model_id for m in loaded.models] == ["cfg-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_without_ptu_config_is_scanned_but_not_priced():
"""It has to stay in the scanned set or its leftover sentinel rows become unprunable."""
entry = _router_entry(model_id="cfg-plain", model_info={"team_id": "t"})
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=_router_holding(entry))
assert loaded.models == ()
assert "cfg-plain" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_no_router_in_the_process_prices_the_database_alone():
"""The rollup is importable and callable outside a running proxy."""
loaded = await ptu_rollup._load_ptu_models(
_prisma_for([_model_row(model_id="db-1", model_info=dict(_VALID_PTU))], _FakeSentinelTable()), router=None
)
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_is_charged_end_to_end():
"""Through the scheduled entry point, so the charge lands in a sentinel row rather than
stopping at the loader."""
table = _FakeSentinelTable()
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
await run_scheduled_ptu_rollup(
_prisma_for([], table),
pod_lock_manager=_pod_lock(acquired=True),
target_date=DAY,
router=_router_holding(entry),
)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-1") in table.rows
@pytest.mark.asyncio
async def test_a_stale_database_backed_router_entry_is_not_treated_as_config(monkeypatch):
async def test_a_stale_database_backed_router_entry_is_not_treated_as_config():
"""The reconcile can leave a deployment on the router after its row is gone. The id
anti-join cannot see that one, so the flag is what keeps it from being priced as though
config.yaml had declared it."""
stale = _router_entry(model_id="db-gone", model_info={**_VALID_PTU, "db_model": True})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(stale))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=_router_holding(stale))
assert loaded.models == ()
def test_the_router_lookup_reads_the_proxys_own_global():
"""Every other config test replaces this helper, so without one test driving the real
body a typo in the module path or the attribute name leaves the whole feature dead in
production with the suite still green."""
import sys
import types as _types
@pytest.mark.asyncio
async def test_a_router_left_on_the_proxy_module_is_not_scanned(monkeypatch):
"""A run scans the router its caller hands it and nothing else. Reading the proxy module's
global instead made every run depend on whatever else in the process had set one, which
is what a caller passing no router is asking not to happen."""
import litellm.proxy.proxy_server as proxy_server
assert ptu_rollup._running_router() is None or "litellm.proxy.proxy_server" in sys.modules
ambient = _router_holding(_router_entry(model_id="ambient-1", model_info=dict(_VALID_PTU)))
monkeypatch.setattr(proxy_server, "llm_router", ambient, raising=False)
sentinel = object()
stub = _types.SimpleNamespace(llm_router=sentinel)
real = sys.modules.get("litellm.proxy.proxy_server")
sys.modules["litellm.proxy.proxy_server"] = stub
try:
assert ptu_rollup._running_router() is sentinel
del stub.llm_router
assert ptu_rollup._running_router() is None
finally:
if real is None:
del sys.modules["litellm.proxy.proxy_server"]
else:
sys.modules["litellm.proxy.proxy_server"] = real
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()), router=None)
def test_the_router_lookup_returns_none_outside_a_proxy():
import sys
real = sys.modules.pop("litellm.proxy.proxy_server", None)
try:
assert ptu_rollup._running_router() is None
finally:
if real is not None:
sys.modules["litellm.proxy.proxy_server"] = real
assert loaded.models == ()
assert loaded.scanned_ids == frozenset()
def test_the_prune_filter_is_a_plain_dict():
@ -2074,7 +2062,7 @@ async def test_a_run_that_scanned_nothing_issues_no_delete_statements():
@pytest.mark.asyncio
async def test_the_catch_up_pass_reaches_a_config_declared_deployment(monkeypatch):
async def test_the_catch_up_pass_reaches_a_config_declared_deployment():
"""The catch-up shares the loader, so config deployments join it without being wired in.
That is what prices the elapsed days of a reservation declared before today."""
table = _FakeSentinelTable()
@ -2084,9 +2072,10 @@ async def test_the_catch_up_pass_reaches_a_config_declared_deployment(monkeypatc
model_id="cfg-back",
model_info={"ptu_count": 100, "cost_per_ptu_per_hour": 0.02, "team_id": "t", "ptu_effective_from": started},
)
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True))
await run_scheduled_ptu_rollup(
_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), router=_router_holding(entry)
)
charged = sorted(day for (_, day, _, model) in table.rows if model == "cfg-back")
yesterday = (now.date() - timedelta(days=1)).isoformat()

View file

@ -11245,6 +11245,35 @@ async def test_ptu_rollup_job_registered_at_startup(monkeypatch):
assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None
@pytest.mark.asyncio
async def test_ptu_rollup_job_hands_the_rollup_the_proxys_router(monkeypatch):
"""The rollup prices PTU deployments declared in config.yaml, which only the router
knows about. It takes the router as an argument, so nothing but this call site puts the
proxy's own router in front of it: without it that half of the feature is dead."""
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
from litellm.proxy.spend_tracking import ptu_flat_cost_rollup
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import PTU_ROLLUP_JOB_ID
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
calls = []
monkeypatch.setattr(
ptu_flat_cost_rollup,
"run_scheduled_ptu_rollup",
AsyncMock(side_effect=lambda *args, **kwargs: calls.append(kwargs)),
)
scheduler = await _run_scheduled_background_jobs()
import litellm.proxy.proxy_server as ps
router = MagicMock()
monkeypatch.setattr(ps, "llm_router", router)
await scheduler.get_job(PTU_ROLLUP_JOB_ID).func()
assert [call["router"] for call in calls] == [router]
@pytest.mark.asyncio
async def test_ptu_rollup_job_not_registered_without_opt_in(monkeypatch):
"""Without LITELLM_ENABLE_PTU_COST_ATTRIBUTION the rollup never runs, so no sentinel row

View file

@ -12,6 +12,8 @@ import os
import sys
from pathlib import Path
import requests
_CODE_COVERAGE_DIR = os.path.join(
os.path.dirname(os.path.abspath(__file__)), "..", "code_coverage_tests"
)
@ -122,6 +124,75 @@ def test_get_license_returns_none_on_request_failure(monkeypatch):
assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None
def test_get_license_retries_connection_error_then_resolves_license():
responses = iter(
(
requests.ConnectionError("connection reset"),
requests.ConnectionError("connection reset"),
_FakeResponse({"info": {"license_expression": "MIT"}}),
)
)
calls = []
sleeps = []
def _fake_get(url, timeout=None):
calls.append((url, timeout))
response = next(responses)
if isinstance(response, Exception):
raise response
return response
checker = check_licenses.LicenseChecker(
config_file=_LICCHECK_INI,
http_get=_fake_get,
sleep=sleeps.append,
)
assert checker.get_package_license_from_pypi("pkg", "1.0.0") == "MIT"
assert len(calls) == 3
assert len(sleeps) == 2
def test_get_license_does_not_retry_not_found_http_error():
response = requests.Response()
response.status_code = 404
calls = []
sleeps = []
def _fake_get(url, timeout=None):
calls.append((url, timeout))
raise requests.HTTPError("not found", response=response)
checker = check_licenses.LicenseChecker(
config_file=_LICCHECK_INI,
http_get=_fake_get,
sleep=sleeps.append,
)
assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None
assert len(calls) == 1
assert sleeps == []
def test_get_license_returns_none_after_connection_retry_limit():
calls = []
sleeps = []
def _fake_get(url, timeout=None):
calls.append((url, timeout))
raise requests.ConnectionError("connection reset")
checker = check_licenses.LicenseChecker(
config_file=_LICCHECK_INI,
http_get=_fake_get,
sleep=sleeps.append,
)
assert checker.get_package_license_from_pypi("pkg", "1.0.0") is None
assert len(calls) == 3
assert len(sleeps) == 2
# --------------------------------------------------------------------------
# is_license_acceptable: SPDX identifiers and compound expressions
# --------------------------------------------------------------------------