mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
commit
0441faadca
8 changed files with 284 additions and 146 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# --------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue