diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 27a8c7992fe..90ce821d62e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -4341,9 +4341,7 @@ class TestAgentMCPPermissions: stack.enter_context(patcher) stack.enter_context( patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here - MCPRequestHandler, - "_get_allowed_mcp_servers_for_key", - AsyncMock(return_value=["server-a", "server-b"]), + MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"]) ) ) stack.enter_context( @@ -4369,9 +4367,7 @@ class TestAgentMCPPermissions: await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth) stack.enter_context( patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here - MCPRequestHandler, - "_get_allowed_mcp_servers_for_key", - AsyncMock(return_value=["server-a", "server-b"]), + MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"]) ) ) stack.enter_context( @@ -4395,15 +4391,9 @@ class TestAgentMCPPermissions: with contextlib.ExitStack() as stack: for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager): stack.enter_context(patcher) - server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( - "server-a", user_api_key_auth - ) - server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( - "server-b", user_api_key_auth - ) - server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server( - "server-c", user_api_key_auth - ) + server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-a", user_api_key_auth) + server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-b", user_api_key_auth) + server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-c", user_api_key_auth) assert sorted(server_a_tools) == ["tool_direct", "tool_via_toolset"] assert server_b_tools == ["tool_b"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 3ccb388116b..d56f08c4e79 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -4535,9 +4535,7 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake( - self, respx_mock, monkeypatch, auth_type, is_byok, scheme - ): + async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4587,28 +4585,14 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - ( - httpx.Response(401, text="secret response content"), - "unhealthy", - "OpenAPI specification request failed (HTTP 401)", - ), + (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - ( - httpx.ConnectError("secret network details"), - "unhealthy", - "OpenAPI specification could not be loaded (ConnectError)", - ), - ( - httpx.Response(200, text="secret invalid JSON body"), - "unhealthy", - "OpenAPI specification could not be loaded (JSONDecodeError)", - ), + (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), + (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), ], ) - async def test_openapi_health_reports_safe_failures( - self, respx_mock, monkeypatch, failure, expected_status, expected_error - ): + async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -12735,12 +12719,7 @@ async def test_debug_resolution_matches_final_header_conflict_winner( from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser from litellm.proxy._experimental.mcp_server.mcp_debug import MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, MCPAuthDiagnostics from litellm.proxy._experimental.mcp_server.outbound_credentials import ( - ApiKeyConfig, - AuthorizationCodeConfig, - NoneConfig, - ServerSpec, - SharedKey, - UpstreamCredentialProvider, + ApiKeyConfig, AuthorizationCodeConfig, NoneConfig, ServerSpec, SharedKey, UpstreamCredentialProvider, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -12756,15 +12735,10 @@ async def test_debug_resolution_matches_final_header_conflict_winner( store = Store() context = MCPAuthenticatedUser(UserAPIKeyAuth(user_id="alice")) diagnostics = MCPAuthDiagnostics() - token = request_ctx.set( - RequestContext( - request_id=1, - meta=None, - session=MagicMock(), - lifespan_context=None, - request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}), - ) - ) + token = request_ctx.set(RequestContext( + request_id=1, meta=None, session=MagicMock(), lifespan_context=None, + request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}), + )) selected = { "stored": AuthorizationCodeConfig(), "static": ApiKeyConfig(key_source=SharedKey(value=SecretStr("static-token"))), @@ -12773,10 +12747,7 @@ async def test_debug_resolution_matches_final_header_conflict_winner( try: auth, remaining = await MCPServerManager()._resolve_v2_auth( server=MCPServer( - server_id="s", - name="s", - transport="http", - url="https://up.example/mcp", + server_id="s", name="s", transport="http", url="https://up.example/mcp", static_headers={"Authorization": "Bearer configured"}, ), spec=ServerSpec(server_id="s", resource="https://up.example/mcp", config=selected), @@ -12806,28 +12777,17 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li from litellm.types.mcp_server.mcp_server_manager import MCPServer diagnostics = MCPAuthDiagnostics() - token = request_ctx.set( - RequestContext( - request_id=1, - meta=None, - session=MagicMock(), - lifespan_context=None, - request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}), - ) - ) + token = request_ctx.set(RequestContext( + request_id=1, meta=None, session=MagicMock(), lifespan_context=None, + request=Request({"type": "http", MCP_AUTH_DIAGNOSTICS_SCOPE_KEY: diagnostics}), + )) try: server = MCPServer( - server_id="signed", - name="signed", - transport=transport, - url="https://up.example/mcp", - auth_type="aws_sigv4", - aws_access_key_id="AKIDEXAMPLE", - aws_secret_access_key="test-signing-secret", - aws_region_name="us-east-1", - aws_service_name="execute-api", - command="python", - args=["-c", "pass"], + server_id="signed", name="signed", transport=transport, + url="https://up.example/mcp", auth_type="aws_sigv4", + aws_access_key_id="AKIDEXAMPLE", aws_secret_access_key="test-signing-secret", + aws_region_name="us-east-1", aws_service_name="execute-api", + command="python", args=["-c", "pass"], ) client = await MCPServerManager()._create_mcp_client(server) if transport == "stdio": @@ -12846,16 +12806,12 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", - name="temporary", - url="https://idp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -12875,18 +12831,13 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", - name="stale", - url="https://idp.example.com/mcp", - transport=MCPTransport.http, - auth_type=auth_type, - oauth2_flow="authorization_code", + server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -12906,20 +12857,13 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", - name="replacement", - url="https://old.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - ) - replacement: Final = original.model_copy( - update={ - "url": "https://new.example.com/mcp", - "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - } + server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", ) + replacement: Final = original.model_copy(update={ + "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -12927,11 +12871,8 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", - name="publication", - url="https://old.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, + server_id="stale-publication", name="publication", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -12947,13 +12888,9 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", - name="temporary", - url="https://idp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -13054,9 +12991,7 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert ( - result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" - ) + assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -13068,11 +13003,8 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", - name="cancelled-cache", - transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", - auth_type=MCPAuth.none, + server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -13130,18 +13062,11 @@ class _DiscoveryUpstream: return httpx.Response(202) self.requests = (*self.requests, (payload.method, request.headers.get("authorization", ""))) if payload.method == "initialize": - return httpx.Response( - 200, - json={ - "jsonrpc": "2.0", - "id": payload.id, - "result": { - "protocolVersion": "2025-03-26", - "serverInfo": {"name": "discovery", "version": "1"}, - "capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}}, - }, - }, - ) + return httpx.Response(200, json={ + "jsonrpc": "2.0", "id": payload.id, + "result": {"protocolVersion": "2025-03-26", "serverInfo": {"name": "discovery", "version": "1"}, + "capabilities": {} if self.outcome == "unsupported" else {"prompts": {}, "resources": {}}}, + }) self.entered.set() await self.release.wait() if self.outcome == "failure": @@ -13149,15 +13074,12 @@ class _DiscoveryUpstream: if self.outcome == "cancelled": raise asyncio.CancelledError() if self.outcome == "rejected": - return httpx.Response( - 200, json={"jsonrpc": "2.0", "id": payload.id, "error": {"code": -32601, "message": "Unsupported"}} - ) + return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, + "error": {"code": -32601, "message": "Unsupported"}}) result: Final = { "prompts/list": {"prompts": [{"name": "example", "description": "original"}]}, "resources/list": {"resources": [{"name": "example", "uri": "test://example", "description": "original"}]}, - "resources/templates/list": { - "resourceTemplates": [{"name": "example", "uriTemplate": "test://{name}", "description": "original"}] - }, + "resources/templates/list": {"resourceTemplates": [{"name": "example", "uriTemplate": "test://{name}", "description": "original"}]}, "tools/list": {"tools": []}, }[payload.method] return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload.id, "result": result}) @@ -13168,9 +13090,7 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer( - server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http - ) + return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) @pytest.mark.asyncio @@ -13181,11 +13101,8 @@ async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None clock: Final = _DiscoveryClock() manager: Final = MCPServerManager(discovery_clock=clock) upstream: Final = _DiscoveryUpstream() - operation: Final = { - "prompts": manager.get_prompts_from_server, - "resources": manager.get_resources_from_server, - "templates": manager.get_resource_templates_from_server, - }[kind] + operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server, + "templates": manager.get_resource_templates_from_server}[kind] server: Final = _discovery_server() with respx.mock(base_url="https://discovery.example") as router: router.route().mock(side_effect=upstream.respond) @@ -13214,11 +13131,8 @@ async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: st manager: Final = MCPServerManager() upstream: Final = _DiscoveryUpstream() upstream.outcome = outcome - operation: Final = { - "prompts": manager.get_prompts_from_server, - "resources": manager.get_resources_from_server, - "templates": manager.get_resource_templates_from_server, - }[kind] + operation: Final = {"prompts": manager.get_prompts_from_server, "resources": manager.get_resources_from_server, + "templates": manager.get_resource_templates_from_server}[kind] with respx.mock(base_url="https://discovery.example") as router: router.route().mock(side_effect=upstream.respond) assert await operation(_discovery_server(), None) == [] @@ -13245,20 +13159,9 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_ assert len(await manager.get_prompts_from_server(server, user)) == 1 assert upstream.initializes == 1 for credential in ("first-secret", "second-secret", "first-secret"): - assert ( - len( - await manager.get_prompts_from_server( - server, first_user, extra_headers={"Authorization": credential} - ) - ) - == 1 - ) + assert len(await manager.get_prompts_from_server(server, first_user, extra_headers={"Authorization": credential})) == 1 assert upstream.initializes == 3 - assert {auth for method, auth in upstream.requests if method == "prompts/list"} == { - "", - "first-secret", - "second-secret", - } + assert {auth for method, auth in upstream.requests if method == "prompts/list"} == {"", "first-secret", "second-secret"} @pytest.mark.asyncio @@ -13270,9 +13173,7 @@ async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> N upstream.release.clear() with respx.mock(base_url="https://discovery.example") as router: router.route().mock(side_effect=upstream.respond) - tasks: Final = tuple( - asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) for _ in range(10) - ) + tasks: Final = tuple(asyncio.create_task(manager.get_prompts_from_server(_discovery_server(), None)) for _ in range(10)) await asyncio.wait_for(upstream.entered.wait(), timeout=5) tasks[0].cancel() with pytest.raises(asyncio.CancelledError): @@ -13321,9 +13222,7 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) assert upstream.initializes == 2 -@pytest.mark.parametrize( - "value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)) -) +@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl @@ -13441,15 +13340,9 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N source: Final = CredentialSource() managers: Final = (MCPServerManager(cred_provider=source), MCPServerManager(cred_provider=source)) server: Final = MCPServer( - server_id="discovery", - name="discovery", - url="https://discovery.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - client_id="discovery-client", - authorization_url="https://discovery.example/authorize", - token_url="https://discovery.example/token", + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="discovery-client", + authorization_url="https://discovery.example/authorize", token_url="https://discovery.example/token", ) user: Final = UserAPIKeyAuth(user_id="same-user", api_key="same-key") upstream: Final = _DiscoveryUpstream() @@ -13468,15 +13361,11 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N with respx.mock(base_url="https://discovery.example") as router: router.route().mock(side_effect=respond) for manager in managers: - assert [item.name for item in await manager.get_prompts_from_server(server, user)] == [ - "discovery-account-a" - ] + assert [item.name for item in await manager.get_prompts_from_server(server, user)] == ["discovery-account-a"] assert upstream.initializes == 2 source.token = "token-b" for manager in managers: - assert [item.name for item in await manager.get_prompts_from_server(server, user)] == [ - "discovery-account-b" - ] + assert [item.name for item in await manager.get_prompts_from_server(server, user)] == ["discovery-account-b"] assert upstream.initializes == 4 source.token = None for manager in managers: @@ -13503,15 +13392,9 @@ async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None store: Final = TokenStore() manager: Final = MCPServerManager(per_user_oauth_token_store=store) server: Final = MCPServer( - server_id="discovery", - name="discovery", - url="https://discovery.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - client_id="discovery-client", - authorization_url="https://discovery.example/authorize", - token_url="https://discovery.example/token", + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", client_id="discovery-client", + authorization_url="https://discovery.example/authorize", token_url="https://discovery.example/token", ) user: Final = UserAPIKeyAuth(user_id="requesting-user") upstream: Final = _DiscoveryUpstream() 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 61e11e72859..395ce68ec54 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 @@ -1886,15 +1886,9 @@ async def test_track_cost_callback_keeps_guardrail_cost_on_cache_hit(): } with ( - patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam - "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock - ) as mock_increment, - patch( # test-quality-ok: same function-body import, no injection seam - "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock - ), - patch( # test-quality-ok: same function-body import, no injection seam - "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, # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam + patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock), # test-quality-ok: same function-body import, no injection seam + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, # test-quality-ok: same function-body import, no injection seam ): mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()