fix(budgets): narrow model access group spend counters to the served deployment

The database writer already intersects the auth-matched groups with the ones
the served deployment declares, but the live spend counters got the unnarrowed
set. A caller granted two pools that both cover a model group debited both
counters while only one row moved, so the in-memory ceiling could block a pool
its persisted spend never touched.

Narrow once at the callback so both consumers read the same set.
This commit is contained in:
ryan-crabbe-berri 2026-08-29 14:25:56 -07:00
parent 7a2f6c4cc2
commit acf3ed7d9b
3 changed files with 93 additions and 5 deletions

View file

@ -112,7 +112,7 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager:
return tx
def _get_llm_router():
def get_llm_router():
"""The proxy's router, or None outside a running proxy.
Injected rather than imported where it is used, so the savings computation stays
@ -371,7 +371,7 @@ class DBSpendUpdateWriter:
routing_decision=metadata.get("routing_decision"),
usage_object=usage_object_raw if isinstance(usage_object_raw, dict) else None,
model_id=payload.get("model_id"),
llm_router=_get_llm_router,
llm_router=get_llm_router,
cost_breakdown=metadata.get("cost_breakdown"),
recorded_autorouter_savings=metadata.get("autorouter_savings"),
)
@ -562,7 +562,7 @@ class DBSpendUpdateWriter:
request_model_access_groups=request_model_access_groups,
served_model_id=payload_copy.get("model_id"),
prisma_client=prisma_client,
router=_get_llm_router(),
router=get_llm_router(),
)
_agent_id_for_spend: Final = payload_copy.get("agent_id")
@ -2004,7 +2004,7 @@ class DBSpendUpdateWriter:
gateway_injected_cache=marks_gateway_injection(_metadata, payload.get("model_id")),
routing_decision=_metadata.get("routing_decision"),
model_id=payload.get("model_id"),
llm_router=_get_llm_router,
llm_router=get_llm_router,
usage_object=usage_obj,
cost_breakdown=_metadata.get("cost_breakdown"),
recorded_autorouter_savings=_metadata.get("autorouter_savings"),

View file

@ -21,6 +21,10 @@ from litellm.proxy.auth.auth_checks import (
log_db_metrics,
)
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.db.db_spend_update_writer import (
debitable_model_access_groups,
get_llm_router,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.spend_tracking.spend_log_error_logger import (
should_suppress_spend_log_tracebacks,
@ -260,7 +264,11 @@ class _ProxyDBLogger(CustomLogger):
sl_object=sl_object,
metadata=metadata,
)
model_access_groups: Final = get_request_model_access_groups(kwargs)
model_access_groups: Final = debitable_model_access_groups(
attributed=get_request_model_access_groups(kwargs),
served_model_id=sl_object.get("model_id") if sl_object is not None else None,
router=get_llm_router(),
)
if response_cost is not None:
user_api_key: Final = metadata.get("user_api_key", None)

View file

@ -1956,3 +1956,83 @@ async def test_track_cost_callback_charges_no_model_access_group_when_none_were_
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."""
def __init__(self, deployments):
self._deployments = deployments
def get_model_info(self, id):
if id not in self._deployments:
return None
return {"model_name": "premium-haiku", "model_info": {"id": id, "access_groups": list(self._deployments[id])}}
def _model_access_group_kwargs(granted, served_model_id):
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),
}
},
"stream": False,
"standard_logging_object": {"response_cost": 0.25, "model_id": served_model_id},
}
async def _run_callback_capturing_groups(kwargs, deployments):
logger = _ProxyDBLogger()
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
patch(
"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)),
):
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
mock_proxy_logging.db_spend_update_writer = MagicMock()
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
return mock_update.await_args.kwargs["model_access_groups"]
@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.
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(
kwargs=_model_access_group_kwargs(granted=["premium", "tier0"], served_model_id="deployment-premium"),
deployments={"deployment-premium": ["premium"], "deployment-tier0": ["tier0"]},
)
assert debited == ("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(
kwargs=_model_access_group_kwargs(granted=["premium", "tier0"], served_model_id="deployment-gone"),
deployments={"deployment-premium": ["premium"]},
)
assert debited == ("premium", "tier0")