diff --git a/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py b/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py index 2896f76eec6..4396dae202b 100644 --- a/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py +++ b/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py @@ -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": []}, diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index c163656219f..5160dd3431e 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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") diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index a0d8ea48770..de5fc96c7c3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -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 ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 65c8c248a27..95067929ac1 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -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, ):