mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test(budgets): clear the test-quality violations this branch added
The four model access group callback tests now share one helper, so nine patches of proxy_server internals become three, and both mock-echo assertions go with them. The delete_access_group tests share a context manager for the same reason. test_group_exactly_at_its_max_budget_passes gained the assertion it was missing: it now proves the group reached the spend comparison, which a group skipped for a missing budget row would not. The route-allowed patch beside it was dead, so it is gone. What is left is suppressed with the collaborator each one cannot inject.
This commit is contained in:
parent
acf3ed7d9b
commit
d7c0bc1e6d
4 changed files with 90 additions and 137 deletions
|
|
@ -343,7 +343,9 @@ async def _enforce(
|
|||
cache: UserApiKeyCache | None = None,
|
||||
) -> list[str]:
|
||||
read, seen = _spend_reader(spend_by_counter_key or {})
|
||||
with patch("litellm.proxy.proxy_server.get_current_spend", read):
|
||||
# The check takes its client and cache as arguments, injected just below. get_current_spend is the
|
||||
# one collaborator it reaches by a lazy `from litellm.proxy.proxy_server import`, with no parameter.
|
||||
with patch("litellm.proxy.proxy_server.get_current_spend", read): # test-quality-ok: get_current_spend is lazily imported inside _model_access_group_max_budget_check and has no injection point
|
||||
await _model_access_group_max_budget_check(
|
||||
matched_model_access_groups=matched,
|
||||
prisma_client=prisma_client if prisma_client is not None else _RecordingPrismaClient(*rows),
|
||||
|
|
@ -363,12 +365,16 @@ async def test_group_under_its_max_budget_passes():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_group_exactly_at_its_max_budget_passes():
|
||||
"""The ceiling is inclusive, matching the tag check it mirrors; only spend strictly above it blocks."""
|
||||
await _enforce(
|
||||
"""The ceiling is inclusive, matching the tag check it mirrors; only spend strictly above it blocks.
|
||||
|
||||
Asserting the counter was read is what keeps this honest: a group that got skipped entirely,
|
||||
because its row never arrived or carried no budget, would also not raise.
|
||||
"""
|
||||
assert await _enforce(
|
||||
("tier-a",),
|
||||
_MagBudgetRow("tier-a", max_budget=10.0),
|
||||
spend_by_counter_key={MODEL_ACCESS_GROUP_COUNTER_KEY: 10.0},
|
||||
)
|
||||
) == [MODEL_ACCESS_GROUP_COUNTER_KEY]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -469,10 +475,11 @@ async def _common_checks_with_over_budget_group(*, skip_budget_checks: bool) ->
|
|||
read, _ = _spend_reader({MODEL_ACCESS_GROUP_COUNTER_KEY: 99.0})
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", cache),
|
||||
patch("litellm.proxy.proxy_server.get_current_spend", read),
|
||||
patch("litellm.proxy.auth.auth_checks._is_api_route_allowed", return_value=True),
|
||||
# common_checks resolves all three off the proxy_server module at call time; its signature
|
||||
# has no client, cache or spend-reader parameter to pass them through instead.
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client), # test-quality-ok: common_checks lazily imports prisma_client from proxy_server and takes no client parameter
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", cache), # test-quality-ok: common_checks lazily imports user_api_key_cache from proxy_server and takes no cache parameter
|
||||
patch("litellm.proxy.proxy_server.get_current_spend", read), # test-quality-ok: get_current_spend is lazily imported inside the budget check and has no injection point
|
||||
):
|
||||
return await common_checks(
|
||||
request_body={"model": "gpt-4o", "messages": []},
|
||||
|
|
|
|||
|
|
@ -1879,85 +1879,6 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request(
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_track_cost_callback_charges_the_model_access_groups_auth_stamped():
|
||||
"""Auth stamps the matched groups onto request metadata; the callback has to carry them through.
|
||||
|
||||
Without this hop nothing writes ``spend:model_access_group:*`` on the normal path, so with
|
||||
reservations disabled the budget check reads a counter no one maintains.
|
||||
"""
|
||||
logger = _ProxyDBLogger()
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"call_type": "acompletion",
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key": "hashed-key",
|
||||
"user_api_key_user_id": "user-1",
|
||||
MODEL_ACCESS_GROUP_METADATA_KEY: ["premium", "starter"],
|
||||
},
|
||||
},
|
||||
"standard_logging_object": {"response_cost": 0.25, "request_tags": None},
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock
|
||||
) as mock_increment_spend_counters,
|
||||
patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock),
|
||||
):
|
||||
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
|
||||
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
|
||||
await logger._PROXY_track_cost_callback(
|
||||
kwargs=kwargs,
|
||||
completion_response=None,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
mock_increment_spend_counters.assert_awaited_once()
|
||||
assert mock_increment_spend_counters.await_args.kwargs["model_access_groups"] == (
|
||||
"premium",
|
||||
"starter",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_track_cost_callback_charges_no_model_access_group_when_none_were_stamped():
|
||||
"""A request no budgeted group authorized must not debit anything."""
|
||||
logger = _ProxyDBLogger()
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"call_type": "acompletion",
|
||||
"litellm_params": {
|
||||
"metadata": {"user_api_key": "hashed-key", "user_api_key_user_id": "user-1"},
|
||||
},
|
||||
"standard_logging_object": {"response_cost": 0.25, "request_tags": None},
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock
|
||||
) as mock_increment_spend_counters,
|
||||
patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock),
|
||||
):
|
||||
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
|
||||
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
|
||||
await logger._PROXY_track_cost_callback(
|
||||
kwargs=kwargs,
|
||||
completion_response=None,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
mock_increment_spend_counters.assert_awaited_once()
|
||||
assert mock_increment_spend_counters.await_args.kwargs["model_access_groups"] == ()
|
||||
|
||||
|
||||
class _FakeDeploymentLookup:
|
||||
"""Deployment lookup returning the access groups each deployment declares."""
|
||||
|
||||
|
|
@ -1970,32 +1891,38 @@ class _FakeDeploymentLookup:
|
|||
return {"model_name": "premium-haiku", "model_info": {"id": id, "access_groups": list(self._deployments[id])}}
|
||||
|
||||
|
||||
def _model_access_group_kwargs(granted, served_model_id):
|
||||
def _model_access_group_kwargs(granted, served_model_id=None):
|
||||
metadata = {"user_api_key": "hashed-key", "user_api_key_user_id": "user-1"}
|
||||
if granted is not None:
|
||||
metadata[MODEL_ACCESS_GROUP_METADATA_KEY] = list(granted)
|
||||
return {
|
||||
"call_type": "acompletion",
|
||||
"model": "premium-haiku",
|
||||
"litellm_call_id": "test-call-id",
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key": "hashed-key",
|
||||
"user_api_key_user_id": "u-1",
|
||||
MODEL_ACCESS_GROUP_METADATA_KEY: list(granted),
|
||||
}
|
||||
},
|
||||
"litellm_params": {"metadata": metadata},
|
||||
"stream": False,
|
||||
"standard_logging_object": {"response_cost": 0.25, "model_id": served_model_id},
|
||||
"standard_logging_object": {"response_cost": 0.25, "request_tags": None, "model_id": served_model_id},
|
||||
}
|
||||
|
||||
|
||||
async def _run_callback_capturing_groups(kwargs, deployments):
|
||||
async def _groups_charged_by_the_callback(kwargs, deployments=None):
|
||||
"""The groups the callback hands the spend counters for one request.
|
||||
|
||||
The callback resolves ``proxy_logging_obj`` and the router by importing them off
|
||||
``proxy_server`` inside its own body, so there is no seam to inject either through.
|
||||
"""
|
||||
logger = _ProxyDBLogger()
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
|
||||
patch(
|
||||
patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj"
|
||||
) as mock_proxy_logging,
|
||||
patch( # test-quality-ok: the arguments to this call are the boundary under test
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback._update_database_and_spend_counters",
|
||||
new=AsyncMock(),
|
||||
) as mock_update,
|
||||
patch("litellm.proxy.proxy_server.llm_router", new=_FakeDeploymentLookup(deployments)),
|
||||
patch( # test-quality-ok: llm_router is a proxy_server global the callback reads lazily, no seam
|
||||
"litellm.proxy.proxy_server.llm_router", new=_FakeDeploymentLookup(deployments or {})
|
||||
),
|
||||
):
|
||||
mock_proxy_logging.failed_tracking_alert = AsyncMock()
|
||||
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
|
||||
|
|
@ -2012,6 +1939,30 @@ async def _run_callback_capturing_groups(kwargs, deployments):
|
|||
return mock_update.await_args.kwargs["model_access_groups"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_track_cost_callback_charges_the_model_access_groups_auth_stamped():
|
||||
"""Auth stamps the matched groups onto request metadata; the callback has to carry them through.
|
||||
|
||||
Without this hop nothing writes ``spend:model_access_group:*`` on the normal path, so with
|
||||
reservations disabled the budget check reads a counter no one maintains.
|
||||
"""
|
||||
charged = await _groups_charged_by_the_callback(
|
||||
kwargs=_model_access_group_kwargs(granted=["premium", "starter"]),
|
||||
)
|
||||
|
||||
assert charged == ("premium", "starter")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_track_cost_callback_charges_no_model_access_group_when_none_were_stamped():
|
||||
"""A request no budgeted group authorized must not debit anything."""
|
||||
charged = await _groups_charged_by_the_callback(
|
||||
kwargs=_model_access_group_kwargs(granted=None),
|
||||
)
|
||||
|
||||
assert charged == ()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_counters_only_debit_the_group_the_served_deployment_belongs_to():
|
||||
"""A caller granted two pools that both cover the model group only draws down the pool that served.
|
||||
|
|
@ -2019,20 +1970,20 @@ async def test_spend_counters_only_debit_the_group_the_served_deployment_belongs
|
|||
The database writer already narrows by served deployment, so passing the unnarrowed set to the
|
||||
live counters let one request block a pool the persisted spend never debited.
|
||||
"""
|
||||
debited = await _run_callback_capturing_groups(
|
||||
charged = await _groups_charged_by_the_callback(
|
||||
kwargs=_model_access_group_kwargs(granted=["premium", "tier0"], served_model_id="deployment-premium"),
|
||||
deployments={"deployment-premium": ["premium"], "deployment-tier0": ["tier0"]},
|
||||
)
|
||||
|
||||
assert debited == ("premium",)
|
||||
assert charged == ("premium",)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_counters_keep_every_granted_group_when_the_deployment_is_unknown():
|
||||
"""An unidentifiable deployment leaves the auth-time set standing, so nothing silently stops billing."""
|
||||
debited = await _run_callback_capturing_groups(
|
||||
charged = await _groups_charged_by_the_callback(
|
||||
kwargs=_model_access_group_kwargs(granted=["premium", "tier0"], served_model_id="deployment-gone"),
|
||||
deployments={"deployment-premium": ["premium"]},
|
||||
)
|
||||
|
||||
assert debited == ("premium", "tier0")
|
||||
assert charged == ("premium", "tier0")
|
||||
|
|
|
|||
|
|
@ -753,7 +753,29 @@ def _seed_budget(prisma, access_group, spend=0.0, budget_id="budget-seed", **bud
|
|||
|
||||
@contextmanager
|
||||
def _proxy(prisma):
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
|
||||
with patch( # test-quality-ok: the endpoints import proxy_server.prisma_client themselves; no parameter to inject
|
||||
"litellm.proxy.proxy_server.prisma_client", prisma
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _proxy_with_stubbed_reload(prisma):
|
||||
"""delete_access_group finishes by reloading the router and judging what it serves afterwards.
|
||||
Both collaborators it reaches for there are module globals it imports itself, so a fake can only
|
||||
get in by patching them; auth_cache and prisma are the ones with a real seam."""
|
||||
never_served_router = MagicMock()
|
||||
never_served_router.get_model_ids.return_value = []
|
||||
with (
|
||||
_proxy(prisma),
|
||||
patch( # test-quality-ok: live_model_ids_snapshot() reads the llm_router global; the endpoint takes no router
|
||||
"litellm.proxy.proxy_server.llm_router", never_served_router
|
||||
),
|
||||
patch( # test-quality-ok: the endpoint calls its module-level clear_cache import; there is no parameter for it
|
||||
"litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache",
|
||||
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
|
||||
),
|
||||
):
|
||||
yield
|
||||
|
||||
|
||||
|
|
@ -1073,16 +1095,7 @@ async def test_deleting_the_access_group_strips_deployments_before_dropping_the_
|
|||
_seed_budget(prisma, "prod-models", spend=3.0, max_budget=100.0)
|
||||
cache = _FakeAuthCache()
|
||||
|
||||
never_served_router = MagicMock()
|
||||
never_served_router.get_model_ids.return_value = []
|
||||
with (
|
||||
_proxy(prisma),
|
||||
patch("litellm.proxy.proxy_server.llm_router", never_served_router),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache",
|
||||
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
|
||||
),
|
||||
):
|
||||
with _proxy_with_stubbed_reload(prisma):
|
||||
response = await delete_access_group(
|
||||
access_group="prod-models", user_api_key_dict=_admin(), auth_cache=cache
|
||||
)
|
||||
|
|
@ -1171,16 +1184,7 @@ async def test_deleting_the_access_group_evicts_both_auth_cache_keys():
|
|||
_seed_budget(prisma, "prod-models", spend=3.0, max_budget=100.0)
|
||||
cache = _FakeAuthCache(journal)
|
||||
|
||||
never_served_router = MagicMock()
|
||||
never_served_router.get_model_ids.return_value = []
|
||||
with (
|
||||
_proxy(prisma),
|
||||
patch("litellm.proxy.proxy_server.llm_router", never_served_router),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache",
|
||||
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
|
||||
),
|
||||
):
|
||||
with _proxy_with_stubbed_reload(prisma):
|
||||
await delete_access_group(access_group="prod-models", user_api_key_dict=_admin(), auth_cache=cache)
|
||||
|
||||
_assert_evicted_after_write(journal, "prod-models", "access_group_budget.delete:prod-models")
|
||||
|
|
@ -1198,16 +1202,7 @@ async def test_deleting_an_access_group_that_never_had_a_budget_still_evicts():
|
|||
prisma = _FakePrismaClient(journal, deployments=[_deployment()])
|
||||
cache = _FakeAuthCache(journal)
|
||||
|
||||
never_served_router = MagicMock()
|
||||
never_served_router.get_model_ids.return_value = []
|
||||
with (
|
||||
_proxy(prisma),
|
||||
patch("litellm.proxy.proxy_server.llm_router", never_served_router),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache",
|
||||
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
|
||||
),
|
||||
):
|
||||
with _proxy_with_stubbed_reload(prisma):
|
||||
response = await delete_access_group(
|
||||
access_group="prod-models", user_api_key_dict=_admin(), auth_cache=cache
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3050,7 +3050,7 @@ async def test_model_access_group_counter_blocks_a_request_over_the_group_budget
|
|||
prisma_client = _ModelAccessGroupBudgetPrisma(premium=1.0)
|
||||
valid_token = UserAPIKeyAuth(api_key="hashed", token="tok", matched_model_access_groups=["premium"])
|
||||
|
||||
with patch(
|
||||
with patch( # test-quality-ok: reserve_budget_for_request takes no estimator, so pinning the estimate needs this attribute
|
||||
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
|
||||
return_value=0.5,
|
||||
):
|
||||
|
|
@ -3081,7 +3081,7 @@ async def _cache_model_access_group_budget(key_cache, group, spend, max_budget=N
|
|||
|
||||
async def _reserve_for_model_access_groups(key_cache, groups, estimate):
|
||||
"""Reserve against the given groups, whose rows are already cached, so nothing hits the DB."""
|
||||
with patch(
|
||||
with patch( # test-quality-ok: reserve_budget_for_request takes no estimator, so pinning the estimate needs this attribute
|
||||
"litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost",
|
||||
return_value=estimate,
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue