mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
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:
parent
7a2f6c4cc2
commit
acf3ed7d9b
3 changed files with 93 additions and 5 deletions
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue