From 26611dd4bd953b4180a55d219d6f92fbd3955ca4 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Thu, 29 Jan 2026 16:40:55 -0800 Subject: [PATCH 01/29] fix: dead code cleanup in MCP server error handler raise e makes the error logging and JSON error response below it unreachable --- litellm/proxy/_experimental/mcp_server/server.py | 1 - 1 file changed, 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 6d54c3871e5..545d5956bd2 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -1898,7 +1898,6 @@ if MCP_AVAILABLE: await session_manager.handle_request(scope, receive, send) except Exception as e: - raise e verbose_logger.exception(f"Error handling MCP request: {e}") # Instead of re-raising, try to send a graceful error response try: From f9e8f8712b1b4eb8b619a9ff18630d84bfdfcc6e Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Tue, 3 Feb 2026 16:51:45 -0800 Subject: [PATCH 02/29] fix: add cache invalidation for _cached_get_model_group_info on deployment changes _cached_get_model_group_info uses @lru_cache but had no invalidation, causing stale model group info (TPM/RPM limits) after dynamic deployment changes. Add cache_clear() at all 5 model_list mutation sites. --- litellm/router.py | 12 ++++++ tests/test_litellm/test_router.py | 67 +++++++++++++++++++++++++++++++ 2 files changed, 79 insertions(+) diff --git a/litellm/router.py b/litellm/router.py index d01c8443dab..5b72c3fb669 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -6055,6 +6055,7 @@ class Router: self.model_list = [] self.model_id_to_deployment_index_map = {} # Reset the index self.model_name_to_deployment_indices = {} # Reset the model_name index + self._invalidate_model_group_info_cache() # we add api_base/api_key each model so load balancing between azure/gpt on api_base1 and api_base2 works for model in original_model_list: @@ -6358,6 +6359,7 @@ class Router: """ idx = len(self.model_list) self.model_list.append(model) + self._invalidate_model_group_info_cache() # Update model_id index for O(1) lookup if model_id is not None: @@ -6405,6 +6407,7 @@ class Router: if removal_idx is not None: self.model_list.pop(removal_idx) + self._invalidate_model_group_info_cache() self._update_deployment_indices_after_removal( model_id=deployment_id, removal_idx=removal_idx ) @@ -6438,6 +6441,7 @@ class Router: if deployment_idx is not None: # Pop the item from the list first item = self.model_list.pop(deployment_idx) + self._invalidate_model_group_info_cache() self._update_deployment_indices_after_removal( model_id=id, removal_idx=deployment_idx ) @@ -7172,6 +7176,7 @@ class Router: """ # First populate the model_list self.model_list = [] + self._invalidate_model_group_info_cache() for _, model in enumerate(model_list): # Extract model_info from the model dict model_info = model.get("model_info", {}) @@ -7508,6 +7513,13 @@ class Router: return returned_models + def _invalidate_model_group_info_cache(self) -> None: + """Invalidate the cached model group info. + + Call this whenever self.model_list is modified to ensure the cache is rebuilt. + """ + self._cached_get_model_group_info.cache_clear() + def get_model_access_groups( self, model_name: Optional[str] = None, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 08ae804ea80..48aa435c3a9 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -925,6 +925,73 @@ def test_router_get_model_access_groups_team_only_models(): assert list(access_groups.keys()) == ["default-models"] +def test_cached_get_model_group_info(): + """ + Test that _cached_get_model_group_info caches results and + invalidates on deployment changes. + """ + from litellm.types.router import Deployment, LiteLLM_Params + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake"}, + "model_info": {"tpm": 1000, "rpm": 100}, + }, + ] + ) + + # First call should compute and cache + result1 = router._cached_get_model_group_info("gpt-4") + assert result1 is not None + assert result1.tpm == 1000 + + # Second call should hit cache (same object) + result2 = router._cached_get_model_group_info("gpt-4") + assert result1 is result2 + + # Add a deployment — cache should be invalidated + router.add_deployment( + Deployment( + model_name="gpt-4", + litellm_params=LiteLLM_Params(model="gpt-4", api_key="fake2"), + model_info={"tpm": 2000, "rpm": 200}, + ) + ) + result3 = router._cached_get_model_group_info("gpt-4") + assert result3 is not result2 + assert result3 is not None + assert result3.tpm == 3000 # 1000 + 2000 + + # Delete a deployment — cache should be invalidated + deployment_id = router.model_list[-1]["model_info"]["id"] + router.delete_deployment(id=deployment_id) + result4 = router._cached_get_model_group_info("gpt-4") + assert result4 is not result3 + assert result4 is not None + assert result4.tpm == 1000 + + # set_model_list — cache should be invalidated + router.set_model_list( + [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake"}, + "model_info": {"tpm": 5000}, + }, + ] + ) + result5 = router._cached_get_model_group_info("gpt-4") + assert result5 is not result4 + assert result5 is not None + assert result5.tpm == 5000 + + # Verify cache still works after invalidation + result6 = router._cached_get_model_group_info("gpt-4") + assert result5 is result6 + + @pytest.mark.asyncio async def test_acompletion_streaming_iterator(): """Test _acompletion_streaming_iterator for normal streaming and fallback behavior.""" From 6743d20de262338750f9b9046606222812915877 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Thu, 5 Feb 2026 10:50:45 -0800 Subject: [PATCH 03/29] docs: add trailing slash to /mcp endpoint URLs The /mcp endpoint requires a trailing slash because the MCP server is mounted as a sub-application using app.mount(). Starlette's mount behavior causes a 307 redirect from /mcp to /mcp/, which many MCP clients fail to handle. Updates documentation examples to use /mcp/ consistently. --- README.md | 2 +- docs/my-website/docs/mcp.md | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index 77adddf8978..e7701b5cf9d 100644 --- a/README.md +++ b/README.md @@ -203,7 +203,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \ { "mcpServers": { "LiteLLM": { - "url": "http://localhost:4000/mcp", + "url": "http://localhost:4000/mcp/", "headers": { "x-litellm-api-key": "Bearer sk-1234" } diff --git a/docs/my-website/docs/mcp.md b/docs/my-website/docs/mcp.md index d63b55ee29e..564a054aae2 100644 --- a/docs/my-website/docs/mcp.md +++ b/docs/my-website/docs/mcp.md @@ -632,7 +632,7 @@ import asyncio config = { "mcpServers": { "mcp_group": { - "url": "http://localhost:4000/mcp", + "url": "http://localhost:4000/mcp/", "headers": { "x-mcp-servers": "dev_group", # assume this gives access to github, zapier and deepwiki "x-litellm-api-key": "Bearer sk-1234", From 7a2c889ec20da38e6878cc62eeda286aa0d0d408 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Tue, 17 Feb 2026 16:10:29 -0800 Subject: [PATCH 04/29] perf: use cached _safe_get_request_headers instead of dict(request.headers) Replace 15 call sites across 9 files that called dict(request.headers) with _safe_get_request_headers(request) which caches the result on request.state. Mutation sites use .copy() to protect the shared cache. --- .../proxy/_experimental/mcp_server/rest_endpoints.py | 5 +++-- litellm/proxy/auth/user_api_key_auth.py | 2 +- litellm/proxy/custom_hooks/custom_ui_sso_hook.py | 3 ++- litellm/proxy/management_endpoints/scim/scim_v2.py | 3 ++- .../llm_passthrough_endpoints.py | 11 ++++++----- .../pass_through_endpoints/pass_through_endpoints.py | 7 +++++-- litellm/proxy/proxy_server.py | 5 +++-- .../proxy/vertex_ai_endpoints/langfuse_endpoints.py | 3 ++- 8 files changed, 24 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index aed81afd254..2fb0aee00ee 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -10,6 +10,7 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import ( ) from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.types.mcp import MCPAuth @@ -645,7 +646,7 @@ if MCP_AVAILABLE: return await _execute_with_mcp_client( new_mcp_server_request, _test_connection_operation, - raw_headers=dict(request.headers), + raw_headers=_safe_get_request_headers(request), ) @router.post("/test/tools/list") @@ -700,5 +701,5 @@ if MCP_AVAILABLE: _list_tools_operation, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, - raw_headers=dict(request.headers), + raw_headers=_safe_get_request_headers(request), ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 77d96e2d39c..2ec1036d6d5 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -559,7 +559,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, parent_otel_span=parent_otel_span, - request_headers=dict(request.headers), + request_headers=_safe_get_request_headers(request), ) is_proxy_admin = result["is_proxy_admin"] diff --git a/litellm/proxy/custom_hooks/custom_ui_sso_hook.py b/litellm/proxy/custom_hooks/custom_ui_sso_hook.py index cad0e62ff56..8bb6b274091 100644 --- a/litellm/proxy/custom_hooks/custom_ui_sso_hook.py +++ b/litellm/proxy/custom_hooks/custom_ui_sso_hook.py @@ -3,6 +3,7 @@ from fastapi_sso.sso.base import OpenID from litellm._logging import verbose_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers class CustomSSOLoginHandler(CustomLogger): @@ -18,7 +19,7 @@ class CustomSSOLoginHandler(CustomLogger): self, request: Request, ) -> OpenID: - request_headers_dict = dict(request.headers) + request_headers_dict = _safe_get_request_headers(request) verbose_logger.debug("inside custom ui sso sign in hook...") return OpenID( id=request_headers_dict.get("x-litellm-user-id") or "123", diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index e67e1eae745..4a5c6b04e5b 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -21,6 +21,7 @@ from typing_extensions import TypedDict import litellm from litellm._logging import verbose_proxy_logger +from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers from litellm._uuid import uuid from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ( @@ -724,7 +725,7 @@ async def get_service_provider_config(request: Request): "SCIM ServiceProviderConfig request: method=%s url=%s headers=%s", request.method, request.url, - dict(request.headers), + _safe_get_request_headers(request), ) meta = { "resourceType": "ServiceProviderConfig", diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 81144ad9f31..339523603c8 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -28,6 +28,7 @@ from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, + _safe_get_request_headers, _safe_set_request_parsed_body, get_form_data, get_request_body, @@ -60,7 +61,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": dict(request.headers), + "headers": _safe_get_request_headers(request), "cookies": request.cookies, "query_params": dict(request.query_params), } @@ -329,7 +330,7 @@ async def vllm_proxy_route( method=request.method, endpoint=endpoint, request_query_params=request.query_params, - request_headers=dict(request.headers), + request_headers=_safe_get_request_headers(request), stream=request_body.get("stream", False), content=None, data=None, @@ -1307,7 +1308,7 @@ async def azure_proxy_route( method=request.method, endpoint=endpoint, request_query_params=request.query_params, - request_headers=dict(request.headers), + request_headers=_safe_get_request_headers(request), stream=request_body.get("stream", False), content=None, data=None, @@ -1505,7 +1506,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict: Returns: dict: Headers dictionary with only allowed headers """ - incoming_headers = dict(request.headers) or {} + incoming_headers = _safe_get_request_headers(request) headers = {} for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: if header_name in incoming_headers: @@ -1621,7 +1622,7 @@ async def _prepare_vertex_auth_headers( if ( vertex_credentials is None or vertex_credentials.vertex_project is None ) and router_credentials is None: - headers = dict(request.headers) or {} + headers = _safe_get_request_headers(request).copy() headers_passed_through = True verbose_proxy_logger.debug( "default_vertex_config not set, incoming request headers %s", headers diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 56b513554a8..f7e0caed449 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -50,7 +50,10 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing -from litellm.proxy.common_utils.http_parsing_utils import _read_request_body +from litellm.proxy.common_utils.http_parsing_utils import ( + _read_request_body, + _safe_get_request_headers, +) from litellm.proxy.utils import get_server_root_path from litellm.secret_managers.main import get_secret_str from litellm.types.llms.custom_http import httpxSpecialProvider @@ -644,7 +647,7 @@ async def pass_through_request( # noqa: PLR0915 url = httpx.URL(target) headers = custom_headers headers = HttpPassThroughEndpointHelpers.forward_headers_from_request( - request_headers=dict(request.headers), + request_headers=_safe_get_request_headers(request).copy(), headers=headers, forward_headers=forward_headers, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e332dd17637..3b6a2143837 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -292,6 +292,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, + _safe_get_request_headers, check_file_size_under_limit, get_form_data, ) @@ -10167,7 +10168,7 @@ async def async_queue_request( data["proxy_server_request"] = { "url": str(request.url), "method": request.method, - "headers": dict(request.headers), + "headers": _safe_get_request_headers(request), "body": copy.copy(data), # use copy instead of deepcopy } @@ -10188,7 +10189,7 @@ async def async_queue_request( data["metadata"] = {} data["metadata"]["user_api_key"] = user_api_key_dict.api_key data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata - _headers = dict(request.headers) + _headers = _safe_get_request_headers(request).copy() _headers.pop( "authorization", None ) # do not store the original `sk-..` api key in the db diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index ce27c830f6f..bff770b7ea3 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -19,6 +19,7 @@ from fastapi import APIRouter, Request, Response import litellm from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( create_pass_through_route, @@ -32,7 +33,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": dict(request.headers), + "headers": _safe_get_request_headers(request), "cookies": request.cookies, "query_params": dict(request.query_params), } From e7175a52129ba05accc6a6f0a3ae85881de5d0a6 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Tue, 17 Feb 2026 17:04:03 -0800 Subject: [PATCH 05/29] perf: add request.state caching to _safe_get_request_headers Cache the dict(request.headers) result on request.state._cached_headers so subsequent calls within the same request return the cached dict instead of re-creating it each time. --- .../proxy/common_utils/http_parsing_utils.py | 22 ++++++++++++++----- 1 file changed, 17 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index e1bca6e905f..8d179a9caed 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -135,17 +135,29 @@ def _safe_set_request_parsed_body( def _safe_get_request_headers(request: Optional[Request]) -> dict: """ - [Non-Blocking] Safely get the request headers + [Non-Blocking] Safely get the request headers. + Caches the result on request.state to avoid re-creating dict(request.headers) per call. + + Warning: Callers must NOT mutate the returned dict — it is shared across + all callers within the same request via the cache. """ + if request is None: + return {} + cached = getattr(request.state, "_cached_headers", None) + if cached is not None: + return cached try: - if request is None: - return {} - return dict(request.headers) + headers = dict(request.headers) except Exception as e: verbose_proxy_logger.debug( "Unexpected error reading request headers - {}".format(e) ) - return {} + headers = {} + try: + request.state._cached_headers = headers + except Exception: + pass # request.state may not be available in all contexts + return headers def check_file_size_under_limit( From f0a542c55d92c6f0b0f37abc34e39d75c2d9b6c1 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Wed, 18 Feb 2026 09:44:35 -0800 Subject: [PATCH 06/29] fix: add .copy() to create_request_copy and tests for header caching Protect cached headers from mutation in create_request_copy sites and add unit tests for _safe_get_request_headers caching behavior. --- .../llm_passthrough_endpoints.py | 2 +- .../vertex_ai_endpoints/langfuse_endpoints.py | 2 +- .../common_utils/test_http_parsing_utils.py | 46 +++++++++++++++++++ 3 files changed, 48 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 339523603c8..028edea8c3f 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -61,7 +61,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": _safe_get_request_headers(request), + "headers": _safe_get_request_headers(request).copy(), "cookies": request.cookies, "query_params": dict(request.query_params), } diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index bff770b7ea3..627618387d5 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -33,7 +33,7 @@ def create_request_copy(request: Request): return { "method": request.method, "url": str(request.url), - "headers": _safe_get_request_headers(request), + "headers": _safe_get_request_headers(request).copy(), "cookies": request.cookies, "query_params": dict(request.query_params), } diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index af366b082a0..05d2ab5d796 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -761,3 +761,49 @@ async def test_request_body_with_html_script_tags(): f"Message content with HTML was modified during parsing: " f"expected={msg['content']!r}, got={result['messages'][2]['content']!r}" ) + + +def test_safe_get_request_headers_caches_on_request_state(): + """ + Test that _safe_get_request_headers caches the result on request.state + and returns the same object on subsequent calls. + """ + mock_request = MagicMock() + mock_request.headers = {"content-type": "application/json", "authorization": "Bearer sk-123"} + mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default + + # First call — should create and cache + result1 = _safe_get_request_headers(mock_request) + assert result1 == {"content-type": "application/json", "authorization": "Bearer sk-123"} + assert mock_request.state._cached_headers is result1 + + # Second call — should return the cached object (same identity) + result2 = _safe_get_request_headers(mock_request) + assert result2 is result1 + + +def test_safe_get_request_headers_none_request(): + """ + Test that _safe_get_request_headers returns empty dict for None request. + """ + result = _safe_get_request_headers(None) + assert result == {} + + +def test_safe_get_request_headers_copy_protects_cache(): + """ + Test that callers using .copy() before mutation do not corrupt the cache. + """ + mock_request = MagicMock() + mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"} + mock_request.state = MagicMock(spec=[]) + + original = _safe_get_request_headers(mock_request) + + # Simulate what mutation call sites do: copy then pop + mutable = _safe_get_request_headers(mock_request).copy() + mutable.pop("authorization", None) + + # Cache must be unaffected + assert "authorization" in _safe_get_request_headers(mock_request) + assert _safe_get_request_headers(mock_request) is original From 7414c09277b63a536b56b870c142fad5c5a1084b Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Thu, 19 Feb 2026 13:41:15 -0800 Subject: [PATCH 07/29] perf: skip throwaway Usage() construction in ModelResponse.__init__ Avoid constructing a default Usage() object that gets immediately overwritten by convert_to_model_response_object. Set usage=None instead; the real Usage is assigned via setattr later. Also fix Bedrock Qwen2/Qwen3 transform_response to assign a new Usage object instead of mutating a potentially missing one. --- .../amazon_qwen2_transformation.py | 10 +++-- .../amazon_qwen3_transformation.py | 10 +++-- litellm/types/utils.py | 2 +- .../test_convert_dict_to_chat_completion.py | 44 +++++++++++++++++++ 4 files changed, 57 insertions(+), 9 deletions(-) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py index c532d8ea27c..2abcc679eef 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen2_transformation.py @@ -11,6 +11,7 @@ from typing import Any, List, Optional import httpx +from litellm.types.utils import Usage from litellm.llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import ( AmazonQwen3Config, ) @@ -79,10 +80,11 @@ class AmazonQwen2Config(AmazonQwen3Config): # Set usage information if available in response if "usage" in response_data: usage_data = response_data["usage"] - if hasattr(model_response, 'usage'): - model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0) - model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0) - model_response.usage.total_tokens = usage_data.get("total_tokens", 0) + model_response.usage = Usage( + prompt_tokens=usage_data.get("prompt_tokens", 0), + completion_tokens=usage_data.get("completion_tokens", 0), + total_tokens=usage_data.get("total_tokens", 0), + ) return model_response diff --git a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py index b3a957ce0f8..12333623f51 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/amazon_qwen3_transformation.py @@ -10,6 +10,7 @@ from typing import Any, List, Optional import httpx +from litellm.types.utils import Usage from litellm.llms.base_llm.chat.transformation import BaseConfig from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, @@ -201,10 +202,11 @@ class AmazonQwen3Config(AmazonInvokeConfig, BaseConfig): # Set usage information if available in response if "usage" in response_data: usage_data = response_data["usage"] - if hasattr(model_response, 'usage'): - model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0) - model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0) - model_response.usage.total_tokens = usage_data.get("total_tokens", 0) + model_response.usage = Usage( + prompt_tokens=usage_data.get("prompt_tokens", 0), + completion_tokens=usage_data.get("completion_tokens", 0), + total_tokens=usage_data.get("total_tokens", 0), + ) return model_response diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 9228b25b03e..9d8d421a91a 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1825,7 +1825,7 @@ class ModelResponse(ModelResponseBase): else: usage = usage elif stream is None or stream is False: - usage = Usage() + usage = None # avoid constructing throwaway Usage; set by convert_to_model_response_object if hidden_params: self._hidden_params = hidden_params diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py index 3b2087d25e9..cb580c63e0d 100644 --- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py +++ b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py @@ -1246,3 +1246,47 @@ def test_convert_to_model_response_object_with_error_code_only(): _response_headers=None, convert_tool_call_to_json_mode=False, ) + + +def test_convert_to_model_response_object_default_usage_overwritten(): + """ + Regression test: convert_to_model_response_object must properly set Usage + on a ModelResponse that only has the default Usage from ModelResponse.__init__() + (i.e. no extra litellm.Usage() set via setattr beforehand). + + This validates the optimization of removing the redundant + `setattr(model_response, "usage", litellm.Usage())` in completion(). + """ + mr = ModelResponse() + # usage is not set by default (optimization: avoid constructing throwaway Usage) + assert not hasattr(mr, "usage") + + response_object = { + "id": "chatcmpl-usage-test", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hello"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 15, + "completion_tokens": 7, + "total_tokens": 22, + }, + "model": "gpt-4o", + } + + result = convert_to_model_response_object( + model_response_object=mr, + response_object=response_object, + stream=False, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert isinstance(result, ModelResponse) + assert result.usage.prompt_tokens == 15 + assert result.usage.completion_tokens == 7 + assert result.usage.total_tokens == 22 From 974392311db4daad7cbbbcc80a35e91538668583 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Thu, 19 Feb 2026 16:57:49 -0800 Subject: [PATCH 08/29] perf: short-circuit is_model_o_series_model with startswith before set lookup MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reorder the check to use str.startswith(tuple) first, which immediately returns False for non-o-series models (the common case), avoiding the genexpr + 198-element set lookup. Line profiling shows 4.75x speedup (13.6µs → 2.9µs per call, 2.45s → 0.52s across 180k calls). --- litellm/llms/openai/chat/o_series_transformation.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/litellm/llms/openai/chat/o_series_transformation.py b/litellm/llms/openai/chat/o_series_transformation.py index 30647f58687..6ef43ec5bfd 100644 --- a/litellm/llms/openai/chat/o_series_transformation.py +++ b/litellm/llms/openai/chat/o_series_transformation.py @@ -131,9 +131,7 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): def is_model_o_series_model(self, model: str) -> bool: model = model.split("/")[-1] # could be "openai/o3" or "o3" - return model in litellm.open_ai_chat_completion_models and any( - model.startswith(pfx) for pfx in ("o1", "o3", "o4") - ) + return model.startswith(("o1", "o3", "o4")) and model in litellm.open_ai_chat_completion_models @overload def _transform_messages( @@ -173,4 +171,4 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): else: return super()._transform_messages( messages, model, is_async=cast(Literal[False], False) - ) + ) \ No newline at end of file From c119adb6dccedc1da78b30d91fa552e174e88191 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 23 Feb 2026 23:06:27 -0800 Subject: [PATCH 09/29] [Fix] UI - Virtual Keys: restrict Edit Settings button to key owners Non-owner Internal Users could see and interact with the "Edit Settings" button in the key Settings tab for keys they don't own. The button was gated by `rolesWithWriteAccess.includes(userRole)` (role-only check) instead of `canModifyKey` (ownership-aware), unlike the Regenerate and Delete buttons which already used the correct check. Replace the condition with `canModifyKey` so the Edit Settings button follows the same proxy-admin / team-admin / key-owner logic as the other action buttons. Add tests covering all permission paths. --- .../templates/key_info_view.test.tsx | 108 +++++++++++++++--- .../components/templates/key_info_view.tsx | 2 +- 2 files changed, 90 insertions(+), 20 deletions(-) diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx index db43d732e0a..bbb212e49b7 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.test.tsx @@ -363,31 +363,101 @@ describe("KeyInfoView", () => { }); - it("should show edit button in settings tab when user has write access", async () => { - vi.mocked(useAuthorized).mockReturnValue({ - ...baseUseAuthorizedMock, - userRole: "Admin", + describe("'Edit Settings' button visibility in the Settings tab", () => { + const renderAndOpenSettingsTab = async (keyData = MOCK_KEY_DATA) => { + render( + {}} + keyId="test-key-id" + onKeyDataUpdate={() => {}} + teams={[]} + />, + ); + await waitFor(() => { + expect(screen.getByRole("tab", { name: /settings/i })).toBeInTheDocument(); + }); + await userEvent.click(screen.getByRole("tab", { name: /settings/i })); + }; + + it("should show the Edit Settings button when the user is a proxy admin for a key they do not own", async () => { + vi.mocked(useAuthorized).mockReturnValue({ + ...baseUseAuthorizedMock, + userId: "proxy-admin-user-id", + userRole: "proxy_admin", + }); + + await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "someone-else-id" }); + + expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument(); }); - render( - { }} - keyId={"test-key-id"} - onKeyDataUpdate={() => { }} - teams={[]} - />, - ); + it("should show the Edit Settings button when the user is the key owner", async () => { + vi.mocked(useAuthorized).mockReturnValue({ + ...baseUseAuthorizedMock, + userId: "owner-user-id", + userRole: "Internal User", + }); - await waitFor(() => { - const settingsTab = screen.getByRole("tab", { name: /settings/i }); - expect(settingsTab).toBeInTheDocument(); + await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "owner-user-id" }); + + expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument(); }); - const settingsTab = screen.getByRole("tab", { name: /settings/i }); - await userEvent.click(settingsTab); + it("should not show the Edit Settings button when an Internal User does not own the key", async () => { + vi.mocked(useAuthorized).mockReturnValue({ + ...baseUseAuthorizedMock, + userId: "non-owner-user-id", + userRole: "Internal User", + }); + + await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "owner-user-id" }); + + expect(screen.queryByRole("button", { name: /edit settings/i })).not.toBeInTheDocument(); + }); + + it("should not show the Edit Settings button when the user is an Internal Viewer even if they own the key", async () => { + vi.mocked(useAuthorized).mockReturnValue({ + ...baseUseAuthorizedMock, + userId: "owner-user-id", + userRole: "Internal Viewer", + }); + + await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, user_id: "owner-user-id" }); + + expect(screen.queryByRole("button", { name: /edit settings/i })).not.toBeInTheDocument(); + }); + + it("should show the Edit Settings button when the user is a team admin for the key's team", async () => { + const teamId = "test-team-id"; + const teamAdminUserId = "team-admin-user"; + vi.mocked(useTeams).mockReturnValue({ + teams: [ + { + team_id: teamId, + team_alias: "Test Team", + models: [], + max_budget: null, + budget_duration: null, + tpm_limit: null, + rpm_limit: null, + organization_id: "org-1", + created_at: "2025-01-01T00:00:00Z", + keys: [], + members_with_roles: [{ user_id: teamAdminUserId, role: "admin" }], + spend: 0, + }, + ], + setTeams: vi.fn(), + }); + vi.mocked(useAuthorized).mockReturnValue({ + ...baseUseAuthorizedMock, + userId: teamAdminUserId, + userRole: "user", + }); + + await renderAndOpenSettingsTab({ ...MOCK_KEY_DATA, team_id: teamId, user_id: "other-user-id" }); - await waitFor(() => { expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx index f2a12c90875..94ca90b9630 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx @@ -595,7 +595,7 @@ export default function KeyInfoView({
Key Settings - {!isEditing && userRole && rolesWithWriteAccess.includes(userRole) && ( + {!isEditing && canModifyKey && ( )}
From ef67b6b53363ea064deba1914568044b54020efa Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 24 Feb 2026 17:48:55 +0530 Subject: [PATCH 10/29] Add support for phase param --- litellm/types/responses/main.py | 3 + .../test_openai_responses_transformation.py | 316 +++++++++++++++++- 2 files changed, 318 insertions(+), 1 deletion(-) diff --git a/litellm/types/responses/main.py b/litellm/types/responses/main.py index 8f6333ff900..bda53bae082 100644 --- a/litellm/types/responses/main.py +++ b/litellm/types/responses/main.py @@ -6,6 +6,7 @@ from typing_extensions import Any, List, Optional, TypedDict from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject +Phase = Optional[Literal["commentary", "final_answer"]] # TODO: Once openai sdk has updated, we can remove this and use the openai sdk type class GenericResponseOutputItemContentAnnotation(BaseLiteLLMOpenAIResponseObject): """Annotation for content in a message""" @@ -35,6 +36,7 @@ class OutputFunctionToolCall(BaseLiteLLMOpenAIResponseObject): type: Optional[str] # "function_call" id: Optional[str] status: Literal["in_progress", "completed", "incomplete"] + phase: Phase = None class OutputImageGenerationCall(BaseLiteLLMOpenAIResponseObject): @@ -57,6 +59,7 @@ class GenericResponseOutputItem(BaseLiteLLMOpenAIResponseObject): status: str # "completed", "in_progress", etc. role: str # "assistant", "user", etc. content: List[OutputText] + phase: Phase = None class DeleteResponseResult(BaseLiteLLMOpenAIResponseObject): diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py index 7c08716c04c..1a5ab808f7b 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_transformation.py @@ -925,4 +925,318 @@ def test_get_supported_openai_params(): assert "temperature" in params assert "stream" in params assert "background" in params - assert "stream" in params \ No newline at end of file + assert "stream" in params + + +class TestPhaseParameter: + """Tests for the `phase` parameter on assistant output items (gpt-5.3-codex).""" + + def setup_method(self): + self.config = OpenAIResponsesAPIConfig() + self.model = "gpt-5.3-codex" + self.logging_obj = MagicMock() + + @staticmethod + def _make_output_text(text: str): + from litellm.types.responses.main import OutputText + + return OutputText(type="output_text", text=text, annotations=[]) + + def test_generic_response_output_item_accepts_phase_commentary(self): + from litellm.types.responses.main import GenericResponseOutputItem + + item = GenericResponseOutputItem( + type="message", + id="msg_001", + status="completed", + role="assistant", + content=[self._make_output_text("Thinking...")], + phase="commentary", + ) + assert item.phase == "commentary" + + def test_generic_response_output_item_accepts_phase_final_answer(self): + from litellm.types.responses.main import GenericResponseOutputItem + + item = GenericResponseOutputItem( + type="message", + id="msg_002", + status="completed", + role="assistant", + content=[self._make_output_text("The answer is 42.")], + phase="final_answer", + ) + assert item.phase == "final_answer" + + def test_generic_response_output_item_phase_defaults_to_none(self): + from litellm.types.responses.main import GenericResponseOutputItem + + item = GenericResponseOutputItem( + type="message", + id="msg_003", + status="completed", + role="assistant", + content=[self._make_output_text("Hello")], + ) + assert item.phase is None + + def test_output_function_tool_call_accepts_phase(self): + from litellm.types.responses.main import OutputFunctionToolCall + + item = OutputFunctionToolCall( + type="function_call", + id="fc_001", + arguments='{"query": "test"}', + call_id="call_001", + name="search", + status="completed", + phase="commentary", + ) + assert item.phase == "commentary" + + def test_input_passthrough_dict_preserves_phase(self): + """Dict input items (the normal HTTP flow) must preserve phase verbatim.""" + input_items = [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Hi"}], + }, + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "Preamble..."}], + "phase": "commentary", + }, + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "Done."}], + "phase": "final_answer", + }, + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "Neutral."}], + "phase": None, + }, + ] + + result = self.config._validate_input_param(input_items) + assert isinstance(result, list) + + assert "phase" not in result[0] + assert result[1]["phase"] == "commentary" + assert result[2]["phase"] == "final_answer" + assert result[3]["phase"] is None + + def test_input_passthrough_pydantic_preserves_non_null_phase(self): + """Pydantic input items must preserve non-null phase values.""" + from litellm.types.responses.main import GenericResponseOutputItem + + item = GenericResponseOutputItem( + type="message", + id="msg_010", + status="completed", + role="assistant", + content=[self._make_output_text("commentary")], + phase="commentary", + ) + + result = self.config._validate_input_param([item]) + assert isinstance(result, list) + assert result[0]["phase"] == "commentary" + + def test_response_parsing_preserves_phase_on_output(self): + """Non-streaming response must preserve phase on output items.""" + raw_json = { + "id": "resp_001", + "created_at": 1700000000, + "model": "gpt-5.3-codex", + "object": "response", + "status": "completed", + "output": [ + { + "type": "message", + "id": "msg_001", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "preamble"}], + "phase": "commentary", + }, + { + "type": "message", + "id": "msg_002", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "answer"}], + "phase": "final_answer", + }, + ], + "usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30}, + } + + response = ResponsesAPIResponse(**raw_json) + assert len(response.output) == 2 + + for idx, output_item in enumerate(response.output): + if isinstance(output_item, dict): + phase = output_item.get("phase") + else: + phase = getattr(output_item, "phase", None) + + expected = "commentary" if idx == 0 else "final_answer" + assert phase == expected, ( + f"output[{idx}] phase={phase!r}, expected {expected!r}" + ) + + def test_streaming_output_item_done_preserves_phase(self): + """OutputItemDoneEvent must preserve phase on its item.""" + from litellm.types.llms.openai import ( + OutputItemDoneEvent, + ResponsesAPIStreamEvents, + ) + + chunk = { + "type": "response.output_item.done", + "output_index": 0, + "sequence_number": 3, + "item": { + "type": "message", + "id": "msg_100", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "done"}], + "phase": "final_answer", + }, + } + + result = self.config.transform_streaming_response( + model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj + ) + + assert isinstance(result, OutputItemDoneEvent) + assert result.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE + assert getattr(result.item, "phase", None) == "final_answer" + + def test_streaming_output_item_added_preserves_phase(self): + """OutputItemAddedEvent must preserve phase on its item.""" + from litellm.types.llms.openai import ( + OutputItemAddedEvent, + ResponsesAPIStreamEvents, + ) + + chunk = { + "type": "response.output_item.added", + "output_index": 0, + "item": { + "type": "message", + "id": "msg_200", + "role": "assistant", + "phase": "commentary", + }, + } + + result = self.config.transform_streaming_response( + model=self.model, parsed_chunk=chunk, logging_obj=self.logging_obj + ) + + assert isinstance(result, OutputItemAddedEvent) + assert result.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + assert getattr(result.item, "phase", None) == "commentary" + + def test_streaming_response_completed_preserves_phase(self): + """ResponseCompletedEvent must preserve phase on output items inside the response.""" + completed_chunk = { + "type": "response.completed", + "response": { + "id": "resp_300", + "created_at": 1700000000, + "model": "gpt-5.3-codex", + "object": "response", + "status": "completed", + "output": [ + { + "type": "message", + "id": "msg_300", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "final"}], + "phase": "final_answer", + } + ], + "usage": { + "input_tokens": 5, + "output_tokens": 10, + "total_tokens": 15, + }, + }, + } + + result = self.config.transform_streaming_response( + model=self.model, + parsed_chunk=completed_chunk, + logging_obj=self.logging_obj, + ) + + assert result.type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + output_item = result.response.output[0] + if isinstance(output_item, dict): + assert output_item["phase"] == "final_answer" + else: + assert getattr(output_item, "phase", None) == "final_answer" + + def test_phase_roundtrip_output_to_input(self): + """Simulate full round-trip: parse response output, then send items back as input.""" + raw_json = { + "id": "resp_rt", + "created_at": 1700000000, + "model": "gpt-5.3-codex", + "object": "response", + "status": "completed", + "output": [ + { + "type": "message", + "id": "msg_rt1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "preamble"}], + "phase": "commentary", + }, + { + "type": "message", + "id": "msg_rt2", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "answer"}], + "phase": "final_answer", + }, + ], + "usage": {"input_tokens": 10, "output_tokens": 20, "total_tokens": 30}, + } + + response = ResponsesAPIResponse(**raw_json) + + input_items = [] + for item in response.output: + if isinstance(item, dict): + input_items.append(item) + else: + input_items.append( + item.model_dump() if hasattr(item, "model_dump") else dict(item) + ) + + input_items.append( + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "next question"}], + } + ) + + validated = self.config._validate_input_param(input_items) + assert isinstance(validated, list) + + assert validated[0]["phase"] == "commentary" + assert validated[1]["phase"] == "final_answer" + assert "phase" not in validated[2] \ No newline at end of file From ac720defc3d7ae338299833055474dedfe667387 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 24 Feb 2026 17:50:38 +0530 Subject: [PATCH 11/29] Add documentation related to phase --- docs/my-website/blog/gpt_5_3_codex/index.md | 145 ++++++++++++++++++++ 1 file changed, 145 insertions(+) create mode 100644 docs/my-website/blog/gpt_5_3_codex/index.md diff --git a/docs/my-website/blog/gpt_5_3_codex/index.md b/docs/my-website/blog/gpt_5_3_codex/index.md new file mode 100644 index 00000000000..321e573cccf --- /dev/null +++ b/docs/my-website/blog/gpt_5_3_codex/index.md @@ -0,0 +1,145 @@ +--- +slug: gpt_5_3_codex +title: "Day 0 Support: GPT-5.3-Codex" +date: 2026-02-24T10:00:00 +authors: + - name: Sameer Kankute + title: SWE @ LiteLLM (LLM Translation) + url: https://www.linkedin.com/in/sameer-kankute/ + image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg + - name: Krrish Dholakia + title: "CEO, LiteLLM" + url: https://www.linkedin.com/in/krish-d/ + image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg + - name: Ishaan Jaff + title: "CTO, LiteLLM" + url: https://www.linkedin.com/in/reffajnaahsi/ + image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg +description: "Day 0 support for GPT-5.3-Codex on LiteLLM, including phase parameter handling for Responses API." +tags: [openai, gpt-5.3-codex, codex, day 0 support] +hide_table_of_contents: false +--- + +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +LiteLLM now supports GPT-5.3-Codex on Day 0, including support for the new assistant `phase` metadata on Responses API output items. + +## Why `phase` matters for GPT-5.3-Codex + +`phase` appears on assistant output items and helps distinguish preamble/commentary turns from final closeout responses. + +Reference: [Phase parameter docs](https://developers-site-git-alphas-venusaur-api-openai.vercel.app/alphas/venusaur-api/phase-parameter) + +Supported values: +- `null` +- `"commentary"` +- `"final_answer"` + +Important: +- Persist assistant output items with `phase` exactly as returned. +- Send those assistant items back on the next turn. +- Do **not** add `phase` to user messages. + +## Docker Image + +```bash +docker pull ghcr.io/berriai/litellm:v1.81.3-stable.sonnet-4-6 +``` + +## Usage + + + + +**1. Setup config.yaml** + +```yaml +model_list: + - model_name: gpt-5.3-codex + litellm_params: + model: openai/gpt-5.3-codex +``` + +**2. Start the proxy** + +```bash +docker run -d \ + -p 4000:4000 \ + -e ANTHROPIC_API_KEY=$ANTHROPIC_API_KEY \ + -v $(pwd)/config.yaml:/app/config.yaml \ + ghcr.io/berriai/litellm:v1.81.3-stable.sonnet-4-6 \ + --config /app/config.yaml +``` + + +**3. Test it** + +```bash +curl -X POST "http://0.0.0.0:4000/v1/responses" \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer $LITELLM_KEY" \ + -d '{ + "model": "gpt-5.3-codex", + "input": "Write a Python script that checks if a number is prime." + }' +``` + + + + +## Python Example: Persist `phase` with OpenAI Client + LiteLLM Base URL + +```python +from openai import OpenAI + +client = OpenAI( + base_url="http://0.0.0.0:4000/v1", # LiteLLM Proxy + api_key="your-litellm-api-key", +) + +items = [] # Persist this per conversation/thread + + +def _item_get(item, key, default=None): + if isinstance(item, dict): + return item.get(key, default) + return getattr(item, key, default) + + +def run_turn(user_text: str): + global items + + # User message: no phase field + items.append( + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": user_text}], + } + ) + + resp = client.responses.create( + model="gpt-5.3-codex", + input=items, + ) + + # Persist assistant output items verbatim, including phase + for out_item in (resp.output or []): + items.append(out_item) + + # Optional: inspect latest phase for UI/telemetry routing + latest_phase = None + for out_item in reversed(resp.output or []): + if _item_get(out_item, "type") == "output_item.done" and _item_get(out_item, "phase") is not None: + latest_phase = _item_get(out_item, "phase") + break + + return resp, latest_phase +``` + +## Notes + +- Use `/v1/responses` for GPT Codex models. +- Preserve full assistant output history for best multi-turn behavior. +- If `phase` metadata is dropped during history reconstruction, output quality can degrade on long-running tasks. From 1c48d8fda7512b386558044cec5bfbe18e22ee5c Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 24 Feb 2026 23:37:09 +0530 Subject: [PATCH 12/29] Add gpt-5.3-codex in model cost map --- ...odel_prices_and_context_window_backup.json | 33 +++++++++++++++++++ model_prices_and_context_window.json | 33 +++++++++++++++++++ 2 files changed, 66 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3909ce4c8b0..00b960a9f7a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -20562,6 +20562,39 @@ "supports_tool_choice": true, "supports_vision": true }, + "gpt-5.3-codex": { + "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_priority": 3.5e-07, + "input_cost_per_token": 1.75e-06, + "input_cost_per_token_priority": 3.5e-06, + "litellm_provider": "openai", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 1.4e-05, + "output_cost_per_token_priority": 2.8e-05, + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": false, + "supports_tool_choice": true, + "supports_vision": true + }, "gpt-5-mini": { "cache_read_input_token_cost": 2.5e-08, "cache_read_input_token_cost_flex": 1.25e-08, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3909ce4c8b0..00b960a9f7a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -20562,6 +20562,39 @@ "supports_tool_choice": true, "supports_vision": true }, + "gpt-5.3-codex": { + "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_priority": 3.5e-07, + "input_cost_per_token": 1.75e-06, + "input_cost_per_token_priority": 3.5e-06, + "litellm_provider": "openai", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 1.4e-05, + "output_cost_per_token_priority": 2.8e-05, + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": false, + "supports_tool_choice": true, + "supports_vision": true + }, "gpt-5-mini": { "cache_read_input_token_cost": 2.5e-08, "cache_read_input_token_cost_flex": 1.25e-08, From c43a8dc84257f07ebbd2211a9f30e3fa8ddb92e5 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 10:17:35 -0800 Subject: [PATCH 13/29] feat(proxy): add warning/error level logging throughout spend tracking lifecycle Elevate silent debug-level and bare except:pass error paths to warning/error so spend tracking failures are visible in production logs. All new log messages are prefixed with "Spend tracking -" for easy filtering. Changes cover the full request-to-DB lifecycle: enqueue, in-memory flush, Redis buffer push/pop, DB commit, cache updates, spend log writes, and pod lock management. Also fixes a copy-paste bug in _update_team_cache that logged "end user" instead of "team". --- litellm/proxy/db/db_spend_update_writer.py | 95 ++++++++++++++++--- .../daily_spend_update_queue.py | 5 + .../db_transaction_queue/pod_lock_manager.py | 14 ++- .../redis_update_buffer.py | 38 ++++++-- .../spend_update_queue.py | 5 + litellm/proxy/proxy_server.py | 36 +++++-- litellm/proxy/utils.py | 19 +++- 7 files changed, 177 insertions(+), 35 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 03628fda47f..c87be7ac819 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -224,8 +224,16 @@ class DBSpendUpdateWriter: verbose_proxy_logger.debug("Runs spend update on all tables") except Exception: - verbose_proxy_logger.debug( - f"Error updating Prisma database: {traceback.format_exc()}" + verbose_proxy_logger.error( + "Spend tracking - update_database failed. All spend updates for this request will be lost. " + "response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s - %s", + response_cost, + token, + user_id, + team_id, + org_id, + end_user_id, + traceback.format_exc(), ) async def _update_key_db( @@ -295,9 +303,14 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.debug( - "\033[91m" - + f"Update User DB call failed to execute {str(e)}\n{traceback.format_exc()}" + verbose_proxy_logger.error( + "Spend tracking - failed to enqueue user spend update. " + "user_id=%s, end_user_id=%s, response_cost=%s - %s\n%s", + user_id, + end_user_id, + response_cost, + str(e), + traceback.format_exc(), ) async def _update_team_db( @@ -334,11 +347,23 @@ class DBSpendUpdateWriter: response_cost=response_cost, ) ) - except Exception: - pass + except Exception as e: + verbose_proxy_logger.error( + "Spend tracking - failed to enqueue team member spend update. " + "team_id=%s, user_id=%s, response_cost=%s - %s", + team_id, + user_id, + response_cost, + str(e), + ) except Exception as e: - verbose_proxy_logger.debug( - f"Update Team DB failed to execute - {str(e)}\n{traceback.format_exc()}" + verbose_proxy_logger.error( + "Spend tracking - failed to enqueue team spend update. " + "team_id=%s, response_cost=%s - %s\n%s", + team_id, + response_cost, + str(e), + traceback.format_exc(), ) raise e @@ -363,8 +388,13 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.debug( - f"Update Org DB failed to execute - {str(e)}\n{traceback.format_exc()}" + verbose_proxy_logger.error( + "Spend tracking - failed to enqueue org spend update. " + "org_id=%s, response_cost=%s - %s\n%s", + org_id, + response_cost, + str(e), + traceback.format_exc(), ) raise e @@ -411,8 +441,13 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.debug( - f"Update Tag DB failed to execute - {str(e)}\n{traceback.format_exc()}" + verbose_proxy_logger.error( + "Spend tracking - failed to enqueue tag spend update. " + "request_tags=%s, response_cost=%s - %s\n%s", + request_tags, + response_cost, + str(e), + traceback.format_exc(), ) raise e @@ -513,6 +548,17 @@ class DBSpendUpdateWriter: await self.redis_update_buffer.get_all_update_transactions_from_redis_buffer() ) if db_spend_update_transactions is not None: + verbose_proxy_logger.info( + "Spend tracking - committing spend updates from Redis to DB: " + "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d", + len(db_spend_update_transactions.get("key_list_transactions") or {}), + len(db_spend_update_transactions.get("user_list_transactions") or {}), + len(db_spend_update_transactions.get("team_list_transactions") or {}), + len(db_spend_update_transactions.get("org_list_transactions") or {}), + len(db_spend_update_transactions.get("end_user_list_transactions") or {}), + len(db_spend_update_transactions.get("team_member_list_transactions") or {}), + len(db_spend_update_transactions.get("tag_list_transactions") or {}), + ) await self._commit_spend_updates_to_db( prisma_client=prisma_client, n_retry_times=n_retry_times, @@ -583,7 +629,12 @@ class DBSpendUpdateWriter: daily_spend_transactions=daily_agent_spend_update_transactions, ) except Exception as e: - verbose_proxy_logger.error(f"Error committing spend updates: {e}") + verbose_proxy_logger.error( + "Spend tracking - failed to commit spend updates from Redis to DB. " + "Data already popped from Redis may be lost. Error: %s\n%s", + str(e), + traceback.format_exc(), + ) finally: await self.pod_lock_manager.release_lock( cronjob_id=DB_SPEND_UPDATE_JOB_NAME, @@ -608,6 +659,22 @@ class DBSpendUpdateWriter: db_spend_update_transactions = ( await self.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() ) + if any( + len(v) > 0 + for v in db_spend_update_transactions.values() + if isinstance(v, dict) + ): + verbose_proxy_logger.info( + "Spend tracking - committing spend updates to DB (no Redis buffer): " + "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d", + len(db_spend_update_transactions.get("key_list_transactions") or {}), + len(db_spend_update_transactions.get("user_list_transactions") or {}), + len(db_spend_update_transactions.get("team_list_transactions") or {}), + len(db_spend_update_transactions.get("org_list_transactions") or {}), + len(db_spend_update_transactions.get("end_user_list_transactions") or {}), + len(db_spend_update_transactions.get("team_member_list_transactions") or {}), + len(db_spend_update_transactions.get("tag_list_transactions") or {}), + ) await self._commit_spend_updates_to_db( prisma_client=prisma_client, n_retry_times=n_retry_times, diff --git a/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py b/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py index 5ba8fb13596..f47b694d44e 100644 --- a/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/daily_spend_update_queue.py @@ -86,6 +86,11 @@ class DailySpendUpdateQueue(BaseUpdateQueue): ) -> Dict[str, BaseDailySpendTransaction]: """Get all updates from the queue and return all updates aggregated by daily_transaction_key. Works for both user and team spend updates.""" updates = await self.flush_all_updates_from_in_memory_queue() + if len(updates) > 0: + verbose_proxy_logger.info( + "Spend tracking - flushed %d daily spend update items from in-memory queue", + len(updates), + ) aggregated_daily_spend_update_transactions = ( DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( updates diff --git a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py index bb5424b0e90..6f86e82cf29 100644 --- a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py +++ b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py @@ -80,6 +80,14 @@ class PodLockManager: ) self._emit_acquired_lock_event(cronjob_id, self.pod_id) return True + else: + verbose_proxy_logger.info( + "Spend tracking - pod %s could not acquire lock for cronjob_id=%s, " + "held by pod %s. Spend updates in Redis will wait for the leader pod to commit.", + self.pod_id, + cronjob_id, + current_value, + ) return False except Exception as e: verbose_proxy_logger.error( @@ -124,10 +132,12 @@ class PodLockManager: pod_id=self.pod_id, ) else: - verbose_proxy_logger.debug( - "Pod %s failed to release Redis lock for cronjob_id=%s", + verbose_proxy_logger.warning( + "Spend tracking - pod %s failed to release Redis lock for cronjob_id=%s. " + "Lock will expire after TTL=%ds.", self.pod_id, cronjob_id, + DEFAULT_CRON_JOB_LOCK_TTL_SECONDS, ) else: verbose_proxy_logger.debug( diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 37b42e26bc9..85e139b4de0 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -96,14 +96,29 @@ class RedisUpdateBuffer: list_of_transactions = [safe_dumps(transactions)] if self.redis_cache is None: return - current_redis_buffer_size = await self.redis_cache.async_rpush( - key=redis_key, - values=list_of_transactions, - ) - await self._emit_new_item_added_to_redis_buffer_event( - queue_size=current_redis_buffer_size, - service=service_type, - ) + try: + current_redis_buffer_size = await self.redis_cache.async_rpush( + key=redis_key, + values=list_of_transactions, + ) + verbose_proxy_logger.info( + "Spend tracking - pushed spend updates to Redis buffer. " + "redis_key=%s, buffer_size=%s", + redis_key, + current_redis_buffer_size, + ) + await self._emit_new_item_added_to_redis_buffer_event( + queue_size=current_redis_buffer_size, + service=service_type, + ) + except Exception as e: + verbose_proxy_logger.error( + "Spend tracking - failed to push spend updates to Redis (redis_key=%s). " + "Error: %s", + redis_key, + str(e), + ) + raise async def store_in_memory_spend_updates_in_redis( self, @@ -305,6 +320,13 @@ class RedisUpdateBuffer: if list_of_transactions is None: return None + verbose_proxy_logger.info( + "Spend tracking - popped %d spend update batches from Redis buffer (key=%s). " + "These items are now removed from Redis and must be committed to DB.", + len(list_of_transactions) if isinstance(list_of_transactions, list) else 1, + REDIS_UPDATE_BUFFER_KEY, + ) + # Parse the list of transactions from JSON strings parsed_transactions = self._parse_list_of_transactions(list_of_transactions) diff --git a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py index b41ff121622..3e059cf8c1f 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py @@ -31,6 +31,11 @@ class SpendUpdateQueue(BaseUpdateQueue): ) -> DBSpendUpdateTransactions: """Flush all updates from the queue and return all updates aggregated by entity type.""" updates = await self.flush_all_updates_from_in_memory_queue() + if len(updates) > 0: + verbose_proxy_logger.info( + "Spend tracking - flushed %d spend update items from in-memory queue", + len(updates), + ) verbose_proxy_logger.debug("Aggregating updates by entity type: %s", updates) return self.get_aggregated_db_spend_update_transactions(updates) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b3b6bf0ccf7..78730ee9d60 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1768,8 +1768,13 @@ async def update_cache( # noqa: PLR0915 ("{}:spend".format(litellm_proxy_admin_name), increment) ) except Exception as e: - verbose_proxy_logger.debug( - f"An error occurred updating user cache: {str(e)}\n\n{traceback.format_exc()}" + verbose_proxy_logger.warning( + "Spend tracking - failed to update user spend in cache. " + "Budget enforcement may use stale spend values. " + "user_id=%s, response_cost=%s - %s", + user_id, + response_cost, + str(e), ) ### UPDATE END-USER SPEND ### @@ -1806,8 +1811,13 @@ async def update_cache( # noqa: PLR0915 existing_spend_obj.spend = new_spend values_to_update_in_cache.append((_id, existing_spend_obj.json())) except Exception as e: - verbose_proxy_logger.exception( - f"An error occurred updating end user cache: {str(e)}" + verbose_proxy_logger.warning( + "Spend tracking - failed to update end user spend in cache. " + "Budget enforcement may use stale spend values. " + "end_user_id=%s, response_cost=%s - %s", + end_user_id, + response_cost, + str(e), ) ### UPDATE TEAM SPEND ### @@ -1848,8 +1858,13 @@ async def update_cache( # noqa: PLR0915 existing_spend_obj.spend = new_spend values_to_update_in_cache.append((_id, existing_spend_obj)) except Exception as e: - verbose_proxy_logger.exception( - f"An error occurred updating end user cache: {str(e)}" + verbose_proxy_logger.warning( + "Spend tracking - failed to update team spend in cache. " + "Budget enforcement may use stale spend values. " + "team_id=%s, response_cost=%s - %s", + team_id, + response_cost, + str(e), ) ### UPDATE TAG SPEND ### @@ -1894,8 +1909,13 @@ async def update_cache( # noqa: PLR0915 existing_tag_obj.spend = new_spend values_to_update_in_cache.append((cache_key, existing_tag_obj)) except Exception as e: - verbose_proxy_logger.exception( - f"An error occurred updating tag cache: {str(e)}" + verbose_proxy_logger.warning( + "Spend tracking - failed to update tag spend in cache. " + "Budget enforcement may use stale spend values. " + "tags=%s, response_cost=%s - %s", + tags, + response_cost, + str(e), ) if token is not None and response_cost is not None: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c4ff325db1f..f6613b5548f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4457,6 +4457,11 @@ class ProxyUpdateSpend: len(logs_to_process) : ] popped_batch = True + if len(logs_to_process) > 0: + verbose_proxy_logger.info( + "Spend tracking - processing %d spend logs for DB write", + len(logs_to_process), + ) start_time = time.time() try: for i in range(n_retry_times + 1): @@ -4503,9 +4508,17 @@ class ProxyUpdateSpend: f"{len(logs_to_process)} logs processed. Remaining in queue: {remaining_count}" ) break - except DB_CONNECTION_ERROR_TYPES: + except DB_CONNECTION_ERROR_TYPES as e: if i is None: i = 0 + verbose_proxy_logger.warning( + "Spend tracking - DB connection error writing spend logs, " + "retry %d/%d. logs_count=%d, error=%s", + i + 1, + n_retry_times, + len(logs_to_process), + str(e), + ) if i >= n_retry_times: raise await asyncio.sleep(2**i) @@ -4620,8 +4633,8 @@ async def update_spend_logs_job( logs_to_process=logs_to_process, ) except Exception as guardrail_tracking_err: - verbose_proxy_logger.debug( - "Guardrail usage tracking failed (non-fatal): %s", + verbose_proxy_logger.warning( + "Spend tracking - guardrail usage tracking failed (non-fatal): %s", guardrail_tracking_err, ) From aded14a55ac973c91559e49fcef4d267cbc17229 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 25 Feb 2026 01:04:12 +0530 Subject: [PATCH 14/29] Fix release version for gpt-5.3-codex --- docs/my-website/blog/gpt_5_3_codex/index.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/my-website/blog/gpt_5_3_codex/index.md b/docs/my-website/blog/gpt_5_3_codex/index.md index 321e573cccf..5c43d287bce 100644 --- a/docs/my-website/blog/gpt_5_3_codex/index.md +++ b/docs/my-website/blog/gpt_5_3_codex/index.md @@ -44,7 +44,7 @@ Important: ## Docker Image ```bash -docker pull ghcr.io/berriai/litellm:v1.81.3-stable.sonnet-4-6 +docker pull ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3 ``` ## Usage @@ -68,7 +68,7 @@ docker run -d \ -p 4000:4000 \ -e ANTHROPIC_API_KEY=$ANTHROPIC_API_KEY \ -v $(pwd)/config.yaml:/app/config.yaml \ - ghcr.io/berriai/litellm:v1.81.3-stable.sonnet-4-6 \ + ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3 \ --config /app/config.yaml ``` From 74abf0c8e6b6144032dc386972626b2b6ee22a98 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 25 Feb 2026 01:19:10 +0530 Subject: [PATCH 15/29] Fix phase docs link --- docs/my-website/blog/gpt_5_3_codex/index.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/my-website/blog/gpt_5_3_codex/index.md b/docs/my-website/blog/gpt_5_3_codex/index.md index 5c43d287bce..5b7773cf8aa 100644 --- a/docs/my-website/blog/gpt_5_3_codex/index.md +++ b/docs/my-website/blog/gpt_5_3_codex/index.md @@ -29,7 +29,7 @@ LiteLLM now supports GPT-5.3-Codex on Day 0, including support for the new assis `phase` appears on assistant output items and helps distinguish preamble/commentary turns from final closeout responses. -Reference: [Phase parameter docs](https://developers-site-git-alphas-venusaur-api-openai.vercel.app/alphas/venusaur-api/phase-parameter) +Reference: [Phase parameter docs](https://developers.openai.com/api/reference/overview) Supported values: - `null` From 5d291c739fbc493e3d2dae593add5ba6eb221c0a Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 25 Feb 2026 01:21:38 +0530 Subject: [PATCH 16/29] Fix phase docs link --- docs/my-website/blog/gpt_5_3_codex/index.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/my-website/blog/gpt_5_3_codex/index.md b/docs/my-website/blog/gpt_5_3_codex/index.md index 5b7773cf8aa..850586538f6 100644 --- a/docs/my-website/blog/gpt_5_3_codex/index.md +++ b/docs/my-website/blog/gpt_5_3_codex/index.md @@ -66,7 +66,7 @@ model_list: ```bash docker run -d \ -p 4000:4000 \ - -e ANTHROPIC_API_KEY=$ANTHROPIC_API_KEY \ + -e ANTHROPIC_API_KEY=$OPENAI_API_KEY \ -v $(pwd)/config.yaml:/app/config.yaml \ ghcr.io/berriai/litellm:v1.81.12-stable.gpt-5.3 \ --config /app/config.yaml From e44b9b6b3584710a948365268d64676e2177adab Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 24 Feb 2026 11:51:42 -0800 Subject: [PATCH 17/29] feat(prometheus): add opt-in stream label to litellm_proxy_total_requests_metric (#22023) Set prometheus_emit_stream_label: true in litellm_settings to emit a stream label (True/False/None) on litellm_proxy_total_requests_metric. Opt-in to avoid breaking cardinality on existing deployments. --- docs/my-website/docs/proxy/prometheus.md | 26 +++++- litellm/__init__.py | 1 + litellm/integrations/prometheus.py | 6 ++ litellm/types/integrations/prometheus.py | 25 ++++-- .../test_prometheus_stream_label.py | 81 +++++++++++++++++++ 5 files changed, 130 insertions(+), 9 deletions(-) create mode 100644 tests/test_litellm/integrations/test_prometheus_stream_label.py diff --git a/docs/my-website/docs/proxy/prometheus.md b/docs/my-website/docs/proxy/prometheus.md index 93a0675f097..18a139d1d29 100644 --- a/docs/my-website/docs/proxy/prometheus.md +++ b/docs/my-website/docs/proxy/prometheus.md @@ -122,7 +122,7 @@ Use this to track overall LiteLLM Proxy usage. | Metric Name | Description | |----------------------|--------------------------------------| | `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "user_email", "exception_status", "exception_class", "route", "model_id"` | -| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"` | +| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route", "model_id"`. Optionally includes `"stream"` — see [Emit Stream Label](#emit-stream-label). | ### Callback Logging Metrics @@ -214,9 +214,31 @@ litellm_settings: ``` +### Emit Stream Label + +Add a `stream` label to `litellm_proxy_total_requests_metric` to split requests by streaming vs. non-streaming. Disabled by default. + +```yaml title="config.yaml" +litellm_settings: + callbacks: ["prometheus"] + prometheus_emit_stream_label: true +``` + +When enabled, `litellm_proxy_total_requests_metric` gains a `stream` label with values `"True"`, `"False"`, or `"None"`. + +``` +litellm_proxy_total_requests_metric{..., stream="True"} 42 +litellm_proxy_total_requests_metric{..., stream="False"} 100 +``` + +:::note +This label is opt-in because adding a new label to an existing metric changes its cardinality and breaks existing Prometheus queries / Grafana dashboards that target this metric. Enable it only on fresh deployments or when you are ready to update your dashboards. +::: + + ## [BETA] Custom Metrics -Track custom metrics on prometheus on all events mentioned above. +Track custom metrics on prometheus on all events mentioned above. ### Custom Metadata Labels diff --git a/litellm/__init__.py b/litellm/__init__.py index 1e74b5692e4..6e42f2c1ea5 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -374,6 +374,7 @@ enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None custom_prometheus_metadata_labels: List[str] = [] custom_prometheus_tags: List[str] = [] prometheus_metrics_config: Optional[List] = None +prometheus_emit_stream_label: bool = False disable_add_prefix_to_prompt: bool = ( False # used by anthropic, to disable adding prefix to prompt ) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 4c7afd5a57c..08db77e8571 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -974,6 +974,9 @@ class PrometheusLogger(CustomLogger): ), client_ip=standard_logging_payload["metadata"].get("requester_ip_address"), user_agent=standard_logging_payload["metadata"].get("user_agent"), + stream=str(standard_logging_payload.get("stream")) + if litellm.prometheus_emit_stream_label + else None, ) if ( @@ -1624,6 +1627,9 @@ class PrometheusLogger(CustomLogger): client_ip=_metadata.get("requester_ip_address"), user_agent=_metadata.get("user_agent"), model_id=model_id, + stream=str(request_data.get("stream")) + if litellm.prometheus_emit_stream_label + else None, ) _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index fd788af9ac1..482b87085dd 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -55,22 +55,21 @@ def _sanitize_prometheus_label_value(value: Optional[Any]) -> Optional[str]: return None # Coerce non-string values (int, bool, etc.) to str before sanitizing - if not isinstance(value, str): - value = str(value) + str_value: str = value if isinstance(value, str) else str(value) # Remove Unicode line/paragraph separators that break text format - value = value.replace("\u2028", "").replace("\u2029", "") + str_value = str_value.replace("\u2028", "").replace("\u2029", "") # Remove carriage returns - value = value.replace("\r", "") + str_value = str_value.replace("\r", "") # Replace newlines with spaces - value = value.replace("\n", " ") + str_value = str_value.replace("\n", " ") # Escape backslashes and double quotes per Prometheus exposition format - value = value.replace("\\", "\\\\").replace('"', '\\"') + str_value = str_value.replace("\\", "\\\\").replace('"', '\\"') - return value + return str_value @dataclass @@ -185,6 +184,7 @@ class UserAPIKeyLabelNames(Enum): CLIENT_IP = "client_ip" USER_AGENT = "user_agent" CALLBACK_NAME = "callback_name" + STREAM = "stream" DEFINED_PROMETHEUS_METRICS = Literal[ @@ -638,6 +638,14 @@ class PrometheusMetricLabels: ] ) + # Conditionally add stream label to litellm_proxy_total_requests_metric + if ( + label_name == "litellm_proxy_total_requests_metric" + and litellm.prometheus_emit_stream_label is True + and UserAPIKeyLabelNames.STREAM.value not in default_labels + ): + custom_labels.append(UserAPIKeyLabelNames.STREAM.value) + return default_labels + custom_labels @@ -709,6 +717,9 @@ class UserAPIKeyLabelValues(BaseModel): user_agent: Annotated[ Optional[str], Field(..., alias=UserAPIKeyLabelNames.USER_AGENT.value) ] = None + stream: Annotated[ + Optional[str], Field(..., alias=UserAPIKeyLabelNames.STREAM.value) + ] = None class PrometheusMetricsConfig(BaseModel): diff --git a/tests/test_litellm/integrations/test_prometheus_stream_label.py b/tests/test_litellm/integrations/test_prometheus_stream_label.py new file mode 100644 index 00000000000..a00a468e0fb --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_stream_label.py @@ -0,0 +1,81 @@ +""" +Unit tests for prometheus_emit_stream_label opt-in setting. + +Tests that: +- stream label is NOT added to litellm_proxy_total_requests_metric by default +- stream label IS added when litellm.prometheus_emit_stream_label = True +- stream value is populated correctly from standard_logging_payload +""" +import pytest + +import litellm +from litellm.types.integrations.prometheus import ( + PrometheusMetricLabels, + UserAPIKeyLabelNames, +) + + +def test_stream_label_not_present_by_default(): + """stream label should NOT appear in litellm_proxy_total_requests_metric unless opted in""" + litellm.prometheus_emit_stream_label = False + labels = PrometheusMetricLabels.get_labels("litellm_proxy_total_requests_metric") + assert UserAPIKeyLabelNames.STREAM.value not in labels + + +def test_stream_label_present_when_opted_in(): + """stream label SHOULD appear in litellm_proxy_total_requests_metric when opted in""" + litellm.prometheus_emit_stream_label = True + try: + labels = PrometheusMetricLabels.get_labels("litellm_proxy_total_requests_metric") + assert UserAPIKeyLabelNames.STREAM.value in labels + finally: + litellm.prometheus_emit_stream_label = False + + +def test_stream_label_not_in_other_metrics_when_opted_in(): + """stream label should NOT be added to other metrics even when opted in""" + litellm.prometheus_emit_stream_label = True + try: + other_metrics = [ + "litellm_proxy_failed_requests_metric", + "litellm_spend_metric", + "litellm_input_tokens_metric", + "litellm_output_tokens_metric", + "litellm_llm_api_latency_metric", + ] + for metric in other_metrics: + labels = PrometheusMetricLabels.get_labels(metric) + assert UserAPIKeyLabelNames.STREAM.value not in labels, ( + f"stream label should not be in {metric}" + ) + finally: + litellm.prometheus_emit_stream_label = False + + +def test_stream_label_name(): + """STREAM label name should be 'stream'""" + assert UserAPIKeyLabelNames.STREAM.value == "stream" + + +def test_user_api_key_label_values_has_stream_field(): + """UserAPIKeyLabelValues should accept stream field""" + from litellm.types.integrations.prometheus import UserAPIKeyLabelValues + + values = UserAPIKeyLabelValues(stream="True") + assert values.stream == "True" + + values_false = UserAPIKeyLabelValues(stream="False") + assert values_false.stream == "False" + + values_none = UserAPIKeyLabelValues() + assert values_none.stream is None + + +def test_stream_label_in_model_dump(): + """stream field appears in model_dump() output for use in prometheus_label_factory""" + from litellm.types.integrations.prometheus import UserAPIKeyLabelValues + + values = UserAPIKeyLabelValues(stream="True") + dumped = values.model_dump() + assert "stream" in dumped + assert dumped["stream"] == "True" From c343bfffdaa3e34e9324733d897a8e4db79ac9ec Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 24 Feb 2026 11:56:16 -0800 Subject: [PATCH 18/29] fix(router): emit x-litellm-overhead-duration-ms header for streaming requests (#22027) * fix(router): preserve _hidden_params in FallbackStreamWrapper so x-litellm-overhead-duration-ms is emitted for streaming requests * test(router): add regression test for FallbackStreamWrapper _hidden_params preservation --- litellm/router.py | 26 ++- .../proxy/test_common_request_processing.py | 201 ++++++++++++++++++ tests/test_litellm/test_router.py | 57 ++++- 3 files changed, 274 insertions(+), 10 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index ac2862da689..46d35352c37 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -163,7 +163,11 @@ from litellm.types.utils import ( ) from litellm.types.utils import ModelInfo from litellm.types.utils import ModelInfo as ModelMapInfo -from litellm.types.utils import ModelResponseStream, StandardLoggingPayload, Usage +from litellm.types.utils import ( + ModelResponseStream, + StandardLoggingPayload, + Usage, +) from litellm.utils import ( CustomStreamWrapper, EmbeddingResponse, @@ -1555,6 +1559,9 @@ class Router: logging_obj=model_response.logging_obj, ) self._async_generator = async_generator + # Preserve hidden params (including litellm_overhead_time_ms) from original response + if hasattr(model_response, "_hidden_params"): + self._hidden_params = model_response._hidden_params.copy() def __aiter__(self): return self @@ -6978,9 +6985,9 @@ class Router: raise ValueError("Deployment not found") ## GET BASE MODEL - base_model = deployment.get("model_info", {}).get("base_model", None) + base_model = (deployment.get("model_info") or {}).get("base_model", None) if base_model is None: - base_model = deployment.get("litellm_params", {}).get("base_model", None) + base_model = (deployment.get("litellm_params") or {}).get("base_model", None) model = base_model @@ -6995,7 +7002,7 @@ class Router: raise ValueError( f"Deployment missing valid litellm_params. " f"Got: {type(litellm_params_data).__name__}, " - f"deployment_id: {deployment.get('model_info', {}).get('id', 'unknown')}" + f"deployment_id: {(deployment.get('model_info') or {}).get('id', 'unknown')}" ) _model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=litellm_params.model, @@ -7015,10 +7022,10 @@ class Router: if potential_models is not None: for potential_model in potential_models: try: - if potential_model.get("model_info", {}).get( + if (potential_model.get("model_info") or {}).get( "id" - ) == deployment.get("model_info", {}).get("id"): - model = potential_model.get("litellm_params", {}).get( + ) == (deployment.get("model_info") or {}).get("id"): + model = (potential_model.get("litellm_params") or {}).get( "model" ) break @@ -7039,9 +7046,10 @@ class Router: model_info = litellm.get_model_info(model=model_info_name) ## CHECK USER SET MODEL INFO - user_model_info = deployment.get("model_info", {}) + user_model_info = deployment.get("model_info") or {} - model_info.update(user_model_info) + if model_info is not None: + model_info.update(user_model_info) return model_info diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 977304f732b..bf794478f10 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -1,4 +1,6 @@ import copy +import datetime +from typing import AsyncGenerator from unittest.mock import AsyncMock, MagicMock import pytest @@ -1348,3 +1350,202 @@ class TestOverrideOpenAIResponseModel: # Verify the model was not changed assert response_obj.model == fallback_model + + +class TestStreamingOverheadHeader: + """ + Tests that x-litellm-overhead-duration-ms is emitted in streaming responses. + + Regression tests for: streaming requests not including overhead header. + """ + + def test_get_custom_headers_includes_overhead_when_set(self): + """ + get_custom_headers() returns x-litellm-overhead-duration-ms + when litellm_overhead_time_ms is in hidden_params. + """ + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.tpm_limit = None + mock_user_api_key_dict.rpm_limit = None + mock_user_api_key_dict.max_budget = None + mock_user_api_key_dict.spend = 0.0 + mock_user_api_key_dict.allowed_model_region = None + + hidden_params = { + "litellm_overhead_time_ms": 42.5, + "_response_ms": 500.0, + "model_id": "test-model-id", + "api_base": "https://api.openai.com", + } + + headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + call_id="test-call-id", + model_id="test-model-id", + cache_key="", + api_base="https://api.openai.com", + version="1.0.0", + response_cost=0.001, + model_region="", + hidden_params=hidden_params, + ) + + assert "x-litellm-overhead-duration-ms" in headers + assert headers["x-litellm-overhead-duration-ms"] == "42.5" + + def test_get_custom_headers_omits_overhead_when_none(self): + """ + get_custom_headers() omits x-litellm-overhead-duration-ms + when litellm_overhead_time_ms is not in hidden_params. + """ + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.tpm_limit = None + mock_user_api_key_dict.rpm_limit = None + mock_user_api_key_dict.max_budget = None + mock_user_api_key_dict.spend = 0.0 + mock_user_api_key_dict.allowed_model_region = None + + hidden_params = { + "_response_ms": 500.0, + "model_id": "test-model-id", + } + + headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + call_id="test-call-id", + model_id="test-model-id", + cache_key="", + api_base="https://api.openai.com", + version="1.0.0", + response_cost=0.001, + model_region="", + hidden_params=hidden_params, + ) + + # Should be absent (None gets filtered by exclude_values) + assert "x-litellm-overhead-duration-ms" not in headers + + def test_update_response_metadata_sets_overhead_on_stream_wrapper(self): + """ + update_response_metadata() sets litellm_overhead_time_ms on + a streaming response's _hidden_params when llm_api_duration_ms is available. + """ + from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( + update_response_metadata, + ) + + # Mock the logging object with llm_api_duration_ms set + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = { + "llm_api_duration_ms": 200.0, + "litellm_params": {}, + } + mock_logging_obj.caching_details = None + mock_logging_obj.callback_duration_ms = None + mock_logging_obj.litellm_call_id = "test-call-id" + mock_logging_obj._response_cost_calculator = MagicMock(return_value=0.001) + + # Simulate a streaming result object with _hidden_params (like CustomStreamWrapper) + stream_result = MagicMock() + stream_result._hidden_params = { + "model_id": "test-model-id", + "api_base": "https://api.openai.com", + "additional_headers": {}, + } + + start_time = datetime.datetime.now() - datetime.timedelta(milliseconds=300) + end_time = datetime.datetime.now() + + update_response_metadata( + result=stream_result, + logging_obj=mock_logging_obj, + model="gpt-4o", + kwargs={}, + start_time=start_time, + end_time=end_time, + ) + + assert "litellm_overhead_time_ms" in stream_result._hidden_params + overhead = stream_result._hidden_params["litellm_overhead_time_ms"] + assert overhead is not None + assert isinstance(overhead, float) + # overhead = total_response_ms (~300ms) - llm_api_duration_ms (200ms) = ~100ms + assert overhead > 0 + + @pytest.mark.asyncio + async def test_streaming_response_includes_overhead_header(self): + """ + StreamingResponse returned by create_response() includes + x-litellm-overhead-duration-ms in its headers. + """ + + async def mock_generator() -> AsyncGenerator[str, None]: + yield 'data: {"id":"chatcmpl-test","choices":[{"delta":{"content":"hi"}}]}\n\n' + yield "data: [DONE]\n\n" + + headers = { + "x-litellm-overhead-duration-ms": "42.5", + "x-litellm-call-id": "test-call-id", + "x-litellm-model-id": "test-model-id", + } + + response = await create_response( + generator=mock_generator(), + media_type="text/event-stream", + headers=headers, + ) + + assert isinstance(response, StreamingResponse) + assert response.headers.get("x-litellm-overhead-duration-ms") == "42.5" + + def test_streaming_overhead_header_in_custom_headers_from_stream_hidden_params( + self, + ): + """ + Verifies that when get_custom_headers() is called with a streaming + response's hidden_params (containing litellm_overhead_time_ms), + the x-litellm-overhead-duration-ms header is correctly populated. + + This tests the critical path: update_response_metadata sets the value + → get_custom_headers reads it → StreamingResponse header is set. + """ + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.tpm_limit = None + mock_user_api_key_dict.rpm_limit = None + mock_user_api_key_dict.max_budget = None + mock_user_api_key_dict.spend = 0.0 + mock_user_api_key_dict.allowed_model_region = None + + # This is what CustomStreamWrapper._hidden_params looks like after + # update_response_metadata() has been called on it + hidden_params = { + "model_id": "openai-gpt4o-deployment", + "api_base": "https://api.openai.com", + "additional_headers": {}, + "litellm_overhead_time_ms": 55.3, # set by update_response_metadata + "_response_ms": 280.0, + "litellm_call_id": "test-call-id", + "response_cost": 0.002, + "cache_key": None, + "fastest_response_batch_completion": None, + "callback_duration_ms": None, + } + + custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + call_id="test-call-id", + model_id=hidden_params.get("model_id"), + cache_key=hidden_params.get("cache_key") or "", + api_base=hidden_params.get("api_base") or "", + version="1.0.0", + response_cost=hidden_params.get("response_cost"), + model_region="", + hidden_params=hidden_params, + ) + + # The overhead header must be present and correct + assert "x-litellm-overhead-duration-ms" in custom_headers, ( + "x-litellm-overhead-duration-ms header must be emitted during streaming. " + "It was missing — this is the streaming overhead header regression." + ) + assert custom_headers["x-litellm-overhead-duration-ms"] == "55.3" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 5542cdf8be4..5732deda6fb 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1297,6 +1297,61 @@ async def test_acompletion_streaming_iterator_edge_cases(): print("✓ Edge case tests passed!") +@pytest.mark.asyncio +async def test_acompletion_streaming_iterator_preserves_hidden_params(): + """ + Regression test: FallbackStreamWrapper must copy _hidden_params from the + original CustomStreamWrapper so that x-litellm-overhead-duration-ms (and + other hidden params) are present in the proxy response headers for streaming. + """ + from unittest.mock import MagicMock + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake-key"}, + } + ], + ) + + # Simulate a CustomStreamWrapper that already has timing metadata set by + # update_response_metadata (litellm_overhead_time_ms, _response_ms, etc.) + mock_response = MagicMock() + mock_response.model = "gpt-4" + mock_response.custom_llm_provider = "openai" + mock_response.logging_obj = MagicMock() + mock_response._hidden_params = { + "litellm_overhead_time_ms": 12.34, + "_response_ms": 500.0, + "litellm_call_id": "test-call-id", + "api_base": "https://api.openai.com", + "additional_headers": {}, + } + + # Make the mock iterable (yields nothing — we only care about hidden_params) + async def _empty(): + return + yield # make it an async generator + + mock_response.__aiter__ = lambda self: _empty().__aiter__() + + result = await router._acompletion_streaming_iterator( + model_response=mock_response, + messages=[{"role": "user", "content": "hi"}], + initial_kwargs={"model": "gpt-4", "stream": True}, + ) + + # The returned FallbackStreamWrapper must carry the original _hidden_params + assert hasattr(result, "_hidden_params"), "result must have _hidden_params" + assert result._hidden_params.get("litellm_overhead_time_ms") == 12.34, ( + "litellm_overhead_time_ms must be preserved — " + "this is what drives x-litellm-overhead-duration-ms in streaming responses" + ) + assert result._hidden_params.get("litellm_call_id") == "test-call-id" + assert result._hidden_params.get("_response_ms") == 500.0 + + @pytest.mark.asyncio async def test_async_function_with_fallbacks_common_utils(): """Test the async_function_with_fallbacks_common_utils method""" @@ -1858,7 +1913,7 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name(): litellm_credential_name to actual credential values (for UI-created models). """ from litellm.types.utils import CredentialItem - + # Setup credential list with a test credential litellm.credential_list = [ CredentialItem( From 235d60eb885210aca0d981a63ccfe708caa8ed81 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 11:58:59 -0800 Subject: [PATCH 19/29] address greptile review feedback (greploop iteration 1) - Add traceback to cache update warning logs (user, end_user, team, tag) - Remove duplicate info log in non-redis commit path --- litellm/proxy/db/db_spend_update_writer.py | 16 ---------------- litellm/proxy/proxy_server.py | 12 ++++++++---- 2 files changed, 8 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index c87be7ac819..e2d365ee465 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -659,22 +659,6 @@ class DBSpendUpdateWriter: db_spend_update_transactions = ( await self.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() ) - if any( - len(v) > 0 - for v in db_spend_update_transactions.values() - if isinstance(v, dict) - ): - verbose_proxy_logger.info( - "Spend tracking - committing spend updates to DB (no Redis buffer): " - "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d", - len(db_spend_update_transactions.get("key_list_transactions") or {}), - len(db_spend_update_transactions.get("user_list_transactions") or {}), - len(db_spend_update_transactions.get("team_list_transactions") or {}), - len(db_spend_update_transactions.get("org_list_transactions") or {}), - len(db_spend_update_transactions.get("end_user_list_transactions") or {}), - len(db_spend_update_transactions.get("team_member_list_transactions") or {}), - len(db_spend_update_transactions.get("tag_list_transactions") or {}), - ) await self._commit_spend_updates_to_db( prisma_client=prisma_client, n_retry_times=n_retry_times, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 78730ee9d60..1983a601e13 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1771,10 +1771,11 @@ async def update_cache( # noqa: PLR0915 verbose_proxy_logger.warning( "Spend tracking - failed to update user spend in cache. " "Budget enforcement may use stale spend values. " - "user_id=%s, response_cost=%s - %s", + "user_id=%s, response_cost=%s - %s\n%s", user_id, response_cost, str(e), + traceback.format_exc(), ) ### UPDATE END-USER SPEND ### @@ -1814,10 +1815,11 @@ async def update_cache( # noqa: PLR0915 verbose_proxy_logger.warning( "Spend tracking - failed to update end user spend in cache. " "Budget enforcement may use stale spend values. " - "end_user_id=%s, response_cost=%s - %s", + "end_user_id=%s, response_cost=%s - %s\n%s", end_user_id, response_cost, str(e), + traceback.format_exc(), ) ### UPDATE TEAM SPEND ### @@ -1861,10 +1863,11 @@ async def update_cache( # noqa: PLR0915 verbose_proxy_logger.warning( "Spend tracking - failed to update team spend in cache. " "Budget enforcement may use stale spend values. " - "team_id=%s, response_cost=%s - %s", + "team_id=%s, response_cost=%s - %s\n%s", team_id, response_cost, str(e), + traceback.format_exc(), ) ### UPDATE TAG SPEND ### @@ -1912,10 +1915,11 @@ async def update_cache( # noqa: PLR0915 verbose_proxy_logger.warning( "Spend tracking - failed to update tag spend in cache. " "Budget enforcement may use stale spend values. " - "tags=%s, response_cost=%s - %s", + "tags=%s, response_cost=%s - %s\n%s", tags, response_cost, str(e), + traceback.format_exc(), ) if token is not None and response_cost is not None: From 70ef4d0d69e3915df64c847582e5b0d60ab1fae3 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 12:10:19 -0800 Subject: [PATCH 20/29] address greptile review feedback (greploop iteration 2) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Remove re-raise in _store_transactions_in_redis so one Redis push failure doesn't drop remaining transaction types - Downgrade per-push success log from info to debug to reduce noise - Fix misleading error message in update_database — entity spend updates run as independent tasks and are not affected by this catch --- litellm/proxy/db/db_spend_update_writer.py | 3 ++- litellm/proxy/db/db_transaction_queue/redis_update_buffer.py | 3 +-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e2d365ee465..4a8c9d33d9f 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -225,7 +225,8 @@ class DBSpendUpdateWriter: verbose_proxy_logger.debug("Runs spend update on all tables") except Exception: verbose_proxy_logger.error( - "Spend tracking - update_database failed. All spend updates for this request will be lost. " + "Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue " + "may not have completed for this request. " "response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s - %s", response_cost, token, diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 85e139b4de0..027f3e639e3 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -101,7 +101,7 @@ class RedisUpdateBuffer: key=redis_key, values=list_of_transactions, ) - verbose_proxy_logger.info( + verbose_proxy_logger.debug( "Spend tracking - pushed spend updates to Redis buffer. " "redis_key=%s, buffer_size=%s", redis_key, @@ -118,7 +118,6 @@ class RedisUpdateBuffer: redis_key, str(e), ) - raise async def store_in_memory_spend_updates_in_redis( self, From 33719e6b38d830c5f63bf6866fd88bb996eebcca Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 24 Feb 2026 12:30:18 -0800 Subject: [PATCH 21/29] docs: update v1.81.12-stable release notes to point to v1.81.12-stable.1 (#22036) --- docs/my-website/release_notes/v1.81.12.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/my-website/release_notes/v1.81.12.md b/docs/my-website/release_notes/v1.81.12.md index a1f1daa2b92..0b7c1e146ab 100644 --- a/docs/my-website/release_notes/v1.81.12.md +++ b/docs/my-website/release_notes/v1.81.12.md @@ -1,5 +1,5 @@ --- -title: "v1.81.12-stable - Guardrail Policy Templates & Action Builder" +title: "v1.81.12-stable.1 - Guardrail Policy Templates & Action Builder" slug: "v1-81-12" date: 2026-02-14T00:00:00 authors: @@ -27,7 +27,7 @@ import Image from '@theme/IdealImage'; docker run \ -e STORE_MODEL_IN_DB=True \ -p 4000:4000 \ -ghcr.io/berriai/litellm:main-v1.81.12-stable +ghcr.io/berriai/litellm:main-v1.81.12-stable.1 ``` From 8b56e1d969385b066f596400e084e754e438fbe5 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 12:37:43 -0800 Subject: [PATCH 22/29] trigger review From 2cabbccf6f8244348dc908e2a2d0a069b3313001 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 12:49:48 -0800 Subject: [PATCH 23/29] address greptile review feedback (greploop iteration 3) - Add missing traceback to team member spend enqueue error log --- litellm/proxy/db/db_spend_update_writer.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 4a8c9d33d9f..d7d6b2b1eb0 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -351,11 +351,12 @@ class DBSpendUpdateWriter: except Exception as e: verbose_proxy_logger.error( "Spend tracking - failed to enqueue team member spend update. " - "team_id=%s, user_id=%s, response_cost=%s - %s", + "team_id=%s, user_id=%s, response_cost=%s - %s\n%s", team_id, user_id, response_cost, str(e), + traceback.format_exc(), ) except Exception as e: verbose_proxy_logger.error( From 9971b67587f8c376923995ee1260429195458124 Mon Sep 17 00:00:00 2001 From: Ryan Crabbe Date: Tue, 24 Feb 2026 15:21:43 -0800 Subject: [PATCH 24/29] fix: convert remaining dict(request.headers) to _safe_get_request_headers Missed conversions in user_api_key_auth.py and litellm_pre_call_utils.py. Both call sites are read-only so no .copy() needed. --- litellm/proxy/auth/user_api_key_auth.py | 2 +- litellm/proxy/litellm_pre_call_utils.py | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ca9af16ff91..8f17440773a 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -483,7 +483,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 parent_otel_span = ( open_telemetry_logger.create_litellm_proxy_request_started_span( start_time=start_time, - headers=dict(request.headers), + headers=_safe_get_request_headers(request), ) ) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index b61dfa5b263..52f0b1d46e9 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -14,6 +14,7 @@ from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors, LitellmDataForBackendLLMCall, LitellmUserRoles, SpecialHeaders, TeamCallbackMetadata, UserAPIKeyAuth) +from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers # Cache special headers as a frozenset for O(1) lookup performance _SPECIAL_HEADERS_CACHE = frozenset( @@ -824,7 +825,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915 from litellm.proxy.proxy_server import llm_router, premium_user from litellm.types.proxy.litellm_pre_call_utils import SecretFields - _raw_headers: Dict[str, str] = dict(request.headers) + _raw_headers: Dict[str, str] = _safe_get_request_headers(request) _headers: Dict[str, str] = clean_headers( request.headers, litellm_key_header_name=( From b3bb744aa4d0fa026e93f89ad6a11ce75b86dbdb Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 15:26:06 -0800 Subject: [PATCH 25/29] [Test] Add unit tests for router_settings components Co-Authored-By: Claude Sonnet 4.6 --- .../LatencyBasedConfiguration.test.tsx | 55 ++++++ .../ReliabilityRetriesSection.test.tsx | 80 +++++++++ .../RouterSettingsForm.test.tsx | 134 ++++++++++++++ .../RoutingStrategySelector.test.tsx | 98 ++++++++++ .../TagFilteringToggle.test.tsx | 113 ++++++++++++ .../components/router_settings/index.test.tsx | 170 ++++++++++++++++++ 6 files changed, 650 insertions(+) create mode 100644 ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx create mode 100644 ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx create mode 100644 ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx create mode 100644 ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx create mode 100644 ui/litellm-dashboard/src/components/router_settings/TagFilteringToggle.test.tsx create mode 100644 ui/litellm-dashboard/src/components/router_settings/index.test.tsx diff --git a/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx b/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx new file mode 100644 index 00000000000..0176be5ab40 --- /dev/null +++ b/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx @@ -0,0 +1,55 @@ +import { describe, it, expect } from "vitest"; +import { render, screen } from "@testing-library/react"; +import LatencyBasedConfiguration from "./LatencyBasedConfiguration"; + +describe("LatencyBasedConfiguration", () => { + it("should render the section heading", () => { + render(); + expect(screen.getByText("Latency-Based Configuration")).toBeInTheDocument(); + }); + + it("should render default params when no args are provided", () => { + render(); + // Default: ttl=3600, lowest_latency_buffer=0 + expect(screen.getByDisplayValue("3600")).toBeInTheDocument(); + expect(screen.getByDisplayValue("0")).toBeInTheDocument(); + }); + + it("should render the provided routing strategy args as inputs", () => { + const args = { ttl: 7200, lowest_latency_buffer: 0.1 }; + render(); + expect(screen.getByDisplayValue("7200")).toBeInTheDocument(); + expect(screen.getByDisplayValue("0.1")).toBeInTheDocument(); + }); + + it("should render an input with the correct name attribute for each param", () => { + const args = { ttl: 3600, lowest_latency_buffer: 0 }; + render(); + expect(screen.getByRole("textbox", { name: /ttl/i })).toBeInTheDocument(); + expect(screen.getByRole("textbox", { name: /lowest latency buffer/i })).toBeInTheDocument(); + }); + + it("should display the TTL parameter explanation", () => { + render(); + expect( + screen.getByText(/sliding window to look back over/i) + ).toBeInTheDocument(); + }); + + it("should display the lowest_latency_buffer parameter explanation", () => { + render(); + expect( + screen.getByText(/shuffle between deployments within this %/i) + ).toBeInTheDocument(); + }); + + it("should render object values stringified into the input", () => { + const args = { ttl: { nested: true } }; + render(); + // HTML input type=text strips newlines, so check that the key/value appears + const input = document.querySelector('input[name="ttl"]') as HTMLInputElement; + expect(input).not.toBeNull(); + expect(input.value).toContain('"nested"'); + expect(input.value).toContain('true'); + }); +}); diff --git a/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx b/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx new file mode 100644 index 00000000000..0892d39fc29 --- /dev/null +++ b/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx @@ -0,0 +1,80 @@ +import { describe, it, expect } from "vitest"; +import { render, screen } from "@testing-library/react"; +import ReliabilityRetriesSection from "./ReliabilityRetriesSection"; + +const baseSettings = { + num_retries: 3, + timeout: 30, + allowed_fails: 2, + fallbacks: ["gpt-3.5"], + context_window_fallbacks: [], + routing_strategy_args: { ttl: 3600 }, + routing_strategy: "simple-shuffle", + enable_tag_filtering: false, +}; + +describe("ReliabilityRetriesSection", () => { + it("should render the section heading", () => { + render(); + expect(screen.getByText("Reliability & Retries")).toBeInTheDocument(); + }); + + it("should render input fields for non-excluded settings", () => { + render(); + expect(screen.getByDisplayValue("3")).toBeInTheDocument(); // num_retries + expect(screen.getByDisplayValue("30")).toBeInTheDocument(); // timeout + expect(screen.getByDisplayValue("2")).toBeInTheDocument(); // allowed_fails + }); + + it("should not render inputs for excluded keys", () => { + render(); + // Each excluded key must not produce a visible input value + const inputs = screen.queryAllByRole("textbox"); + const inputNames = inputs.map((el) => el.getAttribute("name")); + expect(inputNames).not.toContain("fallbacks"); + expect(inputNames).not.toContain("context_window_fallbacks"); + expect(inputNames).not.toContain("routing_strategy_args"); + expect(inputNames).not.toContain("routing_strategy"); + expect(inputNames).not.toContain("enable_tag_filtering"); + }); + + it("should use ui_field_name from metadata as the label", () => { + const metadata = { + num_retries: { ui_field_name: "Number of Retries", field_description: "How many times to retry" }, + }; + render( + + ); + expect(screen.getByText("Number of Retries")).toBeInTheDocument(); + }); + + it("should fall back to the raw param name when no metadata label is available", () => { + render( + + ); + expect(screen.getByText("num_retries")).toBeInTheDocument(); + }); + + it("should render null values as an empty input", () => { + render( + + ); + const input = screen.getByRole("textbox", { name: /timeout/i }) as HTMLInputElement; + expect(input.value).toBe(""); + }); + + it("should render object values stringified into the input", () => { + const settings = { retry_policy: { "rate-limited": 2 } }; + render(); + // HTML input type=text strips newlines, so check that the key/value appears + const input = document.querySelector('input[name="retry_policy"]') as HTMLInputElement; + expect(input).not.toBeNull(); + expect(input.value).toContain('"rate-limited"'); + expect(input.value).toContain('2'); + }); + + it("should render no inputs when routerSettings is empty", () => { + render(); + expect(screen.queryAllByRole("textbox")).toHaveLength(0); + }); +}); diff --git a/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx b/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx new file mode 100644 index 00000000000..767820cd485 --- /dev/null +++ b/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx @@ -0,0 +1,134 @@ +import { describe, it, expect, vi } from "vitest"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import RouterSettingsForm from "./RouterSettingsForm"; +import type { RouterSettingsFormValue } from "./RouterSettingsForm"; + +// Use the same antd mock as RoutingStrategySelector to keep things consistent +vi.mock("antd", () => ({ + Select: Object.assign( + ({ value, onChange, children }: any) => ( + + ), + { + Option: ({ value, children }: any) => ( + + ), + } + ), +})); + +vi.mock("@tremor/react", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + Switch: ({ checked, onChange }: any) => ( + onChange(e.target.checked)} + /> + ), + }; +}); + +const defaultValue: RouterSettingsFormValue = { + routerSettings: {}, + selectedStrategy: null, + enableTagFiltering: false, +}; + +const baseProps = { + value: defaultValue, + onChange: vi.fn(), + routerFieldsMetadata: {}, + availableRoutingStrategies: [], + routingStrategyDescriptions: {}, +}; + +describe("RouterSettingsForm", () => { + it("should render", () => { + render(); + expect(screen.getByText("Routing Settings")).toBeInTheDocument(); + }); + + it("should not show the strategy selector when no strategies are provided", () => { + render(); + expect(screen.queryByTestId("strategy-select")).not.toBeInTheDocument(); + }); + + it("should show the strategy selector when strategies are available", () => { + const props = { + ...baseProps, + availableRoutingStrategies: ["simple-shuffle", "latency-based-routing"], + }; + render(); + expect(screen.getByTestId("strategy-select")).toBeInTheDocument(); + }); + + it("should not render LatencyBasedConfiguration for non-latency strategies", () => { + const props = { + ...baseProps, + value: { ...defaultValue, selectedStrategy: "simple-shuffle" }, + availableRoutingStrategies: ["simple-shuffle"], + }; + render(); + expect(screen.queryByText("Latency-Based Configuration")).not.toBeInTheDocument(); + }); + + it("should render LatencyBasedConfiguration when strategy is latency-based-routing", () => { + const props = { + ...baseProps, + value: { + ...defaultValue, + selectedStrategy: "latency-based-routing", + routerSettings: { routing_strategy_args: { ttl: 3600, lowest_latency_buffer: 0 } }, + }, + availableRoutingStrategies: ["latency-based-routing"], + }; + render(); + expect(screen.getByText("Latency-Based Configuration")).toBeInTheDocument(); + }); + + it("should call onChange with the updated strategy when the selector changes", () => { + const onChange = vi.fn(); + const props = { + ...baseProps, + onChange, + availableRoutingStrategies: ["simple-shuffle", "latency-based-routing"], + }; + render(); + + const select = screen.getByTestId("strategy-select") as HTMLSelectElement; + select.value = "latency-based-routing"; + select.dispatchEvent(new Event("change", { bubbles: true })); + + expect(onChange).toHaveBeenCalledWith( + expect.objectContaining({ selectedStrategy: "latency-based-routing" }) + ); + }); + + it("should call onChange with the updated enableTagFiltering when the toggle changes", async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + render(); + + await user.click(screen.getByRole("switch")); + + expect(onChange).toHaveBeenCalledWith( + expect.objectContaining({ enableTagFiltering: true }) + ); + }); + + it("should show the Reliability & Retries section", () => { + render(); + expect(screen.getByText("Reliability & Retries")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx b/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx new file mode 100644 index 00000000000..85b1dc21acf --- /dev/null +++ b/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx @@ -0,0 +1,98 @@ +import { describe, it, expect, vi } from "vitest"; +import { render, screen } from "@testing-library/react"; +import RoutingStrategySelector from "./RoutingStrategySelector"; + +// Ant Design's Select is complex to drive in JSDOM; swap it for a plain +// onChange(e.target.value)} + > + {children} + + + ), + { + Option: ({ value, children }: any) => ( + + ), + } + ), +})); + +const baseProps = { + selectedStrategy: null, + availableStrategies: ["simple-shuffle", "latency-based-routing", "least-busy"], + routingStrategyDescriptions: { + "simple-shuffle": "Randomly pick a deployment", + "latency-based-routing": "Pick the lowest-latency deployment", + }, + routerFieldsMetadata: {}, + onStrategyChange: vi.fn(), +}; + +describe("RoutingStrategySelector", () => { + it("should render", () => { + render(); + expect(screen.getByTestId("ant-select")).toBeInTheDocument(); + }); + + it("should display default label when no metadata is provided", () => { + render(); + expect(screen.getByText("Routing Strategy")).toBeInTheDocument(); + }); + + it("should display ui_field_name from metadata when provided", () => { + const props = { + ...baseProps, + routerFieldsMetadata: { + routing_strategy: { + ui_field_name: "Strategy", + field_description: "How to pick a deployment", + }, + }, + }; + render(); + expect(screen.getByText("Strategy")).toBeInTheDocument(); + expect(screen.getByText("How to pick a deployment")).toBeInTheDocument(); + }); + + it("should render all available strategies as options", () => { + render(); + expect(screen.getByText("simple-shuffle")).toBeInTheDocument(); + expect(screen.getByText("latency-based-routing")).toBeInTheDocument(); + expect(screen.getByText("least-busy")).toBeInTheDocument(); + }); + + it("should display strategy descriptions alongside option labels", () => { + render(); + expect(screen.getByText("Randomly pick a deployment")).toBeInTheDocument(); + expect(screen.getByText("Pick the lowest-latency deployment")).toBeInTheDocument(); + }); + + it("should not render a description for a strategy that has none", () => { + render(); + // "least-busy" has no entry in routingStrategyDescriptions + const select = screen.getByTestId("strategy-select"); + const leastBusyOption = Array.from(select.querySelectorAll("option")).find( + (o) => o.value === "least-busy" + ); + expect(leastBusyOption).toBeInTheDocument(); + }); + + it("should call onStrategyChange with the selected strategy value", () => { + const onStrategyChange = vi.fn(); + render(); + + const select = screen.getByTestId("strategy-select") as HTMLSelectElement; + select.value = "latency-based-routing"; + select.dispatchEvent(new Event("change", { bubbles: true })); + + expect(onStrategyChange).toHaveBeenCalledWith("latency-based-routing"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/router_settings/TagFilteringToggle.test.tsx b/ui/litellm-dashboard/src/components/router_settings/TagFilteringToggle.test.tsx new file mode 100644 index 00000000000..593721db023 --- /dev/null +++ b/ui/litellm-dashboard/src/components/router_settings/TagFilteringToggle.test.tsx @@ -0,0 +1,113 @@ +import { describe, it, expect, vi } from "vitest"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import TagFilteringToggle from "./TagFilteringToggle"; + +// setupTests.ts mocks @tremor/react but leaves Switch as the real implementation. +// Re-mock Switch as a plain checkbox so toggle interactions are trivially testable. +vi.mock("@tremor/react", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + Switch: ({ checked, onChange, className }: any) => ( + onChange(e.target.checked)} + className={className} + /> + ), + }; +}); + +const baseMetadata = { + enable_tag_filtering: { + ui_field_name: "Tag Filtering", + field_description: "Route requests based on tags", + link: null, + }, +}; + +describe("TagFilteringToggle", () => { + it("should render", () => { + render( + + ); + expect(screen.getByRole("switch")).toBeInTheDocument(); + }); + + it("should display default label when no metadata is provided", () => { + render( + + ); + expect(screen.getByText("Enable Tag Filtering")).toBeInTheDocument(); + }); + + it("should display the label from metadata when provided", () => { + render( + + ); + expect(screen.getByText("Tag Filtering")).toBeInTheDocument(); + }); + + it("should display the description from metadata", () => { + render( + + ); + expect(screen.getByText("Route requests based on tags")).toBeInTheDocument(); + }); + + it("should render a Learn more link when metadata provides one", () => { + const metadata = { + enable_tag_filtering: { + ...baseMetadata.enable_tag_filtering, + link: "https://docs.example.com/tag-filtering", + }, + }; + render( + + ); + const link = screen.getByRole("link", { name: /learn more/i }); + expect(link).toBeInTheDocument(); + expect(link).toHaveAttribute("href", "https://docs.example.com/tag-filtering"); + }); + + it("should not render a Learn more link when metadata has no link", () => { + render( + + ); + expect(screen.queryByRole("link", { name: /learn more/i })).not.toBeInTheDocument(); + }); + + it("should reflect the enabled=true state on the switch", () => { + render( + + ); + expect(screen.getByRole("switch")).toBeChecked(); + }); + + it("should call onToggle with the new value when the switch is toggled", async () => { + const onToggle = vi.fn(); + const user = userEvent.setup(); + render( + + ); + + await user.click(screen.getByRole("switch")); + + expect(onToggle).toHaveBeenCalledWith(true); + }); +}); diff --git a/ui/litellm-dashboard/src/components/router_settings/index.test.tsx b/ui/litellm-dashboard/src/components/router_settings/index.test.tsx new file mode 100644 index 00000000000..1920268207d --- /dev/null +++ b/ui/litellm-dashboard/src/components/router_settings/index.test.tsx @@ -0,0 +1,170 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"; +import userEvent from "@testing-library/user-event"; +import RouterSettings from "./index"; + +vi.mock("antd", () => ({ + Select: Object.assign( + ({ value, onChange, children }: any) => ( + + ), + { + Option: ({ value, children }: any) => ( + + ), + } + ), +})); + +vi.mock("@tremor/react", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + Switch: ({ checked, onChange }: any) => ( + onChange(e.target.checked)} + /> + ), + }; +}); + +vi.mock("@/components/networking", () => ({ + getCallbacksCall: vi.fn(), + getRouterSettingsCall: vi.fn(), + setCallbacksCall: vi.fn(), +})); + +import { + getCallbacksCall, + getRouterSettingsCall, + setCallbacksCall, +} from "@/components/networking"; + +const mockCallbacksResponse = { + router_settings: { + routing_strategy: "simple-shuffle", + num_retries: 3, + timeout: 30, + }, +}; + +const mockRouterSettingsResponse = { + fields: [ + { + field_name: "routing_strategy", + ui_field_name: "Routing Strategy", + field_description: "How requests are distributed", + options: ["simple-shuffle", "latency-based-routing"], + link: null, + }, + { + field_name: "enable_tag_filtering", + ui_field_name: "Tag Filtering", + field_description: "Route by tag", + field_value: false, + link: null, + }, + ], + routing_strategy_descriptions: { + "simple-shuffle": "Randomly pick a deployment", + "latency-based-routing": "Pick the lowest-latency deployment", + }, +}; + +const defaultProps = { + accessToken: "test-token", + userRole: "Admin", + userID: "user-1", + modelData: null, +}; + +describe("RouterSettings", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(getCallbacksCall).mockResolvedValue(mockCallbacksResponse); + vi.mocked(getRouterSettingsCall).mockResolvedValue(mockRouterSettingsResponse); + vi.mocked(setCallbacksCall).mockResolvedValue({}); + }); + + it("should render nothing when accessToken is null", () => { + const { container } = renderWithProviders( + + ); + expect(container).toBeEmptyDOMElement(); + }); + + it("should render the Save Changes and Reset buttons when authenticated", () => { + renderWithProviders(); + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + expect(screen.getByRole("button", { name: /reset/i })).toBeInTheDocument(); + }); + + it("should fetch callbacks and router settings on mount", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(getCallbacksCall).toHaveBeenCalledWith("test-token", "user-1", "Admin"); + }); + expect(getRouterSettingsCall).toHaveBeenCalledWith("test-token"); + }); + + it("should not fetch data when any required prop is missing", () => { + renderWithProviders( + + ); + expect(getCallbacksCall).not.toHaveBeenCalled(); + }); + + it("should render routing strategies loaded from the API", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(screen.getByTestId("strategy-select")).toBeInTheDocument(); + }); + + const select = screen.getByTestId("strategy-select") as HTMLSelectElement; + const optionValues = Array.from(select.options).map((o) => o.value); + expect(optionValues).toContain("simple-shuffle"); + expect(optionValues).toContain("latency-based-routing"); + }); + + it("should call setCallbacksCall with updated settings on Save Changes", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + }); + + await user.click(screen.getByRole("button", { name: /save changes/i })); + + expect(setCallbacksCall).toHaveBeenCalledWith( + "test-token", + expect.objectContaining({ router_settings: expect.any(Object) }) + ); + }); + + it("should show a success notification after saving", async () => { + const NotificationsManager = await import("@/components/molecules/notifications_manager"); + const user = userEvent.setup(); + renderWithProviders(); + + await waitFor(() => + expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument() + ); + await user.click(screen.getByRole("button", { name: /save changes/i })); + + expect(NotificationsManager.default.success).toHaveBeenCalledWith( + "router settings updated successfully" + ); + }); +}); From 784af16cb4064fcd27f407a4b8e268045652ba62 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 15:48:31 -0800 Subject: [PATCH 26/29] address greptile review feedback (greploop iteration 1) - Replace document.querySelector/querySelectorAll with screen.getByRole - Replace raw dispatchEvent with userEvent.selectOptions --- .../LatencyBasedConfiguration.test.tsx | 3 +-- .../ReliabilityRetriesSection.test.tsx | 3 +-- .../router_settings/RouterSettingsForm.test.tsx | 7 +++---- .../RoutingStrategySelector.test.tsx | 16 ++++++---------- 4 files changed, 11 insertions(+), 18 deletions(-) diff --git a/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx b/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx index 0176be5ab40..39d9f8c881c 100644 --- a/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx +++ b/ui/litellm-dashboard/src/components/router_settings/LatencyBasedConfiguration.test.tsx @@ -47,8 +47,7 @@ describe("LatencyBasedConfiguration", () => { const args = { ttl: { nested: true } }; render(); // HTML input type=text strips newlines, so check that the key/value appears - const input = document.querySelector('input[name="ttl"]') as HTMLInputElement; - expect(input).not.toBeNull(); + const input = screen.getByRole("textbox", { name: /ttl/i }) as HTMLInputElement; expect(input.value).toContain('"nested"'); expect(input.value).toContain('true'); }); diff --git a/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx b/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx index 0892d39fc29..101b09af0dc 100644 --- a/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx +++ b/ui/litellm-dashboard/src/components/router_settings/ReliabilityRetriesSection.test.tsx @@ -67,8 +67,7 @@ describe("ReliabilityRetriesSection", () => { const settings = { retry_policy: { "rate-limited": 2 } }; render(); // HTML input type=text strips newlines, so check that the key/value appears - const input = document.querySelector('input[name="retry_policy"]') as HTMLInputElement; - expect(input).not.toBeNull(); + const input = screen.getByRole("textbox", { name: /retry_policy/i }) as HTMLInputElement; expect(input.value).toContain('"rate-limited"'); expect(input.value).toContain('2'); }); diff --git a/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx b/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx index 767820cd485..01f318b909a 100644 --- a/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx +++ b/ui/litellm-dashboard/src/components/router_settings/RouterSettingsForm.test.tsx @@ -97,8 +97,9 @@ describe("RouterSettingsForm", () => { expect(screen.getByText("Latency-Based Configuration")).toBeInTheDocument(); }); - it("should call onChange with the updated strategy when the selector changes", () => { + it("should call onChange with the updated strategy when the selector changes", async () => { const onChange = vi.fn(); + const user = userEvent.setup(); const props = { ...baseProps, onChange, @@ -106,9 +107,7 @@ describe("RouterSettingsForm", () => { }; render(); - const select = screen.getByTestId("strategy-select") as HTMLSelectElement; - select.value = "latency-based-routing"; - select.dispatchEvent(new Event("change", { bubbles: true })); + await user.selectOptions(screen.getByTestId("strategy-select"), "latency-based-routing"); expect(onChange).toHaveBeenCalledWith( expect.objectContaining({ selectedStrategy: "latency-based-routing" }) diff --git a/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx b/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx index 85b1dc21acf..01f839681f1 100644 --- a/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx +++ b/ui/litellm-dashboard/src/components/router_settings/RoutingStrategySelector.test.tsx @@ -1,5 +1,6 @@ import { describe, it, expect, vi } from "vitest"; import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import RoutingStrategySelector from "./RoutingStrategySelector"; // Ant Design's Select is complex to drive in JSDOM; swap it for a plain @@ -77,21 +78,16 @@ describe("RoutingStrategySelector", () => { it("should not render a description for a strategy that has none", () => { render(); - // "least-busy" has no entry in routingStrategyDescriptions - const select = screen.getByTestId("strategy-select"); - const leastBusyOption = Array.from(select.querySelectorAll("option")).find( - (o) => o.value === "least-busy" - ); - expect(leastBusyOption).toBeInTheDocument(); + // "least-busy" has no entry in routingStrategyDescriptions — it still renders without crashing + expect(screen.getByText("least-busy")).toBeInTheDocument(); }); - it("should call onStrategyChange with the selected strategy value", () => { + it("should call onStrategyChange with the selected strategy value", async () => { const onStrategyChange = vi.fn(); + const user = userEvent.setup(); render(); - const select = screen.getByTestId("strategy-select") as HTMLSelectElement; - select.value = "latency-based-routing"; - select.dispatchEvent(new Event("change", { bubbles: true })); + await user.selectOptions(screen.getByTestId("strategy-select"), "latency-based-routing"); expect(onStrategyChange).toHaveBeenCalledWith("latency-based-routing"); }); From 26e5482abb352b16dd57444705a6e44ef206d3d8 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 24 Feb 2026 15:52:15 -0800 Subject: [PATCH 27/29] address greptile review feedback (greploop iteration 2) - Wait for strategy select (API data loaded) before clicking Save - Assert specific payload content in setCallbacksCall - Move NotificationsManager import to top of file --- .../components/router_settings/index.test.tsx | 20 ++++++++++++------- 1 file changed, 13 insertions(+), 7 deletions(-) diff --git a/ui/litellm-dashboard/src/components/router_settings/index.test.tsx b/ui/litellm-dashboard/src/components/router_settings/index.test.tsx index 1920268207d..8edb0ac6e07 100644 --- a/ui/litellm-dashboard/src/components/router_settings/index.test.tsx +++ b/ui/litellm-dashboard/src/components/router_settings/index.test.tsx @@ -48,6 +48,7 @@ import { getRouterSettingsCall, setCallbacksCall, } from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; const mockCallbacksResponse = { router_settings: { @@ -141,29 +142,34 @@ describe("RouterSettings", () => { const user = userEvent.setup(); renderWithProviders(); + // Wait for the strategy select to appear — it only renders after getRouterSettingsCall resolves await waitFor(() => { - expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument(); + expect(screen.getByTestId("strategy-select")).toBeInTheDocument(); }); await user.click(screen.getByRole("button", { name: /save changes/i })); expect(setCallbacksCall).toHaveBeenCalledWith( "test-token", - expect.objectContaining({ router_settings: expect.any(Object) }) + expect.objectContaining({ + router_settings: expect.objectContaining({ + routing_strategy: "simple-shuffle", + }), + }) ); }); it("should show a success notification after saving", async () => { - const NotificationsManager = await import("@/components/molecules/notifications_manager"); const user = userEvent.setup(); renderWithProviders(); - await waitFor(() => - expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument() - ); + // Wait for data to load before interacting + await waitFor(() => { + expect(screen.getByTestId("strategy-select")).toBeInTheDocument(); + }); await user.click(screen.getByRole("button", { name: /save changes/i })); - expect(NotificationsManager.default.success).toHaveBeenCalledWith( + expect(NotificationsManager.success).toHaveBeenCalledWith( "router settings updated successfully" ); }); From 6ee50ff73e47da30c31526c22c8b58979248348c Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 24 Feb 2026 16:27:06 -0800 Subject: [PATCH 28/29] feat(proxy): tool policies - auto-discover tools + policy enforcement guardrail (#22041) * feat(proxy): tool policies - auto-discover tools, manage policies, guardrail enforcement - New LiteLLM_ToolTable in schema.prisma to store discovered tools - Auto-discovery: tools seen in LLM responses get upserted via ToolDiscoveryQueue (hooks into DBSpendUpdateWriter, same pipeline as spend tracking) - Management endpoints: GET /v1/tool/list, GET /v1/tool/{name}, POST /v1/tool/policy - ToolPolicyGuardrail: blocks tool_calls in responses based on policy setting - UI: Tool Policies page under Guardrails section with policy selector, filters by policy/team/key, live tail, sortable table - Unit tests for queue, writer, endpoints, guardrail * feat(tool-policies): track call_count + discover tools from request body and /messages API - Add call_count column to LiteLLM_ToolTable; incremented on every flush - Extract tools from request body too (not just response tool_calls): - OpenAI /chat/completions: tools[].function.name - Anthropic /messages pass-through: request_body.tools[].name - Show call_count column in UI table (sortable) - UI: drop dual_llm option, keep only trusted/blocked * fix: address greptile review feedback - Remove redundant @@index([tool_name]) from schema.prisma (tool_name has @unique which already creates an index) - Replace gen_random_uuid()::text with str(uuid.uuid4()) for portability - Rewrite test_tool_registry_writer.py to mock execute_raw/query_raw (actual implementation) instead of Prisma model methods - Fix test patches in test_tool_management_endpoints.py to target source modules since imports are inside function bodies - Add "Tool Policies" page title to ToolPolicies.tsx * fix: address greptile review round 2 - Replace NOW() with Python datetime parameter in tool_registry_writer (SQLite portability) - Fix cache key collision in tool_policy_guardrail: use null-byte separator instead of colon - Remove type==function filter from request-side tool extraction to match response-side behavior - Clear seen_tool_names on flush so call_count increments per batch cycle not per pod lifetime * fix: address greptile review round 3 - Fix test_seen_names_persist_across_flushes to match actual per-flush-cycle behavior - Update module docstring in tool_discovery_queue.py to accurately describe flush behavior - Add created_at/updated_at to raw SQL INSERT in batch_upsert_tools and update_tool_policy * fix: cache tool policies per tool name not per combination Previously the cache key was built from the full set of tool names in a request, so each unique combination of tools got its own cold cache entry and triggered a separate DB query. With N distinct tools across requests this was effectively a DB hit on every request. Now each tool name is cached individually. Cache hits are checked per tool, only missing tools are fetched from DB in a single batch query, and each result is cached separately. Once a tool's policy is warm, any subsequent request using that tool benefits from the cache regardless of what other tools are in the request. * Update ui/litellm-dashboard/src/components/ToolPolicies.tsx Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- litellm/constants.py | 1 + litellm/proxy/_types.py | 10 + litellm/proxy/db/db_spend_update_writer.py | 151 ++++++- .../tool_discovery_queue.py | 54 +++ litellm/proxy/db/tool_registry_writer.py | 179 ++++++++ .../guardrail_hooks/tool_policy/__init__.py | 16 + .../tool_policy/tool_policy_guardrail.py | 163 +++++++ .../tool_management_endpoints.py | 149 +++++++ litellm/proxy/proxy_server.py | 4 + litellm/proxy/schema.prisma | 20 + litellm/types/tool_management.py | 42 ++ .../test_tool_discovery_queue.py | 75 ++++ .../proxy/db/test_tool_registry_writer.py | 197 +++++++++ .../test_tool_policy_guardrail.py | 181 ++++++++ .../test_tool_management_endpoints.py | 149 +++++++ ui/litellm-dashboard/src/app/page.tsx | 3 + .../src/components/ToolPolicies.tsx | 415 ++++++++++++++++++ .../src/components/leftnav.tsx | 6 + .../src/components/networking.tsx | 54 +++ 19 files changed, 1863 insertions(+), 6 deletions(-) create mode 100644 litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py create mode 100644 litellm/proxy/db/tool_registry_writer.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/tool_policy/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py create mode 100644 litellm/proxy/management_endpoints/tool_management_endpoints.py create mode 100644 litellm/types/tool_management.py create mode 100644 tests/test_litellm/proxy/db/db_transaction_queue/test_tool_discovery_queue.py create mode 100644 tests/test_litellm/proxy/db/test_tool_registry_writer.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py create mode 100644 tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py create mode 100644 ui/litellm-dashboard/src/components/ToolPolicies.tsx diff --git a/litellm/constants.py b/litellm/constants.py index ee79f2fa56f..b1a0021bcc6 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -244,6 +244,7 @@ REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100)) # Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000)) +TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60)) # Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger. # Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire. MAX_SIZE_IN_MEMORY_QUEUE = int( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f354e28acd7..75b9f91acd9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -183,6 +183,7 @@ class LitellmTableNames(str, enum.Enum): KEY_TABLE_NAME = "LiteLLM_VerificationToken" PROXY_MODEL_TABLE_NAME = "LiteLLM_ProxyModelTable" MANAGED_FILE_TABLE_NAME = "LiteLLM_ManagedFileTable" + TOOL_TABLE_NAME = "LiteLLM_ToolTable" class Litellm_EntityType(enum.Enum): @@ -4123,6 +4124,15 @@ class SpendUpdateQueueItem(TypedDict, total=False): response_cost: Optional[float] +class ToolDiscoveryQueueItem(TypedDict, total=False): + tool_name: str + origin: Optional[str] # MCP server name or "user_defined" + created_by: Optional[str] + key_hash: Optional[str] # hash of virtual key that triggered discovery + team_id: Optional[str] # team that triggered discovery + key_alias: Optional[str] # human-readable key alias + + class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): unified_file_id: str file_object: Optional[OpenAIFileObject] = None diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index d7d6b2b1eb0..edf0cf0d397 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -13,7 +13,17 @@ import random import time import traceback from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast, overload +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Optional, + Union, + cast, + overload, +) import litellm from litellm._logging import verbose_proxy_logger @@ -23,18 +33,19 @@ from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy._types import ( DB_CONNECTION_ERROR_TYPES, BaseDailySpendTransaction, - DailyTagSpendTransaction, - DailyOrganizationSpendTransaction, - DailyTeamSpendTransaction, - DailyEndUserSpendTransaction, - DailyUserSpendTransaction, DailyAgentSpendTransaction, + DailyEndUserSpendTransaction, + DailyOrganizationSpendTransaction, + DailyTagSpendTransaction, + DailyTeamSpendTransaction, + DailyUserSpendTransaction, DBSpendUpdateTransactions, Litellm_EntityType, LiteLLM_UserTable, SpendLogsMetadata, SpendLogsPayload, SpendUpdateQueueItem, + ToolDiscoveryQueueItem, ) from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import ( DailySpendUpdateQueue, @@ -42,6 +53,9 @@ from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import ( from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue +from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import ( + ToolDiscoveryQueue, +) from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING if TYPE_CHECKING: @@ -67,6 +81,7 @@ class DBSpendUpdateWriter: self.redis_update_buffer = RedisUpdateBuffer(redis_cache=self.redis_cache) self.pod_lock_manager = PodLockManager() self.spend_update_queue = SpendUpdateQueue() + self.tool_discovery_queue = ToolDiscoveryQueue() self.daily_spend_update_queue = DailySpendUpdateQueue() self.daily_team_spend_update_queue = DailySpendUpdateQueue() self.daily_end_user_spend_update_queue = DailySpendUpdateQueue() @@ -222,6 +237,13 @@ class DBSpendUpdateWriter: ) ) + self._enqueue_tool_registry_upsert( + kwargs=kwargs, + completion_response=completion_response, + hashed_token=hashed_token, + team_id=team_id, + ) + verbose_proxy_logger.debug("Runs spend update on all tables") except Exception: verbose_proxy_logger.error( @@ -237,6 +259,104 @@ class DBSpendUpdateWriter: traceback.format_exc(), ) + def _enqueue_tool_registry_upsert( + self, + kwargs: Optional[dict], + completion_response: Optional[Any], + hashed_token: Optional[str] = None, + team_id: Optional[str] = None, + ) -> None: + """ + Extract tool names from the LLM request and response and enqueue them + for upsert into LiteLLM_ToolTable via ToolDiscoveryQueue. + + Handles four sources: + - MCP tools: standard_logging_object.mcp_tool_call_metadata.namespaced_tool_name + - Response tool_calls (OpenAI / Anthropic pass-through converted to OpenAI format): + completion_response.choices[].message.tool_calls[].function.name + - Request tools array (OpenAI format): kwargs["tools"][].function.name + - Request tools array (Anthropic /messages format): kwargs["passthrough_logging_payload"] + ["request_body"]["tools"][].name + """ + try: + if kwargs is None: + return + + # Extract key_alias from kwargs metadata if available + key_alias: Optional[str] = None + _litellm_params = kwargs.get("litellm_params") or {} + _metadata = _litellm_params.get("metadata") or {} + key_alias = _metadata.get("user_api_key_alias") or None + + def _enqueue(tool_name: str, origin: str = "user_defined") -> None: + self.tool_discovery_queue.add_update( + ToolDiscoveryQueueItem( + tool_name=tool_name, + origin=origin, + key_hash=hashed_token, + team_id=team_id, + key_alias=key_alias, + ) + ) + + # --- MCP tool calls --- + sl_object = kwargs.get("standard_logging_object") + if sl_object is not None: + mcp_metadata = ( + sl_object.get("metadata", {}) or {} + ).get("mcp_tool_call_metadata") + if mcp_metadata and isinstance(mcp_metadata, dict): + tool_name = mcp_metadata.get("namespaced_tool_name") or mcp_metadata.get("name") + mcp_server_name = mcp_metadata.get("mcp_server_name") + if tool_name: + _enqueue(tool_name, origin=mcp_server_name or "user_defined") + + # --- Tools from request body (OpenAI format: tools[].function.name) --- + request_tools = kwargs.get("tools") or [] + for tool_def in request_tools: + if not isinstance(tool_def, dict): + continue + fn = tool_def.get("function") or {} + name = fn.get("name") if isinstance(fn, dict) else None + if name: + _enqueue(name) + + # --- Tools from Anthropic /messages pass-through request body + # (Anthropic format: tools[].name, no "function" wrapper) --- + passthrough_payload = kwargs.get("passthrough_logging_payload") or {} + request_body = ( + passthrough_payload.get("request_body") + if isinstance(passthrough_payload, dict) + else None + ) or {} + for tool_def in request_body.get("tools") or []: + if not isinstance(tool_def, dict): + continue + name = tool_def.get("name") + if name: + _enqueue(name) + + # --- Response tool_calls (OpenAI format; Anthropic pass-through converts tool_use here) --- + if completion_response is not None and hasattr(completion_response, "choices"): + for choice in completion_response.choices or []: + message = getattr(choice, "message", None) + if message is None: + continue + tool_calls = getattr(message, "tool_calls", None) + if not tool_calls: + continue + for tc in tool_calls: + fn = getattr(tc, "function", None) + if fn is None: + continue + tool_name = getattr(fn, "name", None) + if tool_name: + _enqueue(tool_name) + except Exception as e: + verbose_proxy_logger.debug( + "_enqueue_tool_registry_upsert error (non-blocking): %s", e + ) + async def _update_key_db( self, response_cost: Optional[float], @@ -752,6 +872,25 @@ class DBSpendUpdateWriter: daily_spend_transactions=daily_agent_spend_update_transactions, ) + ################## Tool Registry Upserts ################## + await self._flush_tool_discovery_queue(prisma_client=prisma_client) + + async def _flush_tool_discovery_queue( + self, + prisma_client: PrismaClient, + ) -> None: + """Flush ToolDiscoveryQueue and batch-upsert new tools into LiteLLM_ToolTable.""" + from litellm.proxy.db.tool_registry_writer import batch_upsert_tools + + try: + items = self.tool_discovery_queue.flush() + if items: + await batch_upsert_tools(prisma_client=prisma_client, items=items) + except Exception as e: + verbose_proxy_logger.debug( + "_flush_tool_discovery_queue error (non-blocking): %s", e + ) + async def _commit_spend_updates_to_db( # noqa: PLR0915 self, prisma_client: PrismaClient, diff --git a/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py b/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py new file mode 100644 index 00000000000..16a3ada40f2 --- /dev/null +++ b/litellm/proxy/db/db_transaction_queue/tool_discovery_queue.py @@ -0,0 +1,54 @@ +""" +In-memory buffer for tool registry upserts. + +Unlike SpendUpdateQueue (which aggregates increments), ToolDiscoveryQueue +uses set-deduplication: each unique tool_name is only queued once per flush +cycle (~30s). The seen-set is cleared on every flush so that call_count +increments in subsequent cycles rather than stopping after the first flush. +""" + +from typing import List, Set + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import ToolDiscoveryQueueItem + + +class ToolDiscoveryQueue: + """ + In-memory buffer for tool registry upserts. + + Deduplicates by tool_name within each flush cycle: a tool is only queued + once per ~30s batch, so call_count increments once per flush cycle the + tool appears in (not once per invocation, but not once per pod lifetime + either). The seen-set is cleared on flush so subsequent batches can + re-count the same tool. + """ + + def __init__(self) -> None: + self._seen_tool_names: Set[str] = set() + self._pending: List[ToolDiscoveryQueueItem] = [] + + def add_update(self, item: ToolDiscoveryQueueItem) -> None: + """Enqueue a tool discovery item if tool_name has not been seen before.""" + tool_name = item.get("tool_name", "") + if not tool_name: + return + if tool_name in self._seen_tool_names: + verbose_proxy_logger.debug( + "ToolDiscoveryQueue: skipping already-seen tool %s", tool_name + ) + return + self._seen_tool_names.add(tool_name) + self._pending.append(item) + verbose_proxy_logger.debug( + "ToolDiscoveryQueue: queued new tool %s (origin=%s)", + tool_name, + item.get("origin"), + ) + + def flush(self) -> List[ToolDiscoveryQueueItem]: + """Return and clear all pending items. Resets seen-set so the next + flush cycle can re-count the same tools.""" + items, self._pending = self._pending, [] + self._seen_tool_names.clear() + return items diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py new file mode 100644 index 00000000000..4e0a8095a08 --- /dev/null +++ b/litellm/proxy/db/tool_registry_writer.py @@ -0,0 +1,179 @@ +""" +DB helpers for LiteLLM_ToolTable — the global tool registry. + +Tools are auto-discovered from LLM responses and upserted here. +Admins use the management endpoints to read and update call_policy. + +NOTE: Uses raw SQL (query_raw / execute_raw) instead of Prisma model methods +because the generated Prisma Python client may not have LiteLLM_ToolTable +when running against an older generated schema. +""" + +import uuid +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Dict, List, Optional + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import ToolDiscoveryQueueItem +from litellm.types.tool_management import LiteLLM_ToolTableRow, ToolCallPolicy + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + + +def _row_to_model(row: dict) -> LiteLLM_ToolTableRow: + return LiteLLM_ToolTableRow( + tool_id=row.get("tool_id", ""), + tool_name=row.get("tool_name", ""), + origin=row.get("origin"), + call_policy=row.get("call_policy", "untrusted"), + call_count=int(row.get("call_count") or 0), + assignments=row.get("assignments"), + key_hash=row.get("key_hash"), + team_id=row.get("team_id"), + key_alias=row.get("key_alias"), + created_at=row.get("created_at"), + updated_at=row.get("updated_at"), + created_by=row.get("created_by"), + updated_by=row.get("updated_by"), + ) + + +async def batch_upsert_tools( + prisma_client: "PrismaClient", + items: List[ToolDiscoveryQueueItem], +) -> None: + """ + Batch-upsert tool registry rows via raw SQL. + + On first insert: sets call_policy = "untrusted" (schema default), call_count = 1. + On conflict: increments call_count; preserves existing call_policy. + """ + if not items: + return + try: + data = [item for item in items if item.get("tool_name")] + if not data: + return + for item in data: + tool_name = item.get("tool_name", "") + origin = item.get("origin") or "user_defined" + created_by = item.get("created_by") or "system" + key_hash = item.get("key_hash") + team_id = item.get("team_id") + key_alias = item.get("key_alias") + now = datetime.now(timezone.utc).isoformat() + await prisma_client.db.execute_raw( + 'INSERT INTO "LiteLLM_ToolTable" ' + "(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) " + "VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8, $8) " + "ON CONFLICT (tool_name) DO UPDATE SET " + "call_count = \"LiteLLM_ToolTable\".call_count + 1, " + "updated_at = $8", + tool_name, + origin, + created_by, + key_hash, + team_id, + key_alias, + str(uuid.uuid4()), + now, + ) + verbose_proxy_logger.debug( + "tool_registry_writer: upserted %d tool(s)", len(data) + ) + except Exception as e: + verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e) + + +async def list_tools( + prisma_client: "PrismaClient", + call_policy: Optional[ToolCallPolicy] = None, +) -> List[LiteLLM_ToolTableRow]: + """Return all tools, optionally filtered by call_policy.""" + try: + if call_policy is not None: + rows = await prisma_client.db.query_raw( + 'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, ' + 'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by ' + 'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC', + call_policy, + ) + else: + rows = await prisma_client.db.query_raw( + 'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, ' + 'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by ' + 'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC', + ) + return [_row_to_model(row) for row in rows] + except Exception as e: + verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e) + return [] + + +async def get_tool( + prisma_client: "PrismaClient", + tool_name: str, +) -> Optional[LiteLLM_ToolTableRow]: + """Return a single tool row by tool_name.""" + try: + rows = await prisma_client.db.query_raw( + 'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, ' + 'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by ' + 'FROM "LiteLLM_ToolTable" WHERE tool_name = $1', + tool_name, + ) + if not rows: + return None + return _row_to_model(rows[0]) + except Exception as e: + verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e) + return None + + +async def update_tool_policy( + prisma_client: "PrismaClient", + tool_name: str, + call_policy: ToolCallPolicy, + updated_by: Optional[str], +) -> Optional[LiteLLM_ToolTableRow]: + """Update the call_policy for a tool. Upserts the row if it does not exist yet.""" + try: + _updated_by = updated_by or "system" + now = datetime.now(timezone.utc).isoformat() + await prisma_client.db.execute_raw( + 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) ' + "VALUES ($4, $1, $2, $3, $3, $5, $5) " + "ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5", + tool_name, + call_policy, + _updated_by, + str(uuid.uuid4()), + now, + ) + return await get_tool(prisma_client, tool_name) + except Exception as e: + verbose_proxy_logger.error("tool_registry_writer update_tool_policy error: %s", e) + return None + + +async def get_tools_by_names( + prisma_client: "PrismaClient", + tool_names: List[str], +) -> Dict[str, str]: + """ + Return a {tool_name: call_policy} map for the given tool names. + Used by the policy enforcement guardrail — single batch query, never N+1. + """ + if not tool_names: + return {} + try: + placeholders = ", ".join(f"${i+1}" for i in range(len(tool_names))) + rows = await prisma_client.db.query_raw( + f'SELECT tool_name, call_policy FROM "LiteLLM_ToolTable" WHERE tool_name IN ({placeholders})', + *tool_names, + ) + return {row["tool_name"]: row["call_policy"] for row in rows} + except Exception as e: + verbose_proxy_logger.error("tool_registry_writer get_tools_by_names error: %s", e) + return {} diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/__init__.py new file mode 100644 index 00000000000..5a43006e23c --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/__init__.py @@ -0,0 +1,16 @@ +import litellm +from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail): + from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import ( + ToolPolicyGuardrail, + ) + + _callback = ToolPolicyGuardrail( + guardrail_name=guardrail.get("guardrail_name", "tool_policy"), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(_callback) + return _callback diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py new file mode 100644 index 00000000000..87558566c42 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -0,0 +1,163 @@ +""" +Tool Policy Guardrail + +Reads call_policy from LiteLLM_ToolTable and enforces it on LLM requests/responses. + +Policy values: + "trusted" - allow through (no action) + "untrusted" - allow through (no action; default for newly discovered tools) + "blocked" - raise HTTPException, preventing the tool call + "dual_llm" - (Phase 3) send to second LLM for verification; currently treated as allowed + +Configuration in proxy config YAML: + guardrails: + - guardrail_name: "tool_policy" + litellm_params: + guardrail: tool_policy + mode: post_call + +or both pre and post call: + - guardrail_name: "tool_policy" + litellm_params: + guardrail: tool_policy + mode: during_call # runs before LLM and on response +""" + +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional + +from fastapi import HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +GUARDRAIL_NAME = "tool_policy" + + +class ToolPolicyGuardrail(CustomGuardrail): + """ + Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable. + + Tools with call_policy="blocked" are rejected before/after the LLM call. + Tools with call_policy="trusted" or "untrusted" pass through unchanged. + """ + + def __init__(self, **kwargs: Any) -> None: + if "supported_event_hooks" not in kwargs: + kwargs["supported_event_hooks"] = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + GuardrailEventHooks.during_call, + ] + super().__init__(**kwargs) + self._policy_cache: DualCache = DualCache() + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """ + Enforce tool policies on both request tools and response tool_calls. + + - input_type="request": check inputs["tools"] (tool definitions in the LLM request) + - input_type="response": check inputs["tool_calls"] (tool_calls in the LLM response) + + Raises HTTPException (400) if any tool is "blocked". + """ + if input_type == "request": + tools = inputs.get("tools") or [] + tool_names = [ + t["function"]["name"] + for t in tools + if isinstance(t, dict) + and isinstance(t.get("function"), dict) + and t["function"].get("name") + ] + else: # response + tool_calls = inputs.get("tool_calls") or [] + tool_names = [] + for tc in tool_calls: + fn = None + if isinstance(tc, dict): + fn = (tc.get("function") or {}).get("name") + elif hasattr(tc, "function"): + fn = getattr(tc.function, "name", None) + if fn: + tool_names.append(fn) + + if not tool_names: + return inputs + + policy_map = await self._get_policies_cached(tool_names) + + blocked = [name for name in tool_names if policy_map.get(name) == "blocked"] + if blocked: + verbose_proxy_logger.warning( + "ToolPolicyGuardrail: blocking tool(s) %s (policy=blocked)", blocked + ) + raise HTTPException( + status_code=400, + detail={ + "error": "Violated tool policy", + "blocked_tools": blocked, + "message": f"Tool(s) {blocked} are blocked by policy.", + }, + ) + + return inputs + + async def _get_policies_cached(self, tool_names: List[str]) -> Dict[str, str]: + """ + Batch-fetch call_policy for the given tool names. + + Caches per individual tool name (not per combination) so that adding + a new tool to a request doesn't invalidate the cached policies for all + the other tools already in the cache. + """ + from litellm.proxy.db.tool_registry_writer import get_tools_by_names + from litellm.proxy.proxy_server import prisma_client + + if not tool_names or prisma_client is None: + return {} + + result: Dict[str, str] = {} + cache_misses: List[str] = [] + + for name in tool_names: + cached = await self._policy_cache.async_get_cache(f"tool_policy:{name}") + if cached is not None and isinstance(cached, str): + result[name] = cached + else: + cache_misses.append(name) + + if cache_misses: + fetched = await get_tools_by_names( + prisma_client=prisma_client, tool_names=cache_misses + ) + for name, policy in fetched.items(): + result[name] = policy + await self._policy_cache.async_set_cache( + key=f"tool_policy:{name}", + value=policy, + ttl=TOOL_POLICY_CACHE_TTL_SECONDS, + ) + verbose_proxy_logger.debug( + "ToolPolicyGuardrail: fetched %d policies from DB (cache hits: %d)", + len(cache_misses), + len(tool_names) - len(cache_misses), + ) + + return result diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py new file mode 100644 index 00000000000..89880c9a4ec --- /dev/null +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -0,0 +1,149 @@ +""" +TOOL POLICY MANAGEMENT + +All /tool management endpoints + +GET /v1/tool/list - List all discovered tools and their policies +GET /v1/tool/{tool_name} - Get a single tool's details +POST /v1/tool/policy - Update the call_policy for a tool +""" + +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.types.tool_management import ( + LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolListResponse, + ToolPolicyUpdateRequest, + ToolPolicyUpdateResponse, +) + +router = APIRouter() + + +@router.get( + "/v1/tool/list", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ToolListResponse, +) +async def list_tools( + call_policy: Optional[ToolCallPolicy] = None, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + List all auto-discovered tools and their call policies. + + Parameters: + - call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked" + """ + from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, detail=CommonProxyErrors.db_not_connected_error.value + ) + + try: + tools = await db_list_tools(prisma_client=prisma_client, call_policy=call_policy) + return ToolListResponse(tools=tools, total=len(tools)) + except Exception as e: + verbose_proxy_logger.exception("Error listing tools: %s", e) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get( + "/v1/tool/{tool_name:path}", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], + response_model=LiteLLM_ToolTableRow, +) +async def get_tool( + tool_name: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get details for a single tool. + + Parameters: + - tool_name: The tool name (supports namespaced names with slashes) + """ + from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, detail=CommonProxyErrors.db_not_connected_error.value + ) + + try: + tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name) + if tool is None: + raise HTTPException( + status_code=404, detail=f"Tool '{tool_name}' not found" + ) + return tool + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception("Error getting tool: %s", e) + raise HTTPException(status_code=500, detail=str(e)) + + +@router.post( + "/v1/tool/policy", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ToolPolicyUpdateResponse, +) +async def update_tool_policy( + data: ToolPolicyUpdateRequest, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Set the call policy for a tool. + + Parameters: + - tool_name: str - The tool to update + - call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked" + + Setting a tool to "blocked" will cause the ToolPolicyGuardrail to remove + that tool_call from LLM responses before returning them to the client. + """ + from litellm.proxy.db.tool_registry_writer import ( + update_tool_policy as db_update_tool_policy, + ) + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, detail=CommonProxyErrors.db_not_connected_error.value + ) + + try: + updated = await db_update_tool_policy( + prisma_client=prisma_client, + tool_name=data.tool_name, + call_policy=data.call_policy, + updated_by=user_api_key_dict.user_id, + ) + if updated is None: + raise HTTPException( + status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'" + ) + return ToolPolicyUpdateResponse( + tool_name=updated.tool_name, + call_policy=updated.call_policy, + updated=True, + ) + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception("Error updating tool policy: %s", e) + raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 905a4be35e4..36b6bcc7707 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -410,6 +410,9 @@ from litellm.proxy.management_endpoints.team_endpoints import ( update_team, validate_membership, ) +from litellm.proxy.management_endpoints.tool_management_endpoints import ( + router as tool_management_router, +) from litellm.proxy.management_endpoints.ui_sso import ( get_disabled_non_admin_personal_key_creation, ) @@ -12882,6 +12885,7 @@ app.include_router(budget_management_router) app.include_router(model_management_router) app.include_router(model_access_group_management_router) app.include_router(tag_management_router) +app.include_router(tool_management_router) app.include_router(cost_tracking_settings_router) app.include_router(router_settings_router) app.include_router(fallback_management_router) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 50c0a55a875..23917cf7c7f 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1051,6 +1051,26 @@ model LiteLLM_PolicyAttachmentTable { updated_by String? } +// Global tool registry - auto-discovered from LLM responses; admins set call_policy here +model LiteLLM_ToolTable { + tool_id String @id @default(uuid()) + tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space" + origin String? // MCP server name or "user_defined" + call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked" + call_count Int @default(0) // cumulative number of times this tool was seen + assignments Json? @default("{}") + key_hash String? // hash of the virtual key that first called this tool + team_id String? // team that first called this tool + key_alias String? // human-readable alias of the virtual key + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@index([call_policy]) + @@index([team_id]) +} + //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py new file mode 100644 index 00000000000..8704ff27759 --- /dev/null +++ b/litellm/types/tool_management.py @@ -0,0 +1,42 @@ +""" +Pydantic models for Tool Policy management endpoints. +""" + +from datetime import datetime +from typing import Dict, List, Literal, Optional + +from pydantic import BaseModel + +ToolCallPolicy = Literal["trusted", "untrusted", "dual_llm", "blocked"] + + +class LiteLLM_ToolTableRow(BaseModel): + tool_id: str + tool_name: str + origin: Optional[str] = None + call_policy: ToolCallPolicy = "untrusted" + call_count: int = 0 + assignments: Optional[Dict] = None + key_hash: Optional[str] = None + team_id: Optional[str] = None + key_alias: Optional[str] = None + created_at: Optional[datetime] = None + updated_at: Optional[datetime] = None + created_by: Optional[str] = None + updated_by: Optional[str] = None + + +class ToolListResponse(BaseModel): + tools: List[LiteLLM_ToolTableRow] + total: int + + +class ToolPolicyUpdateRequest(BaseModel): + tool_name: str + call_policy: ToolCallPolicy + + +class ToolPolicyUpdateResponse(BaseModel): + tool_name: str + call_policy: ToolCallPolicy + updated: bool diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_tool_discovery_queue.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_tool_discovery_queue.py new file mode 100644 index 00000000000..defdb3834d8 --- /dev/null +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_tool_discovery_queue.py @@ -0,0 +1,75 @@ +""" +Unit tests for ToolDiscoveryQueue. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import ( + ToolDiscoveryQueue, +) + + +@pytest.fixture +def queue(): + return ToolDiscoveryQueue() + + +def test_add_single_tool(queue): + queue.add_update({"tool_name": "my_tool", "origin": "user_defined"}) + items = queue.flush() + assert len(items) == 1 + assert items[0]["tool_name"] == "my_tool" + assert items[0]["origin"] == "user_defined" + + +def test_deduplication_same_name(queue): + """Adding the same tool_name twice should only keep the first.""" + queue.add_update({"tool_name": "tool_a", "origin": "mcp_server"}) + queue.add_update({"tool_name": "tool_a", "origin": "user_defined"}) + items = queue.flush() + assert len(items) == 1 + assert items[0]["origin"] == "mcp_server" # first wins + + +def test_deduplication_different_names(queue): + queue.add_update({"tool_name": "tool_a"}) + queue.add_update({"tool_name": "tool_b"}) + items = queue.flush() + assert len(items) == 2 + names = {i["tool_name"] for i in items} + assert names == {"tool_a", "tool_b"} + + +def test_flush_clears_pending(queue): + queue.add_update({"tool_name": "tool_x"}) + items1 = queue.flush() + assert len(items1) == 1 + items2 = queue.flush() + assert len(items2) == 0 + + +def test_seen_names_reset_after_flush(queue): + """Seen-set is cleared on flush so the same tool can re-enter the next cycle.""" + queue.add_update({"tool_name": "tool_a"}) + queue.flush() + queue.add_update({"tool_name": "tool_a"}) # same tool, new cycle + items = queue.flush() + assert len(items) == 1 + assert items[0]["tool_name"] == "tool_a" + + +def test_empty_tool_name_ignored(queue): + queue.add_update({"tool_name": ""}) + queue.add_update({"tool_name": None}) # type: ignore[arg-type] + items = queue.flush() + assert len(items) == 0 + + +def test_flush_returns_list(queue): + result = queue.flush() + assert isinstance(result, list) diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py new file mode 100644 index 00000000000..44f9e32058a --- /dev/null +++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py @@ -0,0 +1,197 @@ +""" +Unit tests for tool_registry_writer.py — uses a mock prisma client +that exposes execute_raw / query_raw (matching the actual raw-SQL implementation). +""" + +import os +import sys +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.proxy.db.tool_registry_writer import ( + batch_upsert_tools, + get_tool, + get_tools_by_names, + list_tools, + update_tool_policy, +) + + +def _make_prisma(query_rows=None): + """Return a minimal mock prisma_client with execute_raw / query_raw.""" + default_row = { + "tool_id": "uuid-1", + "tool_name": "my_tool", + "origin": "user_defined", + "call_policy": "untrusted", + "call_count": 1, + "assignments": {}, + "key_hash": None, + "team_id": None, + "key_alias": None, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": None, + "updated_by": None, + } + rows = query_rows if query_rows is not None else [default_row] + + prisma = MagicMock() + prisma.db.execute_raw = AsyncMock(return_value=None) + prisma.db.query_raw = AsyncMock(return_value=rows) + return prisma + + +@pytest.mark.asyncio +async def test_batch_upsert_tools_calls_execute_raw(): + prisma = _make_prisma() + items = [{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}] + await batch_upsert_tools(prisma, items) + prisma.db.execute_raw.assert_awaited_once() + call_args = prisma.db.execute_raw.call_args + sql = call_args.args[0] + assert "LiteLLM_ToolTable" in sql + assert "ON CONFLICT" in sql + + +@pytest.mark.asyncio +async def test_batch_upsert_tools_empty_list(): + prisma = _make_prisma() + await batch_upsert_tools(prisma, []) + prisma.db.execute_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_batch_upsert_tools_skips_empty_names(): + prisma = _make_prisma() + items = [{"tool_name": "", "origin": None}, {"tool_name": None}] # type: ignore[list-item] + await batch_upsert_tools(prisma, items) + prisma.db.execute_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_batch_upsert_multiple_tools_calls_execute_raw_per_tool(): + prisma = _make_prisma() + items = [ + {"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}, + {"tool_name": "tool_b", "origin": "user_defined", "created_by": "alice"}, + ] + await batch_upsert_tools(prisma, items) + assert prisma.db.execute_raw.await_count == 2 + + +@pytest.mark.asyncio +async def test_list_tools_no_filter(): + row = { + "tool_id": "id1", + "tool_name": "tool_a", + "origin": "mcp", + "call_policy": "untrusted", + "call_count": 5, + "assignments": {}, + "key_hash": None, + "team_id": None, + "key_alias": None, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": None, + "updated_by": None, + } + prisma = _make_prisma(query_rows=[row]) + result = await list_tools(prisma) + assert len(result) == 1 + assert result[0].tool_name == "tool_a" + assert result[0].call_count == 5 + prisma.db.query_raw.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_list_tools_with_policy_filter(): + row = { + "tool_id": "id1", + "tool_name": "blocked_tool", + "origin": None, + "call_policy": "blocked", + "call_count": 2, + "assignments": None, + "key_hash": None, + "team_id": None, + "key_alias": None, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": None, + "updated_by": None, + } + prisma = _make_prisma(query_rows=[row]) + result = await list_tools(prisma, call_policy="blocked") + assert result[0].call_policy == "blocked" + call_args = prisma.db.query_raw.call_args + sql = call_args.args[0] + assert "WHERE call_policy" in sql + + +@pytest.mark.asyncio +async def test_get_tool_found(): + prisma = _make_prisma() + result = await get_tool(prisma, "my_tool") + assert result is not None + assert result.tool_name == "my_tool" + prisma.db.query_raw.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_get_tool_not_found(): + prisma = _make_prisma(query_rows=[]) + result = await get_tool(prisma, "nonexistent") + assert result is None + + +@pytest.mark.asyncio +async def test_update_tool_policy_calls_execute_raw(): + row = { + "tool_id": "uuid-1", + "tool_name": "my_tool", + "origin": "user_defined", + "call_policy": "blocked", + "call_count": 1, + "assignments": {}, + "key_hash": None, + "team_id": None, + "key_alias": None, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + "created_by": None, + "updated_by": "admin", + } + prisma = _make_prisma(query_rows=[row]) + result = await update_tool_policy(prisma, "my_tool", "blocked", "admin") + assert result is not None + assert result.call_policy == "blocked" + prisma.db.execute_raw.assert_awaited_once() + call_args = prisma.db.execute_raw.call_args + sql = call_args.args[0] + assert "ON CONFLICT" in sql + assert "call_policy" in sql + + +@pytest.mark.asyncio +async def test_get_tools_by_names_returns_policy_map(): + rows = [ + {"tool_name": "tool_a", "call_policy": "trusted"}, + {"tool_name": "tool_b", "call_policy": "blocked"}, + ] + prisma = _make_prisma(query_rows=rows) + result = await get_tools_by_names(prisma, ["tool_a", "tool_b"]) + assert result == {"tool_a": "trusted", "tool_b": "blocked"} + + +@pytest.mark.asyncio +async def test_get_tools_by_names_empty_list(): + prisma = _make_prisma() + result = await get_tools_by_names(prisma, []) + assert result == {} + prisma.db.query_raw.assert_not_awaited() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py new file mode 100644 index 00000000000..c6a81efbf0b --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py @@ -0,0 +1,181 @@ +""" +Unit tests for ToolPolicyGuardrail. +""" + +import os +import sys +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +sys.path.insert(0, os.path.abspath("../../../../../..")) + +from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import ( + ToolPolicyGuardrail, +) +from litellm.types.guardrails import GuardrailEventHooks + + +@pytest.fixture +def guardrail(): + return ToolPolicyGuardrail() + + +# --- helpers --- + +def _tool_request_inputs(tool_names: list) -> dict: + return { + "tools": [ + {"type": "function", "function": {"name": name, "description": ""}} + for name in tool_names + ] + } + + +def _tool_response_inputs(tool_names: list) -> dict: + return { + "tool_calls": [ + {"type": "function", "function": {"name": name}} + for name in tool_names + ] + } + + +# --- tests --- + + +def test_guardrail_supports_pre_and_post_call(guardrail): + hooks = guardrail.supported_event_hooks + assert GuardrailEventHooks.pre_call in hooks + assert GuardrailEventHooks.post_call in hooks + + +@pytest.mark.asyncio +async def test_no_tools_in_request_passes_through(guardrail): + inputs: Any = {"tools": []} + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="request" + ) + assert result is inputs + + +@pytest.mark.asyncio +async def test_no_tool_calls_in_response_passes_through(guardrail): + inputs: Any = {"tool_calls": []} + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="response" + ) + assert result is inputs + + +@pytest.mark.asyncio +async def test_untrusted_tools_pass_through(guardrail): + policy_map = {"search": "untrusted", "read_file": "trusted"} + with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + inputs: Any = _tool_request_inputs(["search", "read_file"]) + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="request" + ) + assert result is inputs + + +@pytest.mark.asyncio +async def test_blocked_tool_in_request_raises_http_exception(guardrail): + policy_map = {"dangerous_tool": "blocked"} + with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + inputs: Any = _tool_request_inputs(["dangerous_tool"]) + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="request" + ) + assert exc_info.value.status_code == 400 + assert "dangerous_tool" in exc_info.value.detail["blocked_tools"] + + +@pytest.mark.asyncio +async def test_blocked_tool_in_response_raises_http_exception(guardrail): + policy_map = {"exfil_tool": "blocked"} + with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + inputs: Any = _tool_response_inputs(["exfil_tool"]) + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="response" + ) + assert exc_info.value.status_code == 400 + assert "exfil_tool" in exc_info.value.detail["blocked_tools"] + + +@pytest.mark.asyncio +async def test_mixed_blocked_and_allowed_raises_for_blocked(guardrail): + policy_map = {"safe_tool": "trusted", "bad_tool": "blocked"} + with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + inputs: Any = _tool_request_inputs(["safe_tool", "bad_tool"]) + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="request" + ) + blocked = exc_info.value.detail["blocked_tools"] + assert "bad_tool" in blocked + assert "safe_tool" not in blocked + + +@pytest.mark.asyncio +async def test_tool_not_in_db_passes_through(guardrail): + """Tools not found in the DB (no entry) should not be blocked.""" + with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value={})): + inputs: Any = _tool_request_inputs(["unknown_tool"]) + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="request" + ) + assert result is inputs + + +@pytest.mark.asyncio +async def test_get_policies_cached_uses_cache(guardrail): + """Second call with same tool names should return the cached result.""" + policy_map = {"tool_a": "trusted"} + with patch( + "litellm.proxy.db.tool_registry_writer.get_tools_by_names", + new=AsyncMock(return_value=policy_map), + ) as mock_db, patch( + "litellm.proxy.proxy_server.prisma_client", + new=MagicMock(), + ): + # first call — should hit DB + result1 = await guardrail._get_policies_cached(["tool_a"]) + assert result1 == policy_map + + # second call — should hit cache, not DB again + result2 = await guardrail._get_policies_cached(["tool_a"]) + assert result2 == policy_map + + assert mock_db.call_count == 1 + + +@pytest.mark.asyncio +async def test_get_policies_cached_no_prisma(guardrail): + """Without a prisma client, returns empty dict.""" + with patch( + "litellm.proxy.proxy_server.prisma_client", + None, + ): + result = await guardrail._get_policies_cached(["tool_a"]) + assert result == {} + + +@pytest.mark.asyncio +async def test_response_tool_calls_as_objects(guardrail): + """tool_calls that are objects (not dicts) with .function.name should work.""" + policy_map = {"obj_tool": "blocked"} + with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + fn = MagicMock() + fn.name = "obj_tool" + tc = MagicMock() + tc.function = fn + inputs: Any = {"tool_calls": [tc]} + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="response" + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py new file mode 100644 index 00000000000..6f1d373fdee --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py @@ -0,0 +1,149 @@ +""" +Unit tests for tool management endpoints (/v1/tool/*). +Uses FastAPI TestClient with mocked DB functions. + +Patches target the source modules (litellm.proxy.db.tool_registry_writer.* +and litellm.proxy.proxy_server.prisma_client) because the endpoint code +imports these inside function bodies to avoid circular imports. +""" + +import os +import sys +from datetime import datetime, timezone +from typing import Optional +from unittest.mock import AsyncMock, MagicMock, patch + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.proxy.management_endpoints.tool_management_endpoints import router +from litellm.types.tool_management import LiteLLM_ToolTableRow + +# --- helpers --- + + +def _make_tool_row( + tool_name: str = "my_tool", + call_policy: str = "untrusted", + origin: Optional[str] = None, +) -> LiteLLM_ToolTableRow: + now = datetime.now(timezone.utc) + return LiteLLM_ToolTableRow( + tool_id="uuid-1", + tool_name=tool_name, + origin=origin, + call_policy=call_policy, # type: ignore[arg-type] + assignments={}, + created_at=now, + updated_at=now, + ) + + +def _make_app() -> FastAPI: + """Build a minimal FastAPI app with the tool management router.""" + app = FastAPI() + app.include_router(router) + return app + + +# Stub the auth dependency so we don't need a real proxy running. +def _override_auth(): + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth(api_key="sk-test", user_id="admin") + + +# A real (non-None) prisma stub for truthiness checks. +_MOCK_PRISMA = MagicMock() + + +# --- test class --- + + +class TestToolManagementEndpoints: + def setup_method(self): + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app = _make_app() + app.dependency_overrides[user_api_key_auth] = _override_auth + self.client = TestClient(app, raise_server_exceptions=True) + + @patch( + "litellm.proxy.db.tool_registry_writer.list_tools", + new_callable=AsyncMock, + ) + @patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA) + def test_list_tools_returns_200(self, mock_db_list): + mock_db_list.return_value = [_make_tool_row()] + + resp = self.client.get("/v1/tool/list") + assert resp.status_code == 200 + body = resp.json() + assert body["total"] == 1 + assert body["tools"][0]["tool_name"] == "my_tool" + + @patch( + "litellm.proxy.db.tool_registry_writer.list_tools", + new_callable=AsyncMock, + ) + @patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA) + def test_list_tools_with_policy_filter(self, mock_db_list): + mock_db_list.return_value = [_make_tool_row(call_policy="blocked")] + + resp = self.client.get("/v1/tool/list?call_policy=blocked") + assert resp.status_code == 200 + assert resp.json()["tools"][0]["call_policy"] == "blocked" + + @patch( + "litellm.proxy.db.tool_registry_writer.get_tool", + new_callable=AsyncMock, + ) + @patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA) + def test_get_tool_found(self, mock_db_get): + mock_db_get.return_value = _make_tool_row(tool_name="tool_a") + + resp = self.client.get("/v1/tool/tool_a") + assert resp.status_code == 200 + assert resp.json()["tool_name"] == "tool_a" + + @patch( + "litellm.proxy.db.tool_registry_writer.get_tool", + new_callable=AsyncMock, + ) + @patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA) + def test_get_tool_not_found_returns_404(self, mock_db_get): + mock_db_get.return_value = None + + resp = self.client.get("/v1/tool/nonexistent", follow_redirects=True) + assert resp.status_code == 404 + + @patch( + "litellm.proxy.db.tool_registry_writer.update_tool_policy", + new_callable=AsyncMock, + ) + @patch("litellm.proxy.proxy_server.prisma_client", _MOCK_PRISMA) + def test_update_tool_policy_blocked(self, mock_db_update): + mock_db_update.return_value = _make_tool_row(call_policy="blocked") + + resp = self.client.post( + "/v1/tool/policy", + json={"tool_name": "my_tool", "call_policy": "blocked"}, + ) + assert resp.status_code == 200 + body = resp.json() + assert body["call_policy"] == "blocked" + assert body["updated"] is True + + @patch("litellm.proxy.proxy_server.prisma_client", None) + def test_list_tools_no_db_returns_500(self): + resp = self.client.get("/v1/tool/list") + assert resp.status_code == 500 + + def test_update_tool_policy_invalid_policy_returns_422(self): + resp = self.client.post( + "/v1/tool/policy", + json={"tool_name": "my_tool", "call_policy": "invalid_value"}, + ) + assert resp.status_code == 422 diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index fb749d7afb0..258c2ccb0e0 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -38,6 +38,7 @@ import Usage from "@/components/usage"; import UserDashboard from "@/components/user_dashboard"; import { AccessGroupsPage } from "@/components/AccessGroups/AccessGroupsPage"; import VectorStoreManagement from "@/components/vector_store_management"; +import ToolPolicies from "@/components/ToolPolicies"; import SpendLogsTable from "@/components/view_logs"; import ViewUserDashboard from "@/components/view_users"; import { ThemeProvider } from "@/contexts/ThemeContext"; @@ -548,6 +549,8 @@ function CreateKeyPageContent() { ) : page == "vector-stores" ? ( + ) : page == "tool-policies" ? ( + ) : page == "guardrails-monitor" ? ( ) : page == "new_usage" ? ( diff --git a/ui/litellm-dashboard/src/components/ToolPolicies.tsx b/ui/litellm-dashboard/src/components/ToolPolicies.tsx new file mode 100644 index 00000000000..0e3f5434e7f --- /dev/null +++ b/ui/litellm-dashboard/src/components/ToolPolicies.tsx @@ -0,0 +1,415 @@ +"use client"; + +import React, { useCallback, useDeferredValue, useEffect, useState } from "react"; +import { Select, Switch, Tooltip } from "antd"; +import { Select, Tooltip } from "antd"; +import { + Table, + TableHead, + TableHeaderCell, + TableBody, + TableRow, + TableCell, +} from "@tremor/react"; +import { TimeCell } from "./view_logs/time_cell"; +import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; +import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; +import FilterComponent, { FilterOption } from "./molecules/filter"; +import { fetchToolsList, updateToolPolicy, ToolRow } from "./networking"; + +const POLICY_OPTIONS = [ + { value: "trusted", label: "trusted", color: "#065f46", bg: "#d1fae5", border: "#6ee7b7" }, + { value: "blocked", label: "blocked", color: "#991b1b", bg: "#fee2e2", border: "#fca5a5" }, +] as const; + +type PolicyValue = "trusted" | "blocked"; + +const policyStyle = (p: string) => + POLICY_OPTIONS.find((o) => o.value === p) ?? POLICY_OPTIONS[1]; + +type SortField = "tool_name" | "call_policy" | "team_id" | "key_alias" | "created_at" | "call_count"; + +interface FilterValues { + [key: string]: string; +} + +interface ToolPoliciesProps { + accessToken: string | null; + userRole?: string; +} + +const PolicySelect: React.FC<{ + value: string; + toolName: string; + saving: boolean; + onChange: (toolName: string, policy: string) => void; +}> = ({ value, toolName, saving, onChange }) => { + const style = policyStyle(value); + return ( + { setSearchTerm(e.target.value); setCurrentPage(1); }} + /> + + + + + +
+ Live Tail + +
+ + + + +
+ + Showing {filtered.length === 0 ? 0 : (currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, filtered.length)} of {filtered.length} results + + Page {currentPage} of {totalPages} +
+ + +
+
+ + + {/* Filter row */} +
+ +
+ + + {/* Auto-refresh banner */} + {isLiveTail && ( +
+ Auto-refreshing every 15 seconds + +
+ )} + + {error && ( +
{error}
+ )} + + {/* Table */} + + + + + + + + + Key Hash + + Origin + + + + {loading ? ( + + Loading tools… + + ) : paginated.length === 0 ? ( + + + No tools discovered yet. Make a chat completion that returns tool_calls to start auto-discovery. + + + ) : ( + paginated.map((tool) => ( + + + + + + + + {tool.tool_name} + + + + + + + + {(tool.call_count ?? 0).toLocaleString()} + + + + {tool.team_id ?? "-"} + + + + + + {tool.key_hash ?? "-"} + + + + + + {tool.key_alias ?? "-"} + + + + + {tool.origin ?? "-"} + + + + )) + )} + +
+ + {/* Bottom pagination (only when > 1 page) */} + {totalPages > 1 && ( +
+ Showing {(currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, sorted.length)} of {sorted.length} +
+ + +
+
+ )} + + + ); +}; + +export default ToolPolicies; diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index da3ca2a8bae..2cbeb22ec81 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -134,6 +134,12 @@ const menuGroups: MenuGroup[] = [ label: "Vector Stores", icon: , }, + { + key: "tool-policies", + page: "tool-policies", + label: "Tool Policies", + icon: , + }, ], }, ], diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 6ffd744cce9..8536a584ee1 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -9854,3 +9854,57 @@ export const checkGdprCompliance = async ( } return response.json(); }; + +export interface ToolRow { + tool_id: string; + tool_name: string; + origin?: string; + call_policy: string; + call_count?: number; + assignments?: Record; + key_hash?: string; + team_id?: string; + key_alias?: string; + created_at?: string; + updated_at?: string; + created_by?: string; + updated_by?: string; +} + +export const fetchToolsList = async (accessToken: string): Promise => { + const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/tool/list` : `/v1/tool/list`; + const response = await fetch(url, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + if (!response.ok) { + const errorData = await response.text(); + throw new Error(errorData); + } + const data = await response.json(); + return data.tools ?? []; +}; + +export const updateToolPolicy = async ( + accessToken: string, + toolName: string, + callPolicy: string +): Promise => { + const url = proxyBaseUrl ? `${proxyBaseUrl}/v1/tool/policy` : `/v1/tool/policy`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ tool_name: toolName, call_policy: callPolicy }), + }); + if (!response.ok) { + const errorData = await response.text(); + throw new Error(errorData); + } + return response.json(); +}; From 360643e21315eb02ca790eb22effb87a6d48167b Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 24 Feb 2026 16:40:04 -0800 Subject: [PATCH 29/29] [Feat] UI - Allow using AI to understand Usage patterns (#22042) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Add Ask AI chat component to Usage page - Create UsageAIChatModal component with streaming chat interface - Integrate with existing model hub for model selection - Pass usage data context (spend, models, providers, keys) to AI - Add Ask AI button next to Export Data button in global view - Add tests for the new component and integration Co-authored-by: Ishaan Jaff * Convert Ask AI from modal to right-side sliding panel - Replace UsageAIChatModal with UsageAIChatPanel - Panel slides in from right side, usage page stays visible - Full-height panel with header, model selector, chat area, and input - Smooth CSS transition for open/close animation - Update tests for new panel component (34 tests passing) Co-authored-by: Ishaan Jaff * Remove build output directory from tracking Co-authored-by: Ishaan Jaff * Add backend AI usage chat endpoint with tool calling Backend: - New /usage/ai/chat SSE streaming endpoint - AI agent has get_usage_data tool that queries /user/daily/activity/aggregated - Follows same architecture as policy AI suggest (litellm.acompletion + tools) - Non-admin users are restricted to their own data - 12 backend unit tests Frontend: - Panel now calls /usage/ai/chat backend endpoint via SSE - Removed direct OpenAI client calls from frontend - Added usageAiChatStream networking function following enrichPolicyTemplateStream pattern Co-authored-by: Ishaan Jaff * Make model selection optional, default to gpt-4o-mini on backend Co-authored-by: Ishaan Jaff * Add team/tag tools, status indicators, and improved AI agent - AI agent now has 3 tools: get_usage_data, get_team_usage_data, get_tag_usage_data - Stream status events (Thinking... Fetching... Analyzing...) to UI - Frontend shows spinner + status text during tool execution - Better system prompt guiding tool selection - Entity summariser for team/tag data with ranked breakdowns - 13 backend tests, 34 frontend tests passing Co-authored-by: Ishaan Jaff * Fix: inject today's date into system prompt so AI resolves relative dates correctly Co-authored-by: Ishaan Jaff * Show tool calls as distinct steps + render markdown in responses - Backend emits tool_call events with tool_name, label, args, and status - Frontend shows each tool call as a step with ✓/spinner/✗ indicator - Tool call steps show icon, label, date range, and filters - AI responses rendered with ReactMarkdown (bold, lists, tables, code) - Cursor-like UX: Thinking → tool calls → Analyzing → streamed answer Co-authored-by: Ishaan Jaff * Refactor backend for code quality: proper types, constants, all functions ≤50 LOC - TypedDict for SSE events (SSEStatusEvent, SSEToolCallEvent, etc.) and ToolHandler - Constants for table names, entity fields, temperature, page sizes, top-N limits - Shared _query_activity() eliminates duplicated fetch logic - _accumulate_breakdown() + _ranked_lines() replace inline aggregation loops - Extracted _process_tool_call() and _stream_final_response() from main stream fn - Black + Ruff clean, all 15 functions verified ≤50 LOC - Replaced Tremor Button with Antd Button in panel (Tremor deprecated per AGENTS.md) Co-authored-by: Ishaan Jaff * Address greptile review: security fixes and input validation - Restrict team/tag tools to admin-only users (non-admins only get get_usage_data) - Constrain ChatMessage.role to Literal['user', 'assistant'] to prevent system prompt injection - Add test for base tools restriction (non-admin gets 1 tool, admin gets 3) - Issues 3 (unused imports) and 4 (inline datetime) were already fixed in prior commit Co-authored-by: Ishaan Jaff * Address greptile round 2: sanitize errors, defense-in-depth allowlist, revert tsconfig - Sanitize error messages: generic 'An internal error occurred' sent to client, full exception logged server-side via verbose_proxy_logger - Defense-in-depth: _process_tool_call validates fn_name against role-based allowlist before dispatch (even though LLM only receives allowed tools) - Revert tsconfig.json jsx back to 'preserve' (Next.js recommended default) Co-authored-by: Ishaan Jaff * Role-scoped system prompt + additional test coverage - System prompt is now role-aware: admin sees all 3 tool descriptions, non-admin only sees get_usage_data (consistent with tool filtering) - Added tests: non-admin prompt excludes team/tag tools, date injection - 15 backend tests, 34 frontend tests passing Co-authored-by: Ishaan Jaff * Fix LLM arg validation + cap conversation size at 20 messages - _resolve_fetch_kwargs uses .get() with ValueError for missing dates (handles malformed LLM tool arguments gracefully) - MAX_CHAT_MESSAGES = 20 constant; backend truncates to last 20 - Frontend also sends only last 20 messages per request - Prevents excessive token usage and context-length errors Co-authored-by: Ishaan Jaff --------- Co-authored-by: Cursor Agent Co-authored-by: Ishaan Jaff --- .../usage_endpoints/__init__.py | 9 + .../usage_endpoints/ai_usage_chat.py | 578 ++++++++++++++++++ .../usage_endpoints/endpoints.py | 65 ++ litellm/proxy/proxy_server.py | 2 + .../usage_endpoints/__init__.py | 0 .../usage_endpoints/test_ai_usage_chat.py | 402 ++++++++++++ ui/litellm-dashboard/package-lock.json | 151 +---- .../components/UsageAIChatPanel.test.tsx | 85 +++ .../UsagePage/components/UsageAIChatPanel.tsx | 402 ++++++++++++ .../components/UsagePageView.test.tsx | 26 + .../UsagePage/components/UsagePageView.tsx | 51 +- .../src/components/networking.tsx | 76 +++ 12 files changed, 1704 insertions(+), 143 deletions(-) create mode 100644 litellm/proxy/management_endpoints/usage_endpoints/__init__.py create mode 100644 litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py create mode 100644 litellm/proxy/management_endpoints/usage_endpoints/endpoints.py create mode 100644 tests/test_litellm/proxy/management_endpoints/usage_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py create mode 100644 ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx create mode 100644 ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.tsx diff --git a/litellm/proxy/management_endpoints/usage_endpoints/__init__.py b/litellm/proxy/management_endpoints/usage_endpoints/__init__.py new file mode 100644 index 00000000000..6e68dcd2a2e --- /dev/null +++ b/litellm/proxy/management_endpoints/usage_endpoints/__init__.py @@ -0,0 +1,9 @@ +""" +Usage endpoints package. + +Re-exports the router from endpoints module. +""" + +from litellm.proxy.management_endpoints.usage_endpoints.endpoints import ( # noqa: F401 + router, +) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py new file mode 100644 index 00000000000..f156be7d2cc --- /dev/null +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -0,0 +1,578 @@ +""" +AI Usage Chat - uses LLM tool calling to answer questions about +usage/spend data by querying the aggregated daily activity endpoints. +""" + +import json +from datetime import date +from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL +from litellm.types.proxy.management_endpoints.common_daily_activity import ( + SpendAnalyticsPaginatedResponse, +) + +from typing_extensions import TypedDict + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +USAGE_AI_TEMPERATURE = 0.2 + +TABLE_DAILY_USER_SPEND = "litellm_dailyuserspend" +TABLE_DAILY_TEAM_SPEND = "litellm_dailyteamspend" +TABLE_DAILY_TAG_SPEND = "litellm_dailytagspend" + +ENTITY_FIELD_USER = "user_id" +ENTITY_FIELD_TEAM = "team_id" +ENTITY_FIELD_TAG = "tag" + +PAGINATED_PAGE_SIZE = 200 +MAX_CHAT_MESSAGES = 20 +TOP_N_MODELS = 15 +TOP_N_PROVIDERS = 10 +TOP_N_KEYS = 10 + +# --------------------------------------------------------------------------- +# Types +# --------------------------------------------------------------------------- + + +class SSEStatusEvent(TypedDict): + type: Literal["status"] + message: str + + +class SSEToolCallEvent(TypedDict, total=False): + type: Literal["tool_call"] + tool_name: str + tool_label: str + arguments: Dict[str, str] + status: Literal["running", "complete", "error"] + error: str + + +class SSEChunkEvent(TypedDict): + type: Literal["chunk"] + content: str + + +class SSEDoneEvent(TypedDict): + type: Literal["done"] + + +class SSEErrorEvent(TypedDict): + type: Literal["error"] + message: str + + +SSEEvent = ( + SSEStatusEvent | SSEToolCallEvent | SSEChunkEvent | SSEDoneEvent | SSEErrorEvent +) + + +class ToolHandler(TypedDict): + fetch: Callable[..., Any] + summarise: Callable[[Dict[str, Any]], str] + label: str + + +# --------------------------------------------------------------------------- +# Tool definitions (OpenAI function-calling schema) +# --------------------------------------------------------------------------- + +_DATE_PARAMS = { + "start_date": {"type": "string", "description": "Start date in YYYY-MM-DD format"}, + "end_date": {"type": "string", "description": "End date in YYYY-MM-DD format"}, +} + +_TOOL_USAGE = { + "type": "function", + "function": { + "name": "get_usage_data", + "description": ( + "Fetch aggregated global usage/spend data. Returns daily spend, " + "token counts, request counts, and breakdowns by model, provider, " + "and API key. Use for overall spend, top models, top providers." + ), + "parameters": { + "type": "object", + "properties": { + **_DATE_PARAMS, + "user_id": { + "type": "string", + "description": "Optional user ID filter. Omit for global view.", + }, + }, + "required": ["start_date", "end_date"], + }, + }, +} + +_TOOL_TEAM = { + "type": "function", + "function": { + "name": "get_team_usage_data", + "description": ( + "Fetch usage/spend data broken down by team. Use for questions " + "like 'which team spends the most' or 'show me team X usage'." + ), + "parameters": { + "type": "object", + "properties": { + **_DATE_PARAMS, + "team_ids": { + "type": "string", + "description": "Optional comma-separated team IDs. Omit for all teams.", + }, + }, + "required": ["start_date", "end_date"], + }, + }, +} + +_TOOL_TAG = { + "type": "function", + "function": { + "name": "get_tag_usage_data", + "description": ( + "Fetch usage/spend data broken down by tag. Tags are labels " + "attached to requests (features, environments, credentials)." + ), + "parameters": { + "type": "object", + "properties": { + **_DATE_PARAMS, + "tags": { + "type": "string", + "description": "Optional comma-separated tag names. Omit for all tags.", + }, + }, + "required": ["start_date", "end_date"], + }, + }, +} + +TOOLS_BASE = [_TOOL_USAGE] +TOOLS_ADMIN = [_TOOL_USAGE, _TOOL_TEAM, _TOOL_TAG] + + +def get_tools_for_role(is_admin: bool) -> List[Dict[str, Any]]: + """Return the tool list appropriate for the user's role.""" + return TOOLS_ADMIN if is_admin else TOOLS_BASE + + +_SYSTEM_PROMPT_BASE = ( + "You are an AI assistant embedded in the LiteLLM Usage dashboard. " + "You help users understand their LLM API spend and usage data.\n\n" + "ALWAYS call the appropriate tool(s) first to fetch data before answering. " + "You may call multiple tools if the question spans different dimensions.\n\n" + "Guidelines:\n" + "- Be concise and specific. Use exact numbers from the data.\n" + "- Format costs as dollar amounts (e.g. $12.34).\n" + "- When comparing entities, show a ranked list.\n" + "- If data is empty or no results found, say so clearly.\n" + "- Do not hallucinate data — only use what the tools return.\n" + "- Today's date will be provided below. Use it to interpret relative dates " + "like 'this week', 'this month', 'last 7 days', etc." +) + +_TOOL_DESCRIPTIONS_ADMIN = ( + "You have access to these tools:\n" + "- `get_usage_data`: Global/user-level usage (spend, models, providers, API keys)\n" + "- `get_team_usage_data`: Team-level usage breakdown\n" + "- `get_tag_usage_data`: Tag-level usage breakdown\n\n" +) + +_TOOL_DESCRIPTIONS_BASE = ( + "You have access to this tool:\n" + "- `get_usage_data`: Your usage data (spend, models, providers, API keys)\n\n" +) + + +def _build_system_prompt(is_admin: bool) -> str: + """Build role-appropriate system prompt with today's date.""" + tool_desc = _TOOL_DESCRIPTIONS_ADMIN if is_admin else _TOOL_DESCRIPTIONS_BASE + return ( + f"{_SYSTEM_PROMPT_BASE}\n\n{tool_desc}" + f"Today's date: {date.today().isoformat()}" + ) + + +# keep a public reference for test assertions +SYSTEM_PROMPT = _SYSTEM_PROMPT_BASE + +# --------------------------------------------------------------------------- +# Data fetchers +# --------------------------------------------------------------------------- + + +def _parse_csv_ids(raw: Optional[str]) -> Optional[List[str]]: + if not raw: + return None + return [t.strip() for t in raw.split(",") if t.strip()] + + +async def _query_activity( + table_name: str, + entity_id_field: str, + entity_id: Optional[Any], + start_date: str, + end_date: str, + *, + use_aggregated: bool = False, +) -> SpendAnalyticsPaginatedResponse: + """Shared helper that calls the daily activity query layer.""" + from litellm.proxy.management_endpoints.common_daily_activity import ( + get_daily_activity, + get_daily_activity_aggregated, + ) + from litellm.proxy.proxy_server import prisma_client + + if use_aggregated: + return await get_daily_activity_aggregated( + prisma_client=prisma_client, + table_name=table_name, + entity_id_field=entity_id_field, + entity_id=entity_id, + entity_metadata_field=None, + start_date=start_date, + end_date=end_date, + model=None, + api_key=None, + ) + return await get_daily_activity( + prisma_client=prisma_client, + table_name=table_name, + entity_id_field=entity_id_field, + entity_id=entity_id, + entity_metadata_field=None, + start_date=start_date, + end_date=end_date, + model=None, + api_key=None, + page=1, + page_size=PAGINATED_PAGE_SIZE, + ) + + +async def _fetch_usage_data( + start_date: str, end_date: str, user_id: Optional[str] = None +) -> Dict[str, Any]: + resp = await _query_activity( + TABLE_DAILY_USER_SPEND, + ENTITY_FIELD_USER, + user_id, + start_date, + end_date, + use_aggregated=True, + ) + return resp.model_dump(mode="json") + + +async def _fetch_team_usage_data( + start_date: str, end_date: str, team_ids: Optional[str] = None +) -> Dict[str, Any]: + resp = await _query_activity( + TABLE_DAILY_TEAM_SPEND, + ENTITY_FIELD_TEAM, + _parse_csv_ids(team_ids), + start_date, + end_date, + ) + return resp.model_dump(mode="json") + + +async def _fetch_tag_usage_data( + start_date: str, end_date: str, tags: Optional[str] = None +) -> Dict[str, Any]: + resp = await _query_activity( + TABLE_DAILY_TAG_SPEND, + ENTITY_FIELD_TAG, + _parse_csv_ids(tags), + start_date, + end_date, + ) + return resp.model_dump(mode="json") + + +# --------------------------------------------------------------------------- +# Summarisers — convert raw JSON to concise text the LLM can reason over +# --------------------------------------------------------------------------- + + +def _accumulate_breakdown( + results: List[Dict[str, Any]], dimension: str, fields: List[str] +) -> Dict[str, Dict[str, float]]: + """Aggregate a single breakdown dimension across days.""" + totals: Dict[str, Dict[str, float]] = {} + for day in results: + for key, entry in day.get("breakdown", {}).get(dimension, {}).items(): + if key not in totals: + totals[key] = {f: 0.0 for f in fields} + m = entry.get("metrics", {}) + for f in fields: + totals[key][f] += m.get(f, 0) + return totals + + +def _ranked_lines( + totals: Dict[str, Dict[str, float]], + fmt: Callable[[str, Dict[str, float]], str], + limit: int, +) -> List[str]: + """Sort by spend descending, format each entry, and truncate.""" + return [ + fmt(name, vals) + for name, vals in sorted(totals.items(), key=lambda x: -x[1].get("spend", 0))[ + :limit + ] + ] + + +def _summarise_usage_data(data: Dict[str, Any]) -> str: + meta = data.get("metadata", {}) + results = data.get("results", []) + + header = ( + f"Total Spend: ${meta.get('total_spend', 0):.4f}\n" + f"Total Requests: {meta.get('total_api_requests', 0)}\n" + f"Successful: {meta.get('total_successful_requests', 0)} | " + f"Failed: {meta.get('total_failed_requests', 0)}\n" + f"Total Tokens: {meta.get('total_tokens', 0)}" + ) + + models = _accumulate_breakdown( + results, "models", ["spend", "api_requests", "total_tokens"] + ) + providers = _accumulate_breakdown(results, "providers", ["spend", "api_requests"]) + + model_lines = _ranked_lines( + models, + lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs, {int(d['total_tokens'])} tokens)", + TOP_N_MODELS, + ) + provider_lines = _ranked_lines( + providers, + lambda n, d: f" - {n}: ${d['spend']:.4f} ({int(d['api_requests'])} reqs)", + TOP_N_PROVIDERS, + ) + + sections = [header, ""] + sections += ["Top Models by Spend:"] + (model_lines or [" (no data)"]) + [""] + sections += ["Top Providers by Spend:"] + (provider_lines or [" (no data)"]) + return "\n".join(sections) + + +def _summarise_entity_data(data: Dict[str, Any], entity_label: str) -> str: + """Summarise team/tag entity usage data.""" + results = data.get("results", []) + if not results: + return f"No {entity_label} usage data found for the given date range." + + totals: Dict[str, Dict[str, Any]] = {} + for day in results: + for eid, entry in day.get("breakdown", {}).get("entities", {}).items(): + if eid not in totals: + alias = entry.get("metadata", {}).get("alias", eid) + totals[eid] = {"alias": alias, "spend": 0.0, "requests": 0, "tokens": 0} + m = entry.get("metrics", {}) + totals[eid]["spend"] += m.get("spend", 0) + totals[eid]["requests"] += m.get("api_requests", 0) + totals[eid]["tokens"] += m.get("total_tokens", 0) + + lines = [f"{entity_label} Usage ({len(totals)} {entity_label.lower()}s):", ""] + for eid, d in sorted(totals.items(), key=lambda x: -x[1]["spend"]): + label = d["alias"] if d["alias"] != eid else eid + lines.append( + f"- {label} (ID: {eid}): ${d['spend']:.4f} | " + f"{int(d['requests'])} reqs | {int(d['tokens'])} tokens" + ) + return "\n".join(lines) + + +# --------------------------------------------------------------------------- +# Tool dispatch registry +# --------------------------------------------------------------------------- + +TOOL_HANDLERS: Dict[str, ToolHandler] = { + "get_usage_data": ToolHandler( + fetch=_fetch_usage_data, + summarise=_summarise_usage_data, + label="global usage data", + ), + "get_team_usage_data": ToolHandler( + fetch=_fetch_team_usage_data, + summarise=lambda data: _summarise_entity_data(data, "Team"), + label="team usage data", + ), + "get_tag_usage_data": ToolHandler( + fetch=_fetch_tag_usage_data, + summarise=lambda data: _summarise_entity_data(data, "Tag"), + label="tag usage data", + ), +} + + +# --------------------------------------------------------------------------- +# SSE streaming +# --------------------------------------------------------------------------- + + +def _sse(event: SSEEvent) -> str: + return f"data: {json.dumps(event)}\n\n" + + +def _resolve_fetch_kwargs( + fn_name: str, + fn_args: Dict[str, str], + user_id: Optional[str], + is_admin: bool, +) -> Dict[str, Any]: + """Build keyword arguments for a tool's fetch function.""" + start_date = fn_args.get("start_date", "") + end_date = fn_args.get("end_date", "") + if not start_date or not end_date: + raise ValueError("Missing required start_date or end_date from tool arguments") + kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date} + if fn_name == "get_usage_data": + if not is_admin: + kwargs["user_id"] = user_id + elif fn_args.get("user_id"): + kwargs["user_id"] = fn_args["user_id"] + elif fn_name == "get_team_usage_data" and fn_args.get("team_ids"): + kwargs["team_ids"] = fn_args["team_ids"] + elif fn_name == "get_tag_usage_data" and fn_args.get("tags"): + kwargs["tags"] = fn_args["tags"] + return kwargs + + +async def _execute_tool_call( + handler: ToolHandler, + fn_name: str, + fn_args: Dict[str, str], + user_id: Optional[str], + is_admin: bool, +) -> str: + """Run a single tool and return the summarised result text.""" + kwargs = _resolve_fetch_kwargs(fn_name, fn_args, user_id, is_admin) + raw_data = await handler["fetch"](**kwargs) + return handler["summarise"](raw_data) + + +async def _process_tool_call( + tc: Any, + chat_messages: List[Dict[str, Any]], + user_id: Optional[str], + is_admin: bool, +) -> AsyncIterator[str]: + """Execute a single tool call, yielding SSE events for status.""" + fn_name = tc.function.name + fn_args = json.loads(tc.function.arguments) + + allowed_names = {t["function"]["name"] for t in get_tools_for_role(is_admin)} + handler = TOOL_HANDLERS.get(fn_name) + + if fn_name not in allowed_names or not handler: + chat_messages.append( + { + "role": "tool", + "tool_call_id": tc.id, + "content": f"Tool not available: {fn_name}", + } + ) + return + + tool_event_base = { + "type": "tool_call", + "tool_name": fn_name, + "tool_label": handler["label"], + "arguments": fn_args, + } + yield _sse({**tool_event_base, "status": "running"}) + + try: + tool_result = await _execute_tool_call( + handler, fn_name, fn_args, user_id, is_admin + ) + yield _sse({**tool_event_base, "status": "complete"}) + except Exception as e: + verbose_proxy_logger.error("Tool %s failed: %s", fn_name, e) + tool_result = f"Error fetching {handler['label']}. Please try again." + yield _sse({**tool_event_base, "status": "error"}) + + chat_messages.append( + {"role": "tool", "tool_call_id": tc.id, "content": tool_result} + ) + + +async def _stream_final_response( + model: str, chat_messages: List[Dict[str, Any]] +) -> AsyncIterator[str]: + """Stream the final LLM response after tool results are appended.""" + yield _sse({"type": "status", "message": "Analyzing results..."}) + + response = await litellm.acompletion( + model=model, + messages=chat_messages, + stream=True, + temperature=USAGE_AI_TEMPERATURE, + ) + async for chunk in response: + delta = chunk.choices[0].delta.content + if delta: + yield _sse({"type": "chunk", "content": delta}) + + +async def stream_usage_ai_chat( + messages: List[Dict[str, str]], + model: Optional[str] = None, + user_id: Optional[str] = None, + is_admin: bool = False, +) -> AsyncIterator[str]: + """Stream SSE events: status → tool_call → chunk → done.""" + resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL + truncated = ( + messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages + ) + chat_messages: List[Dict[str, Any]] = [ + {"role": "system", "content": _build_system_prompt(is_admin)}, + *truncated, + ] + + try: + yield _sse({"type": "status", "message": "Thinking..."}) + tools = get_tools_for_role(is_admin) + response = await litellm.acompletion( + model=resolved_model, + messages=chat_messages, + tools=tools, + temperature=USAGE_AI_TEMPERATURE, + ) + choice = response.choices[0] # type: ignore + + if not choice.message.tool_calls: + if choice.message.content: + yield _sse({"type": "chunk", "content": choice.message.content}) + yield _sse({"type": "done"}) + return + + chat_messages.append(choice.message.model_dump()) + for tc in choice.message.tool_calls: + async for event in _process_tool_call(tc, chat_messages, user_id, is_admin): + yield event + async for event in _stream_final_response(resolved_model, chat_messages): + yield event + yield _sse({"type": "done"}) + + except Exception as e: + verbose_proxy_logger.error("AI usage chat failed: %s", e) + yield _sse( + { + "type": "error", + "message": "An internal error occurred. Please try again.", + } + ) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py new file mode 100644 index 00000000000..0dbe518afb7 --- /dev/null +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -0,0 +1,65 @@ +""" +USAGE AI CHAT ENDPOINTS + +/usage/ai/chat - Stream AI chat responses about usage data +""" + +from typing import List, Literal, Optional + +from fastapi import APIRouter, Depends, Request +from fastapi.responses import StreamingResponse +from pydantic import BaseModel, Field + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router = APIRouter() + + +class ChatMessage(BaseModel): + role: Literal["user", "assistant"] + content: str + + +class UsageAIChatRequest(BaseModel): + messages: List[ChatMessage] = Field( + ..., description="Chat messages (user/assistant history)" + ) + model: Optional[str] = Field(default=None, description="Model to use for AI chat") + + +@router.post( + "/usage/ai/chat", + tags=["Budget & Spend Tracking"], + dependencies=[Depends(user_api_key_auth)], +) +async def usage_ai_chat( + data: UsageAIChatRequest, + request: Request, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + AI chat about usage data. Streams SSE events with the AI response. + The AI agent has access to tools that query aggregated daily activity data. + """ + from litellm.proxy.management_endpoints.common_utils import ( + _user_has_admin_view, + ) + from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( + stream_usage_ai_chat, + ) + + is_admin = _user_has_admin_view(user_api_key_dict) + user_id = user_api_key_dict.user_id + messages = [{"role": m.role, "content": m.content} for m in data.messages] + + return StreamingResponse( + stream_usage_ai_chat( + messages=messages, + model=data.model, + user_id=user_id, + is_admin=is_admin, + ), + media_type="text/event-stream", + headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, + ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 36b6bcc7707..607306f3806 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -392,6 +392,7 @@ from litellm.proxy.management_endpoints.organization_endpoints import ( router as organization_router, ) from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router +from litellm.proxy.management_endpoints.usage_endpoints import router as usage_ai_router from litellm.proxy.management_endpoints.project_endpoints import ( router as project_router, ) @@ -12872,6 +12873,7 @@ app.include_router(caching_router) app.include_router(analytics_router) app.include_router(guardrails_router) app.include_router(policy_router) +app.include_router(usage_ai_router) app.include_router(policy_crud_router) app.include_router(policy_resolve_router) app.include_router(search_tool_management_router) diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/__init__.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py new file mode 100644 index 00000000000..f9303bd13a6 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py @@ -0,0 +1,402 @@ +""" +Tests for AI Usage Chat module. +""" + +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( + TOOL_HANDLERS, + TOOLS_ADMIN, + TOOLS_BASE, + _build_system_prompt, + _summarise_entity_data, + _summarise_usage_data, + stream_usage_ai_chat, +) + + +SAMPLE_AGGREGATED_RESPONSE = { + "results": [ + { + "date": "2025-01-15", + "metrics": { + "spend": 50.25, + "prompt_tokens": 20000, + "completion_tokens": 10000, + "total_tokens": 30000, + "api_requests": 500, + "successful_requests": 480, + "failed_requests": 20, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + "breakdown": { + "models": { + "gpt-4": { + "metrics": { + "spend": 40.0, + "api_requests": 300, + "total_tokens": 25000, + }, + "metadata": {}, + "api_key_breakdown": {}, + }, + }, + "providers": { + "openai": { + "metrics": {"spend": 50.25, "api_requests": 500}, + "metadata": {}, + "api_key_breakdown": {}, + }, + }, + "api_keys": { + "sk-test123": { + "metrics": {"spend": 50.25}, + "metadata": {"key_alias": "Production Key"}, + }, + }, + "model_groups": {}, + "mcp_servers": {}, + "entities": {}, + }, + }, + ], + "metadata": { + "total_spend": 50.25, + "total_api_requests": 500, + "total_successful_requests": 480, + "total_failed_requests": 20, + "total_tokens": 30000, + }, +} + +SAMPLE_TEAM_RESPONSE = { + "results": [ + { + "date": "2025-01-15", + "metrics": {"spend": 100.0, "api_requests": 1000, "total_tokens": 50000}, + "breakdown": { + "entities": { + "team-1": { + "metrics": { + "spend": 60.0, + "api_requests": 600, + "total_tokens": 30000, + }, + "metadata": {"alias": "Engineering"}, + "api_key_breakdown": {}, + }, + "team-2": { + "metrics": { + "spend": 40.0, + "api_requests": 400, + "total_tokens": 20000, + }, + "metadata": {"alias": "Marketing"}, + "api_key_breakdown": {}, + }, + }, + "models": {}, + "providers": {}, + "api_keys": {}, + "model_groups": {}, + "mcp_servers": {}, + }, + }, + ], + "metadata": {"total_spend": 100.0, "total_api_requests": 1000}, +} + + +class TestToolSchemas: + def test_admin_tools_include_all(self): + assert len(TOOLS_ADMIN) == 3 + names = {t["function"]["name"] for t in TOOLS_ADMIN} + assert "get_usage_data" in names + assert "get_team_usage_data" in names + assert "get_tag_usage_data" in names + + def test_base_tools_restricted_to_usage_only(self): + assert len(TOOLS_BASE) == 1 + assert TOOLS_BASE[0]["function"]["name"] == "get_usage_data" + + def test_admin_prompt_mentions_all_tools(self): + prompt = _build_system_prompt(is_admin=True) + assert "get_usage_data" in prompt + assert "get_team_usage_data" in prompt + assert "get_tag_usage_data" in prompt + + def test_non_admin_prompt_only_mentions_usage_tool(self): + prompt = _build_system_prompt(is_admin=False) + assert "get_usage_data" in prompt + assert "get_team_usage_data" not in prompt + assert "get_tag_usage_data" not in prompt + + def test_system_prompt_includes_todays_date(self): + from datetime import date + + prompt = _build_system_prompt(is_admin=True) + assert date.today().isoformat() in prompt + + +class TestSummariseUsageData: + def test_summarise_includes_totals(self): + summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE) + assert "$50.25" in summary + assert "500" in summary + + def test_summarise_includes_models(self): + summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE) + assert "gpt-4" in summary + + def test_summarise_includes_providers(self): + summary = _summarise_usage_data(SAMPLE_AGGREGATED_RESPONSE) + assert "openai" in summary + + def test_summarise_handles_empty_data(self): + empty = {"results": [], "metadata": {}} + summary = _summarise_usage_data(empty) + assert "no data" in summary.lower() + + +class TestSummariseEntityData: + def test_team_summary_includes_teams(self): + summary = _summarise_entity_data(SAMPLE_TEAM_RESPONSE, "Team") + assert "Engineering" in summary + assert "Marketing" in summary + assert "$60.0" in summary + assert "$40.0" in summary + + def test_team_summary_empty(self): + empty = {"results": [], "metadata": {}} + summary = _summarise_entity_data(empty, "Team") + assert "No Team usage data" in summary + + +class TestStreamUsageAiChat: + @pytest.mark.asyncio + async def test_stream_emits_status_events(self): + mock_tool_call = MagicMock() + mock_tool_call.id = "call_123" + mock_tool_call.function.name = "get_usage_data" + mock_tool_call.function.arguments = json.dumps( + { + "start_date": "2025-01-01", + "end_date": "2025-01-31", + } + ) + + mock_first_response = MagicMock() + mock_first_response.choices = [MagicMock()] + mock_first_response.choices[0].message.tool_calls = [mock_tool_call] + mock_first_response.choices[0].message.model_dump.return_value = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": { + "name": "get_usage_data", + "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', + }, + } + ], + } + + async def mock_stream(): + chunk = MagicMock() + chunk.choices = [MagicMock()] + chunk.choices[0].delta.content = "Total spend is $50.25" + yield chunk + + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm, patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_usage_data", + new_callable=AsyncMock, + ) as mock_fetch: + mock_litellm.acompletion = AsyncMock( + side_effect=[ + mock_first_response, + mock_stream(), + ] + ) + mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "What is my total spend?"}], + model="gpt-4o-mini", + user_id="user-123", + is_admin=True, + ): + events.append(json.loads(event.replace("data: ", "").strip())) + + status_events = [e for e in events if e["type"] == "status"] + tool_call_events = [e for e in events if e["type"] == "tool_call"] + chunk_events = [e for e in events if e["type"] == "chunk"] + done_events = [e for e in events if e["type"] == "done"] + + assert len(status_events) >= 1 + assert "Thinking" in status_events[0]["message"] + assert len(tool_call_events) >= 1 + assert tool_call_events[0]["tool_name"] == "get_usage_data" + assert tool_call_events[0]["status"] in ("running", "complete") + assert len(chunk_events) >= 1 + assert len(done_events) == 1 + + @pytest.mark.asyncio + async def test_stream_handles_team_tool(self): + mock_tool_call = MagicMock() + mock_tool_call.id = "call_team" + mock_tool_call.function.name = "get_team_usage_data" + mock_tool_call.function.arguments = json.dumps( + { + "start_date": "2025-01-01", + "end_date": "2025-01-31", + } + ) + + mock_first_response = MagicMock() + mock_first_response.choices = [MagicMock()] + mock_first_response.choices[0].message.tool_calls = [mock_tool_call] + mock_first_response.choices[0].message.model_dump.return_value = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_team", + "type": "function", + "function": { + "name": "get_team_usage_data", + "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', + }, + } + ], + } + + async def mock_stream(): + chunk = MagicMock() + chunk.choices = [MagicMock()] + chunk.choices[0].delta.content = "Engineering is the top team." + yield chunk + + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm, patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_team_usage_data", + new_callable=AsyncMock, + ) as mock_fetch: + mock_litellm.acompletion = AsyncMock( + side_effect=[ + mock_first_response, + mock_stream(), + ] + ) + mock_fetch.return_value = SAMPLE_TEAM_RESPONSE + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "Which team spends the most?"}], + model="gpt-4o-mini", + is_admin=True, + ): + events.append(json.loads(event.replace("data: ", "").strip())) + + chunk_events = [e for e in events if e["type"] == "chunk"] + assert len(chunk_events) >= 1 + assert "Engineering" in chunk_events[0]["content"] + + @pytest.mark.asyncio + async def test_stream_handles_error(self): + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm: + mock_litellm.acompletion = AsyncMock(side_effect=Exception("LLM error")) + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "test"}], + ): + events.append(json.loads(event.replace("data: ", "").strip())) + + error_events = [e for e in events if e["type"] == "error"] + assert len(error_events) == 1 + assert "internal error" in error_events[0]["message"].lower() + + @pytest.mark.asyncio + async def test_non_admin_enforces_user_id(self): + mock_tool_call = MagicMock() + mock_tool_call.id = "call_456" + mock_tool_call.function.name = "get_usage_data" + mock_tool_call.function.arguments = json.dumps( + { + "start_date": "2025-01-01", + "end_date": "2025-01-31", + "user_id": "other-user", + } + ) + + mock_first_response = MagicMock() + mock_first_response.choices = [MagicMock()] + mock_first_response.choices[0].message.tool_calls = [mock_tool_call] + mock_first_response.choices[0].message.model_dump.return_value = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_456", + "type": "function", + "function": { + "name": "get_usage_data", + "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31","user_id":"other-user"}', + }, + } + ], + } + + async def mock_stream(): + chunk = MagicMock() + chunk.choices = [MagicMock()] + chunk.choices[0].delta.content = "Data." + yield chunk + + mock_fetch = AsyncMock(return_value=SAMPLE_AGGREGATED_RESPONSE) + + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm, patch.dict( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.TOOL_HANDLERS", + { + "get_usage_data": { + "fetch": mock_fetch, + "summarise": _summarise_usage_data, + "label": "global usage data", + } + }, + ): + mock_litellm.acompletion = AsyncMock( + side_effect=[ + mock_first_response, + mock_stream(), + ] + ) + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "Show data"}], + model="gpt-4o-mini", + user_id="my-user-id", + is_admin=False, + ): + events.append(event) + + mock_fetch.assert_called_once_with( + start_date="2025-01-01", + end_date="2025-01-31", + user_id="my-user-id", + ) diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 489a39a7ee2..3787f451ad3 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -1757,29 +1757,6 @@ "url": "https://opencollective.com/libvips" } }, - "node_modules/@isaacs/balanced-match": { - "version": "4.0.1", - "resolved": "https://registry.npmjs.org/@isaacs/balanced-match/-/balanced-match-4.0.1.tgz", - "integrity": "sha512-yzMTt9lEb8Gv7zRioUilSglI0c0smZ9k5D65677DLWLtWJaXIS3CqcGyUFByYKlnUj6TkjLVs54fBl6+TiGQDQ==", - "dev": true, - "license": "MIT", - "engines": { - "node": "20 || >=22" - } - }, - "node_modules/@isaacs/brace-expansion": { - "version": "5.0.1", - "resolved": "https://registry.npmjs.org/@isaacs/brace-expansion/-/brace-expansion-5.0.1.tgz", - "integrity": "sha512-WMz71T1JS624nWj2n2fnYAuPovhv7EUhk69R6i9dsVyzxt5eM3bjwvgk9L+APE1TRscGysAVMANkB0jh0LQZrQ==", - "dev": true, - "license": "MIT", - "dependencies": { - "@isaacs/balanced-match": "^4.0.1" - }, - "engines": { - "node": "20 || >=22" - } - }, "node_modules/@istanbuljs/schema": { "version": "0.1.3", "resolved": "https://registry.npmjs.org/@istanbuljs/schema/-/schema-0.1.3.tgz", @@ -3696,32 +3673,6 @@ "typescript": ">=4.8.4 <6.0.0" } }, - "node_modules/@typescript-eslint/typescript-estree/node_modules/brace-expansion": { - "version": "2.0.2", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz", - "integrity": "sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==", - "dev": true, - "license": "MIT", - "dependencies": { - "balanced-match": "^1.0.0" - } - }, - "node_modules/@typescript-eslint/typescript-estree/node_modules/minimatch": { - "version": "9.0.5", - "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-9.0.5.tgz", - "integrity": "sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==", - "dev": true, - "license": "ISC", - "dependencies": { - "brace-expansion": "^2.0.1" - }, - "engines": { - "node": ">=16 || 14 >=14.17" - }, - "funding": { - "url": "https://github.com/sponsors/isaacs" - } - }, "node_modules/@typescript-eslint/utils": { "version": "8.54.0", "resolved": "https://registry.npmjs.org/@typescript-eslint/utils/-/utils-8.54.0.tgz", @@ -4749,11 +4700,14 @@ } }, "node_modules/balanced-match": { - "version": "1.0.2", - "resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-1.0.2.tgz", - "integrity": "sha512-3oSeUO0TMV67hN1AmbXsK4yaqU7tjiHlbxRDZOpH0KW9+CeX4bRAaX0Anxt0tx2MrpRpWwQaPwIlISEJhYU5Pw==", + "version": "4.0.4", + "resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-4.0.4.tgz", + "integrity": "sha512-BLrgEcRTwX2o6gGxGOCNyMvGSp35YofuYzw9h1IMTRmKqttAZZVU67bdb9Pr2vUHA8+j3i2tJfjO6C6+4myGTA==", "dev": true, - "license": "MIT" + "license": "MIT", + "engines": { + "node": "18 || 20 || >=22" + } }, "node_modules/baseline-browser-mapping": { "version": "2.9.19", @@ -4787,14 +4741,16 @@ } }, "node_modules/brace-expansion": { - "version": "1.1.12", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-1.1.12.tgz", - "integrity": "sha512-9T9UjW3r0UW5c1Q7GTwllptXwhvYmEzFhzMfZ9H7FQWt+uZePjZPjBP/W1ZEyZ1twGWom5/56TF4lPcqjnDHcg==", + "version": "5.0.3", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.3.tgz", + "integrity": "sha512-fy6KJm2RawA5RcHkLa1z/ScpBeA762UF9KmZQxwIbDtRJrgLzM10depAiEQ+CXYcoiqW1/m96OAAoke2nE9EeA==", "dev": true, "license": "MIT", "dependencies": { - "balanced-match": "^1.0.0", - "concat-map": "0.0.1" + "balanced-match": "^4.0.2" + }, + "engines": { + "node": "18 || 20 || >=22" } }, "node_modules/braces": { @@ -5149,13 +5105,6 @@ "integrity": "sha512-VRhuHOLoKYOy4UbilLbUzbYg93XLjv2PncJC50EuTWPA3gaja1UjBsUP/D/9/juV3vQFr6XBEzn9KCAHdUvOHw==", "license": "MIT" }, - "node_modules/concat-map": { - "version": "0.0.1", - "resolved": "https://registry.npmjs.org/concat-map/-/concat-map-0.0.1.tgz", - "integrity": "sha512-/Srv4dswyQNBfohGpz9o6Yb3Gz3SrUDqBH5rTuhGR7ahtlbYKnVxw2bCFMRljaA7EXHaXZ8wsHdodFvbkhKmqg==", - "dev": true, - "license": "MIT" - }, "node_modules/copy-to-clipboard": { "version": "3.3.3", "resolved": "https://registry.npmjs.org/copy-to-clipboard/-/copy-to-clipboard-3.3.3.tgz", @@ -6924,22 +6873,6 @@ "node": ">=10.13.0" } }, - "node_modules/glob/node_modules/minimatch": { - "version": "10.1.1", - "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.1.1.tgz", - "integrity": "sha512-enIvLvRAFZYXJzkCYG5RKmPfrFArdLv+R+lbQ53BmIMLIry74bjKzX6iHAm8WYamJkhSSEabrWN5D97XnKObjQ==", - "dev": true, - "license": "BlueOak-1.0.0", - "dependencies": { - "@isaacs/brace-expansion": "^5.0.0" - }, - "engines": { - "node": "20 || >=22" - }, - "funding": { - "url": "https://github.com/sponsors/isaacs" - } - }, "node_modules/globals": { "version": "14.0.0", "resolved": "https://registry.npmjs.org/globals/-/globals-14.0.0.tgz", @@ -9035,16 +8968,19 @@ } }, "node_modules/minimatch": { - "version": "3.1.2", - "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-3.1.2.tgz", - "integrity": "sha512-J7p63hRiAjw1NDEww1W7i37+ByIrOWO5XQQAzZ3VOcL0PNybwpfmV/N05zFAzwQ9USyEcX6t3UO+K5aqBQOIHw==", + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.2.2.tgz", + "integrity": "sha512-+G4CpNBxa5MprY+04MbgOw1v7So6n5JY166pFi9KfYwT78fxScCeSNQSNzp6dpPSW2rONOps6Ocam1wFhCgoVw==", "dev": true, - "license": "ISC", + "license": "BlueOak-1.0.0", "dependencies": { - "brace-expansion": "^1.1.7" + "brace-expansion": "^5.0.2" }, "engines": { - "node": "*" + "node": "18 || 20 || >=22" + }, + "funding": { + "url": "https://github.com/sponsors/isaacs" } }, "node_modules/minimist": { @@ -12004,32 +11940,6 @@ "node": ">=18" } }, - "node_modules/test-exclude/node_modules/brace-expansion": { - "version": "2.0.2", - "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-2.0.2.tgz", - "integrity": "sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==", - "dev": true, - "license": "MIT", - "dependencies": { - "balanced-match": "^1.0.0" - } - }, - "node_modules/test-exclude/node_modules/minimatch": { - "version": "9.0.5", - "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-9.0.5.tgz", - "integrity": "sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==", - "dev": true, - "license": "ISC", - "dependencies": { - "brace-expansion": "^2.0.1" - }, - "engines": { - "node": ">=16 || 14 >=14.17" - }, - "funding": { - "url": "https://github.com/sponsors/isaacs" - } - }, "node_modules/thenify": { "version": "3.3.1", "resolved": "https://registry.npmjs.org/thenify/-/thenify-3.3.1.tgz", @@ -13085,21 +12995,6 @@ "type": "github", "url": "https://github.com/sponsors/wooorm" } - }, - "node_modules/@next/swc-win32-ia32-msvc": { - "version": "14.2.33", - "resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.33.tgz", - "integrity": "sha512-pc9LpGNKhJ0dXQhZ5QMmYxtARwwmWLpeocFmVG5Z0DzWq5Uf0izcI8tLc+qOpqxO1PWqZ5A7J1blrUIKrIFc7Q==", - "cpu": [ - "ia32" - ], - "optional": true, - "os": [ - "win32" - ], - "engines": { - "node": ">= 10" - } } } } diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx new file mode 100644 index 00000000000..85ce11e605a --- /dev/null +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx @@ -0,0 +1,85 @@ +import { screen } from "@testing-library/react"; +import { beforeAll, describe, expect, it, vi } from "vitest"; +import { renderWithProviders } from "../../../../tests/test-utils"; +import UsageAIChatPanel from "./UsageAIChatPanel"; + +beforeAll(() => { + if (typeof window !== "undefined" && !window.ResizeObserver) { + window.ResizeObserver = class ResizeObserver { + observe() {} + unobserve() {} + disconnect() {} + } as any; + } +}); + +vi.mock("../../networking", () => ({ + modelHubCall: vi.fn().mockResolvedValue({ + data: [ + { model_group: "gpt-4" }, + { model_group: "claude-3-opus" }, + ], + }), + usageAiChatStream: vi.fn(), +})); + +const defaultProps = { + open: true, + onClose: vi.fn(), + accessToken: "test-token", +}; + +describe("UsageAIChatPanel", () => { + it("should render the panel when open", () => { + renderWithProviders(); + + expect(screen.getByText("Ask AI")).toBeInTheDocument(); + expect( + screen.getByText("Ask about your spend, models, keys, and trends") + ).toBeInTheDocument(); + }); + + it("should render model selector", () => { + renderWithProviders(); + + expect(screen.getByText("Select a model (optional, defaults to gpt-4o-mini)")).toBeInTheDocument(); + }); + + it("should render empty state message when no conversation", () => { + renderWithProviders(); + + expect(screen.getByText("Ask a question about your usage")).toBeInTheDocument(); + }); + + it("should render the send button", () => { + renderWithProviders(); + + expect(screen.getByText("Send")).toBeInTheDocument(); + }); + + it("should render input placeholder", () => { + renderWithProviders(); + + expect(screen.getByPlaceholderText("Ask about your usage...")).toBeInTheDocument(); + }); + + it("should render clear chat button", () => { + renderWithProviders(); + + expect(screen.getByText("Clear chat")).toBeInTheDocument(); + }); + + it("should have the panel element even when closed (just off-screen)", () => { + renderWithProviders(); + + expect(screen.getByTestId("usage-ai-chat-panel")).toBeInTheDocument(); + expect(screen.getByTestId("usage-ai-chat-panel")).toHaveClass("translate-x-full"); + }); + + it("should not have translate-x-full class when open", () => { + renderWithProviders(); + + expect(screen.getByTestId("usage-ai-chat-panel")).not.toHaveClass("translate-x-full"); + expect(screen.getByTestId("usage-ai-chat-panel")).toHaveClass("translate-x-0"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.tsx new file mode 100644 index 00000000000..85f46bfa346 --- /dev/null +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.tsx @@ -0,0 +1,402 @@ +import React, { useEffect, useRef, useState } from "react"; +import { Button, Select, Input, Spin } from "antd"; +import ReactMarkdown from "react-markdown"; +import { modelHubCall, usageAiChatStream, UsageAiToolCallEvent } from "../../networking"; + +const { TextArea } = Input; + +interface ToolCallStep { + tool_name: string; + tool_label: string; + arguments: Record; + status: "running" | "complete" | "error"; + error?: string; +} + +interface ChatMessage { + role: "user" | "assistant"; + content: string; + toolCalls?: ToolCallStep[]; +} + +interface UsageAIChatPanelProps { + open: boolean; + onClose: () => void; + accessToken: string | null; +} + +const TOOL_ICONS: Record = { + get_usage_data: "📊", + get_team_usage_data: "👥", + get_tag_usage_data: "🏷️", +}; + +const ToolCallDisplay: React.FC<{ step: ToolCallStep }> = ({ step }) => { + const icon = TOOL_ICONS[step.tool_name] || "🔧"; + const args = step.arguments; + const dateRange = args.start_date && args.end_date + ? `${args.start_date} → ${args.end_date}` + : ""; + const filter = args.team_ids || args.tags || args.user_id || ""; + + return ( +
+ + {step.status === "running" ? ( + + ) : step.status === "error" ? ( + ✗ + ) : ( + ✓ + )} + +
+
+ {icon} {step.tool_label} +
+ {dateRange && ( +
{dateRange}
+ )} + {filter && ( +
Filter: {filter}
+ )} + {step.status === "error" && step.error && ( +
{step.error}
+ )} +
+
+ ); +}; + +const MarkdownContent: React.FC<{ content: string }> = ({ content }) => ( +

{children}

, + strong: ({ children }) => {children}, + ul: ({ children }) =>
    {children}
, + ol: ({ children }) =>
    {children}
, + li: ({ children }) =>
  • {children}
  • , + h1: ({ children }) =>

    {children}

    , + h2: ({ children }) =>

    {children}

    , + h3: ({ children }) =>

    {children}

    , + code: ({ children, className }) => { + const isBlock = className?.includes("language-"); + return isBlock ? ( +
    +            {children}
    +          
    + ) : ( + {children} + ); + }, + table: ({ children }) => ( +
    + {children}
    +
    + ), + th: ({ children }) => {children}, + td: ({ children }) => {children}, + }} + > + {content} +
    +); + +const UsageAIChatPanel: React.FC = ({ + open, + onClose, + accessToken, +}) => { + const [messages, setMessages] = useState([]); + const [inputText, setInputText] = useState(""); + const [isLoading, setIsLoading] = useState(false); + const [selectedModel, setSelectedModel] = useState(undefined); + const [availableModels, setAvailableModels] = useState([]); + const [isLoadingModels, setIsLoadingModels] = useState(false); + const [streamingContent, setStreamingContent] = useState(""); + const [statusMessage, setStatusMessage] = useState(null); + const [activeToolCalls, setActiveToolCalls] = useState([]); + const messagesEndRef = useRef(null); + const abortControllerRef = useRef(null); + + useEffect(() => { + if (open && availableModels.length === 0) { + loadModels(); + } + }, [open]); + + useEffect(() => { + if (typeof messagesEndRef.current?.scrollIntoView === "function") { + messagesEndRef.current.scrollIntoView({ behavior: "smooth" }); + } + }, [messages, streamingContent, activeToolCalls, statusMessage]); + + const loadModels = async () => { + if (!accessToken) return; + setIsLoadingModels(true); + try { + const fetchedModels = await modelHubCall(accessToken); + if (fetchedModels?.data?.length > 0) { + const models = fetchedModels.data + .map((item: any) => item.model_group as string) + .sort(); + setAvailableModels(models); + } + } catch (error) { + console.error("Failed to load models:", error); + } finally { + setIsLoadingModels(false); + } + }; + + const handleSend = async () => { + if (!accessToken || !inputText.trim() || isLoading) return; + + const userMessage: ChatMessage = { role: "user", content: inputText.trim() }; + const updatedMessages = [...messages, userMessage]; + setMessages(updatedMessages); + setInputText(""); + setIsLoading(true); + setStreamingContent(""); + setStatusMessage(null); + setActiveToolCalls([]); + + const abortController = new AbortController(); + abortControllerRef.current = abortController; + + let accumulated = ""; + const toolCalls: ToolCallStep[] = []; + + try { + await usageAiChatStream( + accessToken, + updatedMessages.slice(-20).map((m) => ({ role: m.role, content: m.content })), + selectedModel || "", + (content: string) => { + setStatusMessage(null); + accumulated += content; + setStreamingContent(accumulated); + }, + () => { + setStatusMessage(null); + setActiveToolCalls([]); + setMessages((prev) => [ + ...prev, + { role: "assistant", content: accumulated, toolCalls: toolCalls.length > 0 ? [...toolCalls] : undefined }, + ]); + setStreamingContent(""); + }, + (errorMsg: string) => { + setStatusMessage(null); + setActiveToolCalls([]); + setMessages((prev) => [ + ...prev, + { role: "assistant", content: `Error: ${errorMsg}` }, + ]); + setStreamingContent(""); + }, + (status: string) => { + setStatusMessage(status); + }, + (event: UsageAiToolCallEvent) => { + const idx = toolCalls.findIndex((tc) => tc.tool_name === event.tool_name); + if (idx >= 0) { + toolCalls[idx] = { ...event }; + } else { + toolCalls.push({ ...event }); + } + setActiveToolCalls([...toolCalls]); + }, + abortController.signal, + ); + } catch (error: any) { + if (error?.name === "AbortError" || abortController.signal.aborted) { + return; + } + const errorMsg = error?.message || "Failed to get response. Please try again."; + setMessages((prev) => [ + ...prev, + { role: "assistant", content: `Error: ${errorMsg}` }, + ]); + setStreamingContent(""); + } finally { + setIsLoading(false); + abortControllerRef.current = null; + } + }; + + const handleKeyDown = (e: React.KeyboardEvent) => { + if (e.key === "Enter" && !e.shiftKey) { + e.preventDefault(); + handleSend(); + } + }; + + const handleClose = () => { + if (abortControllerRef.current) { + abortControllerRef.current.abort(); + } + onClose(); + }; + + const handleClear = () => { + setMessages([]); + setStreamingContent(""); + setActiveToolCalls([]); + setStatusMessage(null); + }; + + return ( +
    + {/* Header */} +
    +
    +
    + + + +

    Ask AI

    +
    + +
    +

    + Ask about your spend, models, keys, and trends +

    +
    + + {/* Model selector */} +
    +