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:
ryan-crabbe-berri 2026-08-29 14:44:00 -07:00
parent acf3ed7d9b
commit d7c0bc1e6d
4 changed files with 90 additions and 137 deletions

View file

@ -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": []},

View file

@ -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")

View file

@ -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
)

View file

@ -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,
):