From bae58eb4e0aa015d5085264e8b0d9a342f163ce5 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Fri, 31 Jul 2026 13:03:13 -0400 Subject: [PATCH 001/139] fix(anthropic): preserve mid-turn system messages Generated with AI Co-Authored-By: Claude Code --- .../chat/guardrail_translation/handler.py | 209 ++++-- .../adapters/transformation.py | 48 +- .../responses_adapters/transformation.py | 49 +- litellm/types/guardrails.py | 47 +- litellm/types/llms/anthropic.py | 12 + .../test_anthropic_guardrail_handler.py | 699 ++++++++++++++++++ ...al_pass_through_adapters_transformation.py | 218 ++++++ .../context_management/test_compact.py | 54 ++ .../test_responses_adapters_transformation.py | 138 ++++ 9 files changed, 1380 insertions(+), 94 deletions(-) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index a549db94224..82707e741c0 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -13,6 +13,7 @@ Pattern Overview: """ import json +from copy import deepcopy from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast from litellm._logging import verbose_proxy_logger @@ -24,7 +25,6 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTra from litellm.llms.base_llm.guardrail_translation.utils import ( effective_skip_system_message_for_guardrail, effective_skip_tool_message_for_guardrail, - openai_messages_without_system, openai_messages_without_tool, ) from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( @@ -59,14 +59,10 @@ if TYPE_CHECKING: class AnthropicMessagesHandler(BaseTranslation): - """ - Handler for processing Anthropic messages with guardrails. + """Process Anthropic messages with guardrails. - This class provides methods to: - 1. Process input messages (pre-call hook) - 2. Process output responses (post-call hook) - - Methods can be overridden to customize behavior for different message formats. + In-sequence system entries are untrusted client input. This handler scans and preserves + them through guardrail rewrites; downstream provider handling is out of scope. """ def __init__(self): @@ -279,14 +275,26 @@ class AnthropicMessagesHandler(BaseTranslation): skip_system = effective_skip_system_message_for_guardrail(guardrail_to_apply) skip_tool = effective_skip_tool_message_for_guardrail(guardrail_to_apply) - chat_completion_compatible_request = self._translate_to_openai(data) + # Exclude only the trusted top-level prompt. In-sequence system entries are untrusted + # and must stay aligned with texts_to_check for positional masking. When the top-level + # prompt is included, the pre-existing count mismatch disables positional masking. + translation_source = { # mutable-ok: API message payload + key: value for key, value in data.items() if key != "system" + } # mutable-ok: API message payload + chat_completion_compatible_request = self._translate_to_openai(translation_source) structured_messages = cast( List[AllMessageValues], chat_completion_compatible_request.get("messages", []), ) - if skip_system: - structured_messages = openai_messages_without_system(structured_messages) + has_midturn_system_message = any( + str(message.get("role") or "").lower() == "system" for message in structured_messages + ) + hoisted_system_message: AllMessageValues | None = None + if not skip_system: + hoisted_system_message = self._hoisted_top_level_system_message(data) + if hoisted_system_message is not None: + structured_messages.insert(0, hoisted_system_message) if skip_tool: structured_messages = openai_messages_without_tool(structured_messages) @@ -346,7 +354,12 @@ class AnthropicMessagesHandler(BaseTranslation): guardrailed_structured_messages is not None and guardrailed_structured_messages is not original_structured_messages ): - self._write_back_structured_messages(data, guardrailed_structured_messages) + self._write_back_structured_messages( + data, + guardrailed_structured_messages, + hoisted_system_message=hoisted_system_message, + preserve_system_messages=has_midturn_system_message, + ) else: # Step 3: Map guardrail responses back to original message structure await self._apply_guardrail_responses_to_input( @@ -359,36 +372,120 @@ class AnthropicMessagesHandler(BaseTranslation): return data - @staticmethod - def _write_back_structured_messages(data: dict, structured_messages: list) -> None: - """Convert compressed structured_messages back to Anthropic format and write to data. + def _hoisted_top_level_system_message( + self, data: dict + ) -> AllMessageValues | None: # mutable-ok: API message payload + """Return the system message produced by translating the top-level prompt.""" + system = data.get("system") + if not system: + return None + probe = self._translate_to_openai( + { # mutable-ok: API message payload + "model": data.get("model") or "", + "messages": [], # mutable-ok: API message payload + "system": system, + } + ) + hoisted = probe.get("messages") or [] # mutable-ok: API message payload + return hoisted[0] if hoisted else None - ``anthropic_messages_pt`` merges every run of consecutive user/tool rows - into a single message, so a turn carrying only tool results and the user - turn that follows it come back fused, and the request the model sees no - longer has the boundaries the client sent. Converting a row at a time - would keep them apart but breaks tool pairing: an assistant row whose - tool results sit outside its own call reads as an orphaned tool call, - and under ``modify_params`` the sanitizer answers it with a synthetic - "tool execution skipped" result and drops the real one. Converting each - assistant row together with the tool rows that answer it, and every - other row on its own, satisfies both. - """ + @staticmethod + def _openai_system_message_to_anthropic( + message: dict[str, Any], + ) -> dict[str, Any] | None: # mutable-ok: API message payload + """Convert an OpenAI system message to the client's Anthropic-shaped entry.""" + content = message.get("content") + if isinstance(content, str): + return ( + {"role": "system", "content": content} if content else None # mutable-ok: API message payload + ) # mutable-ok: API message payload + if not isinstance(content, list): + return None + blocks: list[dict[str, Any]] = [] # mutable-ok: API message payload + for block in content: + if not isinstance(block, dict) or block.get("type") != "text": + continue + text = block.get("text") + if not isinstance(text, str) or not text: + continue + anthropic_block: dict[str, Any] = { # mutable-ok: API message payload + "type": "text", + "text": text, + } # mutable-ok: API message payload + cache_control = block.get("cache_control") + if cache_control: + anthropic_block["cache_control"] = deepcopy(cache_control) + blocks.append(anthropic_block) + return ( + {"role": "system", "content": blocks} if blocks else None # mutable-ok: API message payload + ) # mutable-ok: API message payload + + @staticmethod + def _is_hoisted_top_level_system(message: Any, hoisted_system_message: Any) -> bool: + """Match the hoisted prompt by identity, or by value after serialization.""" + if hoisted_system_message is None: + return False + if message is hoisted_system_message: + return True + return ( + isinstance(message, dict) and isinstance(hoisted_system_message, dict) and message == hoisted_system_message + ) + + @staticmethod + def _write_back_structured_messages( + data: dict, # mutable-ok: API message payload + structured_messages: list, # mutable-ok: API message payload + hoisted_system_message: Any = None, + preserve_system_messages: bool = False, + ) -> None: + """Write a guardrail's structured-message rewrite back without losing corrections.""" from litellm.litellm_core_utils.prompt_templates.factory import ( anthropic_messages_pt, group_tool_exchanges, ) + def _is_system(message: Any) -> bool: + return isinstance(message, dict) and str(message.get("role") or "").lower() == "system" + model = str(data.get("model") or "") - non_system = [m for m in structured_messages if m.get("role") != "system"] - groups = tuple([non_system[index] for index in group] for group in group_tool_exchanges(non_system)) or ( - non_system, - ) - converted = [ - message - for group in groups - for message in anthropic_messages_pt(messages=group, model=model, llm_provider="anthropic") - ] + converted: list = [] # mutable-ok: API message payload + + def _convert_run(run: list) -> None: # mutable-ok: API message payload + for group in group_tool_exchanges(run): + converted.extend( + anthropic_messages_pt( + messages=[ # mutable-ok: API message payload + run[index] for index in group + ], # mutable-ok: API message payload + model=model, + llm_provider="anthropic", + ) + ) + + run: list = [] # mutable-ok: API message payload + hoisted_dropped = False + for message in structured_messages: + if not _is_system(message): + run.append(message) + continue + _convert_run(run) + run = [] # mutable-ok: API message payload + if not hoisted_dropped and AnthropicMessagesHandler._is_hoisted_top_level_system( + message, hoisted_system_message + ): + hoisted_dropped = True + continue + if preserve_system_messages: + anthropic_system = AnthropicMessagesHandler._openai_system_message_to_anthropic(message) + if anthropic_system is not None: + converted.append(anthropic_system) + _convert_run(run) + if not any(not _is_system(message) for message in converted): + converted.extend( + anthropic_messages_pt( + messages=[], model=model, llm_provider="anthropic" + ) # mutable-ok: API message payload + ) # mutable-ok: API message payload for msg in converted: content = msg.get("content") if isinstance(content, list): @@ -397,6 +494,29 @@ class AnthropicMessagesHandler(BaseTranslation): block.pop("cache_control", None) data["messages"] = converted + @staticmethod + def _extract_midturn_system_text( + message: dict[str, Any], # mutable-ok: API message payload + msg_idx: int, + texts_to_check: list[str], # mutable-ok: API message payload + task_mappings: list[tuple[int, int | None]], # mutable-ok: API message payload + ) -> None: + content = message.get("content") + if isinstance(content, str): + if content: + texts_to_check.append(content) + task_mappings.append((msg_idx, None)) + return + if not isinstance(content, list): + return + for content_idx, content_item in enumerate(content): + if not isinstance(content_item, dict) or content_item.get("type") != "text": + continue + text_str = content_item.get("text") + if isinstance(text_str, str) and text_str: + texts_to_check.append(text_str) + task_mappings.append((msg_idx, content_idx)) + def extract_request_tool_names(self, data: dict) -> List[str]: """Extract tool names from Anthropic messages request (tools[].name).""" names: List[str] = [] @@ -415,15 +535,18 @@ class AnthropicMessagesHandler(BaseTranslation): skip_system_message: bool = False, skip_tool_message: bool = False, ) -> None: - """ - Extract text content and images from a message. - - Override this method to customize text/image extraction logic. - """ - role = str(message.get("role") or "").lower() - if skip_system_message and role == "system": + """Extract text content and images from a message.""" + role = str(message.get("role") or "") + if role == "system": + # Match the adapter's filtering so positional guardrail write-back stays aligned. + self._extract_midturn_system_text( + message=message, + msg_idx=msg_idx, + texts_to_check=texts_to_check, + task_mappings=task_mappings, + ) return - if skip_tool_message and role == "tool": + if skip_tool_message and role.lower() == "tool": return content = message.get("content", None) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 86c9c1db481..707d53e9006 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -85,12 +85,12 @@ from litellm.llms.anthropic.experimental_pass_through.context_management import ) from litellm.types.llms.anthropic import ( ANTHROPIC_HOSTED_TOOLS, + AllAnthropicPassThroughMessageValues, AllAnthropicToolsValues, - AnthopicMessagesAssistantMessageParam, AnthropicFinishReason, AnthropicMessagesRequest, + AnthropicMessagesSystemMessageParam, AnthropicMessagesToolChoice, - AnthropicMessagesUserMessageParam, AnthropicResponseContentBlockRedactedThinking, AnthropicResponseContentBlockText, AnthropicResponseContentBlockThinking, @@ -354,12 +354,7 @@ class LiteLLMAnthropicMessagesAdapter: def translate_anthropic_messages_to_openai( self, - messages: List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], + messages: List[AllAnthropicPassThroughMessageValues], # mutable-ok: API message payload model: Optional[str] = None, ) -> List: new_messages: List[AllMessageValues] = [] @@ -367,6 +362,11 @@ class LiteLLMAnthropicMessagesAdapter: user_message: Optional[ChatCompletionUserMessage] = None tool_message_list: List[ChatCompletionToolMessage] = [] new_user_content_list: List[Union[ChatCompletionTextObject, ChatCompletionImageObject]] = [] + if m["role"] == "system": + system_message = self._translate_midturn_system_message_to_openai(m, model) + if system_message is not None: + new_messages.append(system_message) + continue ## USER MESSAGE ## if m["role"] == "user": ## translate user message @@ -867,6 +867,29 @@ class LiteLLMAnthropicMessagesAdapter: for def_schema in schema[key].values(): LiteLLMAnthropicMessagesAdapter._add_additional_properties_false(def_schema) + def _translate_midturn_system_message_to_openai( + self, + message: AnthropicMessagesSystemMessageParam, + model: str | None, + ) -> ChatCompletionSystemMessage | None: + """Translate an in-sequence system entry without changing its role or position.""" + content = message.get("content") + if isinstance(content, str): + return ChatCompletionSystemMessage(role="system", content=content) if content else None + if not isinstance(content, list): + return None + text_parts: list[ChatCompletionTextObject] = [] # mutable-ok: API message payload + for block in content: + if not isinstance(block, dict) or block.get("type") != "text": + continue + text = block.get("text") + if not text: + continue + text_obj = ChatCompletionTextObject(type="text", text=text) + self._add_cache_control_if_applicable(block, text_obj, model) + text_parts.append(text_obj) + return ChatCompletionSystemMessage(role="system", content=text_parts) if text_parts else None + def _add_system_message_to_messages( self, new_messages: List[AllMessageValues], @@ -1068,13 +1091,8 @@ class LiteLLMAnthropicMessagesAdapter: tool_name_mapping: Dict[str, str] = {} ## CONVERT ANTHROPIC MESSAGES TO OPENAI - messages_list: List[Union[AnthropicMessagesUserMessageParam, AnthopicMessagesAssistantMessageParam]] = cast( - List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], + messages_list = cast( + List[AllAnthropicPassThroughMessageValues], anthropic_message_request["messages"], ) new_messages = self.translate_anthropic_messages_to_openai( diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 172e54de98e..cbe36100eb0 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -6,6 +6,7 @@ path used for OpenAI and Azure models. """ import json +from collections.abc import Iterable from typing import Any, Dict, List, Optional, Union, cast from litellm.litellm_core_utils.reasoning_effort_utils import ( @@ -15,15 +16,15 @@ from litellm.llms.anthropic.experimental_pass_through.utils import ( is_reasoning_auto_summary_enabled, ) from litellm.types.llms.anthropic import ( + AllAnthropicPassThroughMessageValues, AllAnthropicToolsValues, - AnthopicMessagesAssistantMessageParam, AnthropicFinishReason, AnthropicMessagesRequest, AnthropicMessagesToolChoice, - AnthropicMessagesUserMessageParam, AnthropicResponseContentBlockText, AnthropicResponseContentBlockThinking, AnthropicResponseContentBlockToolUse, + AnthropicSystemMessageContent, ) from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, @@ -54,19 +55,32 @@ class LiteLLMAnthropicToResponsesAPIAdapter: return source.get("url") return None + @staticmethod + def _translate_midturn_system_content_to_responses( + content: Union[str, Iterable[AnthropicSystemMessageContent]], + ) -> list[dict[str, str]]: # mutable-ok: API message payload + """Convert in-sequence system content to Responses input-text parts.""" + if isinstance(content, str): + return ( + [{"type": "input_text", "text": content}] if content else [] # mutable-ok: API message payload + ) # mutable-ok: API message payload + if not isinstance(content, list): + return [] # mutable-ok: API message payload + return [ # mutable-ok: API message payload + {"type": "input_text", "text": text} # mutable-ok: API message payload + for block in content + if isinstance(block, dict) and block.get("type") == "text" and (text := block.get("text")) + ] + def translate_messages_to_responses_input( self, - messages: List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], + messages: List[AllAnthropicPassThroughMessageValues], # mutable-ok: API message payload ) -> List[Dict[str, Any]]: """ Convert Anthropic messages list to Responses API `input` items. Mapping: + system text -> message(role=system, input_text) user text -> message(role=user, input_text) user image -> message(role=user, input_image) user tool_result -> function_call_output @@ -76,6 +90,18 @@ class LiteLLMAnthropicToResponsesAPIAdapter: input_items: List[Dict[str, Any]] = [] for m in messages: + if m["role"] == "system": + system_parts = self._translate_midturn_system_content_to_responses(m.get("content")) + if system_parts: + input_items.append( + { # mutable-ok: API message payload + "type": "message", + "role": "system", + "content": system_parts, + } + ) + continue + role = m["role"] content = m.get("content") @@ -287,12 +313,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: """ model: str = anthropic_request["model"] messages_list = cast( - List[ - Union[ - AnthropicMessagesUserMessageParam, - AnthopicMessagesAssistantMessageParam, - ] - ], + List[AllAnthropicPassThroughMessageValues], anthropic_request["messages"], ) diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index af419d8cb6f..2e0da24ccda 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -11,12 +11,24 @@ from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( from litellm.types.proxy.guardrails.guardrail_hooks.block_code_execution import ( BlockCodeExecutionGuardrailConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import ( + CiscoAIDefenseGuardrailConfigModel, +) +from litellm.types.proxy.guardrails.guardrail_hooks.compresr import ( + CompresrGuardrailConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import ( EnkryptAIGuardrailConfigs, ) from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import ( GraySwanGuardrailConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.headroom import ( + HeadroomGuardrailConfigModel, +) +from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( + HiddenlayerGuardrailConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( IBMGuardrailsBaseConfigModel, ) @@ -29,38 +41,26 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( from litellm.types.proxy.guardrails.guardrail_hooks.promptguard import ( PromptGuardConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.xecguard import ( - XecGuardConfigModel, +from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, ) from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import ( QualifireGuardrailConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( - ToolPermissionGuardrailConfigModel, -) -from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( - HiddenlayerGuardrailConfigModel, -) -from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( - QostodianNexusConfigModel, -) from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import ( RepelloAIGuardrailConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import ( - VigilGuardGuardrailConfigModel, -) -from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import ( - CiscoAIDefenseGuardrailConfigModel, -) from litellm.types.proxy.guardrails.guardrail_hooks.singulr import ( SingulrGuardrailConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.headroom import ( - HeadroomGuardrailConfigModel, +from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( + ToolPermissionGuardrailConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.compresr import ( - CompresrGuardrailConfigModel, +from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import ( + VigilGuardGuardrailConfigModel, +) +from litellm.types.proxy.guardrails.guardrail_hooks.xecguard import ( + XecGuardConfigModel, ) """ @@ -743,7 +743,10 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up "When True, unified guardrails skip system-role messages when building " "evaluation inputs (texts and structured_messages). When False, system " "messages are included even if litellm_settings sets a global skip. When " - "None, use the global litellm.skip_system_message_in_guardrail setting." + "None, use the global litellm.skip_system_message_in_guardrail setting. " + "For Anthropic /v1/messages, the flag applies only to the trusted top-level " + "system prompt. In-sequence system entries are untrusted client input and remain " + "in texts and structured_messages." ), ) diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index c24d072217a..29faf500b3d 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -365,8 +365,20 @@ class AnthropicSystemMessageContent(TypedDict, total=False): cache_control: Optional[Union[dict, ChatCompletionCachedContent]] +class AnthropicMessagesSystemMessageParam(TypedDict, total=False): + role: Required[Literal["system"]] + content: Required[Union[str, Iterable[AnthropicSystemMessageContent]]] + + AllAnthropicMessageValues = Union[AnthropicMessagesUserMessageParam, AnthopicMessagesAssistantMessageParam] +# System is not a native Anthropic message role; only pass-through adapters use this union. +AllAnthropicPassThroughMessageValues = Union[ + AnthropicMessagesUserMessageParam, + AnthopicMessagesAssistantMessageParam, + AnthropicMessagesSystemMessageParam, +] + class AnthropicMessagesRequestOptionalParams(TypedDict, total=False): max_tokens: Optional[int] diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 48acdd348e9..7757b0fa5a4 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -75,6 +75,51 @@ class MockRecordingGuardrail(CustomGuardrail): return inputs +class MockMaskingGuardrail(CustomGuardrail): + """Capture request inputs and mask one known prohibited value.""" + + def __init__(self, skip_system_message_in_guardrail: Optional[bool] = True): + super().__init__(guardrail_name="masking-test") + self.skip_system_message_in_guardrail = skip_system_message_in_guardrail + self.inputs: Optional[GenericGuardrailAPIInputs] = None + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + self.inputs = inputs.copy() + masked_inputs = inputs.copy() + masked_inputs["texts"] = [ + "[MASKED]" if text == "prohibited correction" else text for text in inputs.get("texts", []) + ] + return masked_inputs + + +class MockCompactingGuardrail(CustomGuardrail): + """Stand in for a compaction guardrail that rewrites `structured_messages` wholesale.""" + + def __init__(self, replacement_messages: list): + super().__init__(guardrail_name="compacting-test") + self.replacement_messages = replacement_messages + self.inputs: Optional[GenericGuardrailAPIInputs] = None + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + self.inputs = inputs.copy() + rewritten = inputs.copy() + # A new list object -- this is what signals a rewrite to the handler. + rewritten["structured_messages"] = list(self.replacement_messages) + return rewritten + + class TestAnthropicMessagesHandlerStreamingRequestData: """Post-call guardrails on streaming /v1/messages receive the response and identity metadata""" @@ -210,6 +255,660 @@ class TestAnthropicMessagesHandlerInputProcessing: assert data.get("litellm_metadata", {}).get("guardrails") assert guardrail.dynamic_params == {"policy_id": "policy-123"} + @pytest.mark.asyncio + async def test_midturn_system_correction_is_guardrailed_when_top_level_system_is_skipped( + self, + ): + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail() + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "safe text"}, + { + "role": "system", + "content": [ + {"type": "unsupported", "text": "discarded text"}, + {"type": "text", "text": "prohibited correction"}, + ], + }, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"] + assert "trusted top-level system prompt" not in guardrail.inputs["texts"] + assert data["messages"][1]["content"][0]["text"] == "discarded text" + assert data["messages"][1]["content"][1]["text"] == "[MASKED]" + + @pytest.mark.asyncio + async def test_string_midturn_system_correction_is_guardrailed(self): + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail() + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [{"role": "system", "content": "prohibited correction"}], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + assert guardrail.inputs["texts"] == ["prohibited correction"] + assert data["messages"][0]["content"] == "[MASKED]" + + @pytest.mark.asyncio + async def test_unsupported_midturn_system_content_is_not_guardrailed(self): + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail() + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + { + "role": "system", + "content": [{"type": "image", "source": {"type": "url"}}], + } + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is None + + @pytest.mark.asyncio + async def test_skip_system_message_excludes_only_hoisted_top_level_system(self): + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail() + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "safe text"}, + {"role": "system", "content": "prohibited correction"}, + {"role": "user", "content": "continue"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + structured = guardrail.inputs["structured_messages"] + assert [m["role"] for m in structured] == ["user", "system", "user"] + assert structured[1]["content"] == "prohibited correction" + + @pytest.mark.asyncio + async def test_default_skip_false_scans_midturn_system_and_hoists_top_level_system( + self, + ): + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail(skip_system_message_in_guardrail=None) + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "safe text"}, + {"role": "system", "content": "prohibited correction"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"] + structured = guardrail.inputs["structured_messages"] + assert [m["role"] for m in structured] == ["system", "user", "system"] + assert structured[0]["content"] == "trusted top-level system prompt" + assert data["messages"][1]["content"] == "[MASKED]" + + @pytest.mark.asyncio + async def test_bedrock_masking_slice_is_unavailable_when_top_level_system_is_included( + self, + ): + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( + BedrockGuardrail, + ) + + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail(skip_system_message_in_guardrail=None) + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "safe text"}, + {"role": "system", "content": "prohibited correction"}, + {"role": "user", "content": "latest question"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + texts = guardrail.inputs["texts"] + structured = guardrail.inputs["structured_messages"] + + bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1") + assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + 1 + latest_user_index = bedrock._find_latest_message_index(structured, target_role="user") + assert ( + bedrock._locate_message_texts_slice( + structured_messages=structured, + target_index=latest_user_index, + texts=texts, + ) + is None + ) + assert ( + bedrock._merge_masked_texts( + masked_texts=["{MASKED}"], + texts=texts, + scanned_slice=None, + scanned_role_subset=True, + ) + == texts + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize("skip_system_message_in_guardrail", [True, None]) + async def test_midturn_system_text_extraction_matches_translation_in_both_skip_modes( + self, + skip_system_message_in_guardrail: Optional[bool], + ): + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( + BedrockGuardrail, + ) + + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail(skip_system_message_in_guardrail=skip_system_message_in_guardrail) + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "safe text"}, + { + "role": "system", + "content": [ + {"type": "text", "text": ""}, + {"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}, + {"type": "text", "text": "prohibited correction"}, + ], + }, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + texts = guardrail.inputs["texts"] + structured = guardrail.inputs["structured_messages"] + assert texts == ["safe text", "prohibited correction"] + bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1") + assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + assert data["messages"][1]["content"][2]["text"] == "[MASKED]" + + @pytest.mark.asyncio + async def test_bedrock_masking_slice_stays_aligned_with_midturn_system(self): + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( + BedrockGuardrail, + ) + + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail() + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "safe text"}, + { + "role": "system", + "content": [ + {"type": "text", "text": "prohibited correction"}, + {"type": "text", "text": "second correction"}, + ], + }, + {"role": "user", "content": "latest question"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + texts = guardrail.inputs["texts"] + structured = guardrail.inputs["structured_messages"] + + bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1") + total = sum(bedrock._count_message_texts(m) for m in structured) + assert total == len(texts) + + latest_user_index = bedrock._find_latest_message_index(structured, target_role="user") + assert latest_user_index == 2 + scanned_slice = bedrock._locate_message_texts_slice( + structured_messages=structured, + target_index=latest_user_index, + texts=texts, + ) + assert scanned_slice == (3, 1) + + merged = bedrock._merge_masked_texts( + masked_texts=["{MASKED}"], + texts=texts, + scanned_slice=scanned_slice, + scanned_role_subset=True, + ) + assert merged == [ + "safe text", + "prohibited correction", + "second correction", + "{MASKED}", + ] + + @pytest.mark.asyncio + async def test_compaction_rewrite_keeps_midturn_system_messages(self): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "user", "content": "compacted history"}, + { + "role": "system", + "content": [{"type": "text", "text": "use the corrected result"}], + }, + {"role": "user", "content": "continue"}, + ] + ) + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "continue"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["user", "system", "user"] + assert data["messages"][1]["content"] == [{"type": "text", "text": "use the corrected result"}] + assert data["messages"][0]["content"] == [{"type": "text", "text": "compacted history"}] + assert data["messages"][2]["content"] == [{"type": "text", "text": "continue"}] + assert data["system"] == "trusted top-level system prompt" + + @pytest.mark.asyncio + async def test_compaction_rewrite_does_not_duplicate_hoisted_top_level_system(self): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "system", "content": "trusted top-level system prompt"}, + {"role": "user", "content": "compacted history"}, + {"role": "system", "content": "use the corrected result"}, + ] + ) + guardrail.skip_system_message_in_guardrail = None + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["user", "system"] + assert data["messages"][1]["content"] == "use the corrected result" + assert data["system"] == "trusted top-level system prompt" + + @pytest.mark.asyncio + async def test_compaction_rewrite_keeps_leading_midturn_system_when_system_is_skipped( + self, + ): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "compacted history"}, + ] + ) + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "trusted top-level system prompt", + "messages": [ + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "original history"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["system", "user"] + assert data["messages"][0]["content"] == "use the corrected result" + + @pytest.mark.asyncio + async def test_compaction_rewrite_keeps_leading_correction_when_top_level_system_hoists_nothing( + self, + ): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "compacted history"}, + ] + ) + guardrail.skip_system_message_in_guardrail = None + data = { + "model": "claude-3-5-sonnet-20241022", + "system": [{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}], + "messages": [ + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "original history"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["system", "user"] + assert data["messages"][0]["content"] == "use the corrected result" + + @pytest.mark.asyncio + async def test_compaction_rewrite_keeps_leading_correction_when_hoisted_prompt_is_dropped( + self, + ): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "system", "content": "CLIENT CORRECTION"}, + {"role": "user", "content": "compacted history"}, + ] + ) + guardrail.skip_system_message_in_guardrail = None + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "TRUSTED", + "messages": [ + {"role": "system", "content": "CLIENT CORRECTION"}, + {"role": "user", "content": "original history"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + assert guardrail.inputs["structured_messages"][0] == { + "role": "system", + "content": "TRUSTED", + } + assert [m["role"] for m in data["messages"]] == ["system", "user"] + assert data["messages"][0]["content"] == "CLIENT CORRECTION" + assert data["system"] == "TRUSTED" + + @pytest.mark.asyncio + async def test_compaction_rewrite_drops_hoisted_prompt_matched_by_content_copy(self): + import json + + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + json.loads(json.dumps({"role": "system", "content": "TRUSTED"})), + {"role": "user", "content": "compacted history"}, + {"role": "system", "content": "CLIENT CORRECTION"}, + ] + ) + guardrail.skip_system_message_in_guardrail = None + data = { + "model": "claude-3-5-sonnet-20241022", + "system": "TRUSTED", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "CLIENT CORRECTION"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["user", "system"] + assert data["messages"][1]["content"] == "CLIENT CORRECTION" + assert data["system"] == "TRUSTED" + + @pytest.mark.asyncio + async def test_compaction_rewrite_preserves_cache_control_on_system_blocks(self): + """ + `cache_control` on an in-sequence system text block survives the write-back, and is + copied rather than aliased into the guardrail's own returned list. + """ + handler = AnthropicMessagesHandler() + source_cache_control = {"type": "ephemeral"} + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "user", "content": "compacted history"}, + { + "role": "system", + "content": [ + { + "type": "text", + "text": "use the corrected result", + "cache_control": source_cache_control, + } + ], + }, + ] + ) + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert data["messages"][1]["content"] == [ + { + "type": "text", + "text": "use the corrected result", + "cache_control": {"type": "ephemeral"}, + } + ] + assert data["messages"][1]["content"][0]["cache_control"] is not source_cache_control + + @pytest.mark.asyncio + async def test_compaction_rewrite_rstrips_trailing_assistant_in_each_run(self): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "user", "content": "compacted history"}, + {"role": "assistant", "content": "earlier "}, + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "continue"}, + {"role": "assistant", "content": "prefill "}, + ] + ) + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == [ + "user", + "assistant", + "system", + "user", + "assistant", + ] + assert data["messages"][1]["content"] == [{"type": "text", "text": "earlier"}] + assert data["messages"][-1]["content"] == [{"type": "text", "text": "prefill"}] + + @pytest.mark.asyncio + async def test_compaction_rewrite_drops_text_free_system_message(self): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[ + {"role": "user", "content": "compacted history"}, + {"role": "system", "content": [{"type": "text", "text": ""}]}, + {"role": "system", "content": ""}, + {"role": "user", "content": "continue"}, + ] + ) + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "continue"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["user", "user"] + assert data["messages"][0]["content"] == [{"type": "text", "text": "compacted history"}] + assert data["messages"][1]["content"] == [{"type": "text", "text": "continue"}] + + @pytest.mark.asyncio + async def test_noncanonical_system_role_casing_is_still_scanned(self): + handler = AnthropicMessagesHandler() + guardrail = MockMaskingGuardrail() + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "safe text"}, + {"role": "System", "content": "prohibited correction"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert guardrail.inputs is not None + assert "prohibited correction" in guardrail.inputs["texts"] + assert data["messages"][1]["content"] == "[MASKED]" + + @pytest.mark.asyncio + async def test_tool_result_turns_have_a_preexisting_alignment_gap(self): + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( + BedrockGuardrail, + ) + + handler = AnthropicMessagesHandler() + bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1") + tool_loop = [ + {"role": "user", "content": "call the tool"}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "tu_1", "name": "get", "input": {"a": 1}}], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "tu_1", + "content": [{"type": "text", "text": "tool output"}], + } + ], + }, + ] + + async def _slice_for(messages: list): + guardrail = MockMaskingGuardrail() + data = {"model": "claude-3-5-sonnet-20241022", "messages": messages} + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + assert guardrail.inputs is not None + texts = guardrail.inputs["texts"] + structured = guardrail.inputs["structured_messages"] + target_index = bedrock._find_latest_message_index(structured, target_role="user") + return ( + sum(bedrock._count_message_texts(m) for m in structured) - len(texts), + bedrock._locate_message_texts_slice( + structured_messages=structured, + target_index=target_index, + texts=texts, + ), + ) + + with_system = await _slice_for( + tool_loop + + [ + {"role": "system", "content": "use the corrected result"}, + {"role": "user", "content": "latest question"}, + ] + ) + without_system = await _slice_for(tool_loop + [{"role": "user", "content": "latest question"}]) + + assert with_system == without_system == (1, None) + + @pytest.mark.asyncio + async def test_compaction_rewrite_to_only_system_messages_is_rejected(self): + import litellm + + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[{"role": "system", "content": "use the corrected result"}] + ) + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + ], + } + + with patch.object(litellm, "modify_params", False): + with pytest.raises(litellm.BadRequestError, match="at least one non-system message"): + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + @pytest.mark.asyncio + async def test_compaction_rewrite_to_only_system_messages_repaired_with_modify_params( + self, + ): + import litellm + + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail( + replacement_messages=[{"role": "system", "content": "use the corrected result"}] + ) + guardrail.skip_system_message_in_guardrail = True + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "original history"}, + {"role": "system", "content": "use the corrected result"}, + ], + } + + with patch.object(litellm, "modify_params", True): + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert [m["role"] for m in data["messages"]] == ["system", "user"] + assert data["messages"][0]["content"] == "use the corrected result" + + @pytest.mark.asyncio + async def test_compaction_rewrite_without_system_messages_is_unchanged(self): + handler = AnthropicMessagesHandler() + guardrail = MockCompactingGuardrail(replacement_messages=[{"role": "user", "content": "compacted history"}]) + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [ + {"role": "user", "content": "a"}, + {"role": "assistant", "content": "b"}, + {"role": "user", "content": "c"}, + ], + } + + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert data["messages"] == [{"role": "user", "content": [{"type": "text", "text": "compacted history"}]}] + @pytest.mark.asyncio async def test_process_output_streaming_response_empty_choices(self): """Test that streaming response with empty choices doesn't raise IndexError diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index c0c6e315b5b..b72620f9918 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -413,6 +413,224 @@ def test_translate_anthropic_messages_to_openai_tool_message_placement(): ), "Tool message should be placed before user message" +@pytest.mark.parametrize( + ("system_content", "expected_content"), + [ + ("Use the corrected result.", "Use the corrected result."), + ( + [{"type": "text", "text": "Use the corrected result."}], + [{"type": "text", "text": "Use the corrected result."}], + ), + ( + [ + { + "type": "image", + "source": {"type": "url", "url": "https://example.com/a.png"}, + }, + {"type": "text", "text": "Use the corrected result."}, + ], + [{"type": "text", "text": "Use the corrected result."}], + ), + ( + [ + {"type": "text", "text": "First correction."}, + {"type": "text", "text": "Second correction."}, + ], + [ + {"type": "text", "text": "First correction."}, + {"type": "text", "text": "Second correction."}, + ], + ), + ], +) +def test_translate_anthropic_messages_to_openai_preserves_midturn_system_correction( + system_content: object, + expected_content: object, +): + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_01234", + "name": "get_weather", + "input": {"location": "Boston"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01234", + "content": "Rainy, 55°F", + } + ], + }, + {"role": "system", "content": system_content}, + {"role": "user", "content": "Continue."}, + ] + + result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( + messages=messages, + model="claude-3-5-sonnet-20240620", + ) + + assert result == [ + { + "role": "assistant", + "content": None, + "thinking_blocks": None, + "tool_calls": [ + { + "id": "toolu_01234", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Boston"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "toolu_01234", + "content": "Rainy, 55°F", + }, + {"role": "system", "content": expected_content}, + {"role": "user", "content": "Continue."}, + ] + + +def test_translate_anthropic_messages_to_openai_preserves_midturn_system_cache_control(): + """ + `cache_control` on an in-sequence system text block survives, matching how the + hoisted top-level `system` prompt and user text blocks are already handled. + """ + messages = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Use the corrected result.", + "cache_control": {"type": "ephemeral"}, + } + ], + } + ] + + result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( + messages=messages, + model="claude-3-5-sonnet-20240620", + ) + + assert result == [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Use the corrected result.", + "cache_control": {"type": "ephemeral"}, + } + ], + } + ] + + +def test_translate_anthropic_messages_to_openai_drops_midturn_system_cache_control_for_non_claude(): + """ + `cache_control` goes through the same `_add_cache_control_if_applicable` gate as the + hoisted top-level prompt and user text blocks, so a non-Claude *requested model name* + does not get it. That gate is a best-effort check of the requested name before routing + (behind the proxy it is often a public alias), not a guarantee about the backend that + ultimately serves the request. + """ + messages = [ + { + "role": "system", + "content": [ + { + "type": "text", + "text": "Use the corrected result.", + "cache_control": {"type": "ephemeral"}, + } + ], + } + ] + + result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( + messages=messages, + model="gpt-4o", + ) + + assert result == [ + { + "role": "system", + "content": [{"type": "text", "text": "Use the corrected result."}], + } + ] + + +@pytest.mark.parametrize( + "system_content", + [ + "", + [{"type": "text", "text": ""}], + [ + { + "type": "image", + "source": {"type": "url", "url": "https://example.com/a.png"}, + } + ], + None, + ], +) +def test_translate_anthropic_messages_to_openai_drops_empty_midturn_system( + system_content: object, +): + messages = [{"role": "system", "content": system_content}] + + result = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( + messages=messages, + model="claude-3-5-sonnet-20240620", + ) + + assert result == [] + + +def test_translate_anthropic_to_openai_orders_top_level_and_midturn_system(): + """ + Request level: the trusted top-level prompt is hoisted to index 0 exactly once and the + in-sequence correction keeps its own position and `role: "system"` -- no duplication of + either, and no reordering of the surrounding turns. + """ + openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( + anthropic_message_request={ + "model": "claude-3-5-sonnet-20240620", + "max_tokens": 100, + "system": "Trusted top-level prompt.", + "messages": [ + {"role": "user", "content": "First question."}, + {"role": "assistant", "content": "First answer."}, + {"role": "system", "content": "Use the corrected result."}, + {"role": "user", "content": "Continue."}, + ], + } + ) + + assert openai_request["messages"] == [ + {"role": "system", "content": "Trusted top-level prompt."}, + {"role": "user", "content": "First question."}, + {"role": "assistant", "content": "First answer.", "thinking_blocks": None}, + {"role": "system", "content": "Use the corrected result."}, + {"role": "user", "content": "Continue."}, + ] + + def test_translate_openai_content_to_anthropic_empty_function_arguments(): """Test that empty function arguments are handled safely and don't cause JSON parsing errors.""" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_compact.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_compact.py index 9c8df1c79f9..6cc1d9e5add 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_compact.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/context_management/test_compact.py @@ -2475,3 +2475,57 @@ def test_endpoint_runs_failure_hook_on_500_context_management_error(): body = response.json() assert body["type"] == "error" failure_hook.assert_awaited_once() + + +def test_count_effective_tokens_counts_midturn_system_correction(): + from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import ( + _count_effective_tokens, + ) + + base: List[Dict[str, Any]] = [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi"}, + ] + correction = { + "role": "system", + "content": [{"type": "text", "text": "use the corrected result " * 20}], + } + + without_correction = _count_effective_tokens( + model=MODEL, effective_messages=base, compaction_block=None, tools=None + ) + with_correction = _count_effective_tokens( + model=MODEL, + effective_messages=base + [correction], + compaction_block=None, + tools=None, + ) + + assert with_correction > without_correction + + +def test_build_summary_messages_keeps_midturn_system_correction_in_place(): + from litellm.llms.anthropic.experimental_pass_through.context_management.editors.compact import ( + _build_summary_messages, + ) + + summary_messages = _build_summary_messages( + effective_messages=[ + {"role": "user", "content": "original question"}, + {"role": "system", "content": "use the corrected result"}, + {"role": "assistant", "content": "acknowledged"}, + ], + prompt="summarize the conversation", + system="caller system prompt", + ) + + assert [m["role"] for m in summary_messages] == [ + "system", + "user", + "system", + "assistant", + "user", + ] + assert summary_messages[0]["content"] == "caller system prompt" + assert summary_messages[2]["content"] == "use the corrected result" + assert summary_messages[-1]["content"] == "summarize the conversation" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py index 606ff39b35e..8963012ecd5 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py @@ -9,6 +9,8 @@ import sys from typing import Any, Dict, List from unittest.mock import MagicMock +import pytest + sys.path.insert(0, os.path.abspath("../../../../../../..")) from litellm.constants import ( @@ -221,6 +223,106 @@ class TestTranslateMessagesToResponsesInput: {"type": "input_text", "text": "Second part."}, ] + @pytest.mark.parametrize( + "system_content", + [ + "Use the corrected result.", + [{"type": "text", "text": "Use the corrected result."}], + [ + {"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}, + {"type": "text", "text": "Use the corrected result."}, + ], + ], + ) + def test_midturn_system_correction_stays_system_in_sequence(self, system_content: object): + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_01234", + "name": "get_weather", + "input": {"location": "Boston"}, + } + ], + }, + { + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "toolu_01234", + "content": "Rainy, 55°F", + } + ], + }, + {"role": "system", "content": system_content}, + {"role": "user", "content": "Continue."}, + ] + + result = _translate_messages(messages) + + assert result == [ + { + "type": "function_call", + "call_id": "toolu_01234", + "name": "get_weather", + "arguments": '{"location": "Boston"}', + }, + { + "type": "function_call_output", + "call_id": "toolu_01234", + "output": "Rainy, 55°F", + }, + { + "type": "message", + "role": "system", + "content": [{"type": "input_text", "text": "Use the corrected result."}], + }, + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Continue."}], + }, + ] + + def test_midturn_system_correction_keeps_multiple_text_blocks(self): + messages = [ + { + "role": "system", + "content": [ + {"type": "text", "text": "First correction."}, + {"type": "text", "text": "Second correction."}, + ], + } + ] + + assert _translate_messages(messages) == [ + { + "type": "message", + "role": "system", + "content": [ + {"type": "input_text", "text": "First correction."}, + {"type": "input_text", "text": "Second correction."}, + ], + } + ] + + @pytest.mark.parametrize( + "system_content", + [ + "", + [{"type": "text", "text": ""}], + [{"type": "image", "source": {"type": "url", "url": "https://example.com/a.png"}}], + None, + ], + ) + def test_empty_or_unsupported_midturn_system_correction_is_dropped(self, system_content: object): + messages = [{"role": "system", "content": system_content}] + + assert _translate_messages(messages) == [] + def test_user_base64_image(self): """User message with base64 image source becomes input_image with data URL.""" messages = [ @@ -722,6 +824,42 @@ class TestTranslateRequestBroaderCoverage: kwargs = _ADAPTER.translate_request(req) assert kwargs["instructions"] == "You are a helpful assistant." + def test_top_level_system_and_midturn_correction_are_not_duplicated(self): + """ + Request level: the trusted top-level prompt goes to `instructions` only, and the + in-sequence correction stays a `role: "system"` input item in its original position. + Neither appears twice, and the surrounding turns keep their order. + """ + req = _make_request( + system="Trusted top-level prompt.", + messages=[ + {"role": "user", "content": "First question."}, + {"role": "system", "content": "Use the corrected result."}, + {"role": "user", "content": "Continue."}, + ], + ) + + kwargs = _ADAPTER.translate_request(req) + + assert kwargs["instructions"] == "Trusted top-level prompt." + assert kwargs["input"] == [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "First question."}], + }, + { + "type": "message", + "role": "system", + "content": [{"type": "input_text", "text": "Use the corrected result."}], + }, + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Continue."}], + }, + ] + def test_system_list_of_text_blocks_joined(self): req = _make_request( system=[ From 5a5bb8c9d844870c25684e169960f1571d08e5ce Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 15:29:39 +0000 Subject: [PATCH 002/139] fix(proxy): stop /{provider}/v1/files from capturing /openai_passthrough The native files and batches routes declare /{provider}/v1/... and their routers are mounted before the passthrough router, so /openai_passthrough/v1/files and /openai_passthrough/v1/batches matched them with provider="openai_passthrough" and 500'd on the LlmProviders lookup instead of reaching openai_proxy_route. Move the dedicated /openai_passthrough prefix onto its own router mounted ahead of the batches and files routers. /openai/... and every other provider prefix keep their current behavior. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llm_passthrough_endpoints.py | 3 +- litellm/proxy/proxy_server.py | 2 + .../test_llm_pass_through_endpoints.py | 57 +++++++++++++++++++ 3 files changed, 61 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 38da00a3bb9..baa74c19182 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -60,6 +60,7 @@ from .passthrough_endpoint_router import PassthroughEndpointRouter vertex_llm_base: Final = VertexBase() router: Final = APIRouter() +openai_passthrough_router: Final = APIRouter() default_vertex_config: Final = None passthrough_endpoint_router: Final = PassthroughEndpointRouter() @@ -1875,7 +1876,7 @@ async def vertex_proxy_route( ) -@router.api_route( +@openai_passthrough_router.api_route( "/openai_passthrough/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], tags=["OpenAI Pass-through", "pass-through"], diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index fb9c4e67aad..e75277e7f0a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -522,6 +522,7 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import ( set_files_config, ) from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + openai_passthrough_router, passthrough_endpoint_router, vertex_ai_live_websocket_passthrough, ) @@ -16433,6 +16434,7 @@ app.include_router(search_router) app.include_router(image_router) app.include_router(fine_tuning_router) app.include_router(credential_router) +app.include_router(openai_passthrough_router) app.include_router(batches_router) app.include_router(openai_files_router) app.include_router(llm_passthrough_router) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 181846fe289..27d6e4c8585 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -2814,6 +2814,63 @@ class TestOpenAIPassthroughRoute: assert result == {"id": "asst_123", "object": "assistant"} +def _resolve_route_name(method: str, path: str) -> str | None: + from starlette.routing import Match + + from litellm.proxy.proxy_server import app + + scope = { + "type": "http", + "method": method, + "path": path, + "headers": [], + "query_string": b"", + "root_path": "", + } + for route in app.router.routes: + if route.matches(scope)[0] == Match.FULL: + return getattr(route, "name", None) + return None + + +@pytest.mark.parametrize( + "method, path", + [ + ("POST", "/openai_passthrough/v1/files"), + ("GET", "/openai_passthrough/v1/files"), + ("GET", "/openai_passthrough/v1/files/file-abc123"), + ("DELETE", "/openai_passthrough/v1/files/file-abc123"), + ("GET", "/openai_passthrough/v1/files/file-abc123/content"), + ("POST", "/openai_passthrough/v1/batches"), + ("GET", "/openai_passthrough/v1/batches"), + ("GET", "/openai_passthrough/v1/batches/batch_abc123"), + ("POST", "/openai_passthrough/v1/batches/batch_abc123/cancel"), + ("POST", "/openai_passthrough/v1/responses"), + ], +) +def test_openai_passthrough_prefix_wins_over_native_provider_routes(method, path): + """ + /openai_passthrough exists to guarantee passthrough, so the native + /{provider}/v1/files and /{provider}/v1/batches routes must never capture it + with provider="openai_passthrough" (which 500s on the LlmProviders lookup). + """ + assert _resolve_route_name(method, path) == "openai_proxy_route" + + +@pytest.mark.parametrize( + "method, path, expected_name", + [ + ("POST", "/openai/v1/files", "create_file"), + ("GET", "/azure/v1/files", "list_files"), + ("POST", "/v1/files", "create_file"), + ("POST", "/v1/batches", "create_batch"), + ("POST", "/openai/v1/chat/completions", "openai_proxy_route"), + ], +) +def test_native_provider_routes_are_unchanged(method, path, expected_name): + assert _resolve_route_name(method, path) == expected_name + + class TestCursorProxyRoute: """Tests for the Cursor Cloud Agents pass-through route.""" From 357f90fa39d18c9a158a978ebd1ed0fecac6044d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 15:33:40 +0000 Subject: [PATCH 003/139] fix(proxy): scope file list pagination cursors to the caller GET /v1/files filters data down to the caller's own managed files but left first_id and last_id as the upstream page's, so a non-owner got back file ids belonging to other users even with an empty data array --- .../proxy/hooks/managed_files.py | 15 +++ .../proxy/hooks/test_managed_files.py | 100 ++++++++++++++++++ 2 files changed, 115 insertions(+) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 0036603bcd1..851e202e2fb 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1270,10 +1270,25 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) ## Filter the response to only include the files created by the user response.data = user_created_file_ids # type: ignore + self._scope_list_page_cursors(response, user_created_file_ids) return response return response return response + @staticmethod + def _scope_list_page_cursors(response: AsyncCursorPage, data: List[OpenAIFileObject]) -> None: + """Rebuild ``first_id`` / ``last_id`` from the caller-scoped page. + + The upstream cursors point at rows that were just filtered out, so + leaving them in place discloses other callers' file ids. + """ + if hasattr(response, "first_id"): + response.first_id = data[0].id if data else None + if hasattr(response, "last_id"): + response.last_id = data[-1].id if data else None + if not data and hasattr(response, "has_more"): + response.has_more = False + async def afile_retrieve( self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router=None ) -> OpenAIFileObject: diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 50af6465d06..3384c553740 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -2861,3 +2861,103 @@ async def test_same_user_different_keys_can_access_batch(): assert "batch_id" in result2 # Both keys should get the same result assert result1["batch_id"] == result2["batch_id"] + + +@pytest.mark.asyncio +async def test_file_list_cursors_are_scoped_to_the_caller(): + """A non-owner must not learn other callers' file ids through the page cursors.""" + from openai.pagination import AsyncCursorPage + from openai.types import FileObject + + from litellm.proxy._types import UserAPIKeyAuth + + owner_file = FileObject( + id="file-owner-1", + bytes=100, + created_at=1, + filename="owner.jsonl", + object="file", + purpose="batch", + status="processed", + ) + upstream_page = AsyncCursorPage[FileObject].construct( + data=[owner_file], + has_more=True, + first_id=owner_file.id, + last_id=owner_file.id, + object="list", + ) + + prisma_client = AsyncMock() + prisma_client.db.litellm_managedfiletable.find_many.return_value = [] + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + response = await proxy_managed_files.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth( + user_id="other-user", team_id="other-team", parent_otel_span=MagicMock() + ), + response=upstream_page, + ) + + assert response.data == [] + assert response.first_id is None + assert response.last_id is None + assert response.has_more is False + + +@pytest.mark.asyncio +async def test_file_list_cursors_follow_the_owner_scoped_page(): + from openai.pagination import AsyncCursorPage + from openai.types import FileObject + + from litellm.proxy._types import UserAPIKeyAuth + + def _raw_file(file_id: str) -> FileObject: + return FileObject( + id=file_id, + bytes=100, + created_at=1, + filename=f"{file_id}.jsonl", + object="file", + purpose="batch", + status="processed", + ) + + upstream_page = AsyncCursorPage[FileObject].construct( + data=[_raw_file("file-someone-else"), _raw_file("file-mine")], + has_more=False, + first_id="file-someone-else", + last_id="file-mine", + object="list", + ) + + managed_row = MagicMock() + managed_row.file_object = { + "id": "litellm_proxy:mine", + "bytes": 100, + "created_at": 1, + "filename": "mine.jsonl", + "object": "file", + "purpose": "batch", + "status": "processed", + } + prisma_client = AsyncMock() + prisma_client.db.litellm_managedfiletable.find_many.return_value = [managed_row] + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + response = await proxy_managed_files.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth( + user_id="mine-user", parent_otel_span=MagicMock() + ), + response=upstream_page, + ) + + assert [file_object.id for file_object in response.data] == ["litellm_proxy:mine"] + assert response.first_id == "litellm_proxy:mine" + assert response.last_id == "litellm_proxy:mine" From 9d69fdac72b860627f627f9802ceb4758a7c372d Mon Sep 17 00:00:00 2001 From: daleselaji-dev <265319989+daleselaji-dev@users.noreply.github.com> Date: Fri, 7 Aug 2026 13:14:56 +0800 Subject: [PATCH 004/139] fix(bedrock): use deployment credentials for AWS requests --- .../llms/bedrock/batches/transformation.py | 11 ++- litellm/llms/bedrock/files/transformation.py | 11 ++- .../test_bedrock_request_credentials.py | 95 +++++++++++++++++++ 3 files changed, 109 insertions(+), 8 deletions(-) create mode 100644 tests/litellm/llms/bedrock/test_bedrock_request_credentials.py diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index 9e3dec26673..8b437af8581 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -130,7 +130,8 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): Get the complete URL for Bedrock batch creation. Bedrock batch jobs are created via the model invocation job API. """ - aws_region_name: Final = self._get_aws_region_name(optional_params, model) + request_params: Final = {**litellm_params, **optional_params} + aws_region_name: Final = self._get_aws_region_name(request_params, model) # Bedrock model invocation job endpoint # Format: https://bedrock.{region}.amazonaws.com/model-invocation-job @@ -232,14 +233,15 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): # For Bedrock, we need to return a pre-signed request with AWS auth headers # Use common utility for AWS signing + request_params: Final = {**litellm_params, **optional_params} endpoint_url: Final = ( - f"https://bedrock.{self._get_aws_region_name(optional_params, model)}.amazonaws.com/model-invocation-job" + f"https://bedrock.{self._get_aws_region_name(request_params, model)}.amazonaws.com/model-invocation-job" ) signed_headers, signed_data = self.common_utils.sign_aws_request( service_name="bedrock", data=bedrock_request, endpoint_url=endpoint_url, - optional_params=optional_params, + optional_params=request_params, method="POST", ) @@ -387,11 +389,12 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): endpoint_url: Final = f"https://bedrock.{region}.amazonaws.com/model-invocation-job/{encoded_arn}" # Use common utility for AWS signing + request_params: Final = {**litellm_params, **optional_params} signed_headers, _ = self.common_utils.sign_aws_request( service_name="bedrock", data={}, # GET request has no body endpoint_url=endpoint_url, - optional_params=optional_params, + optional_params=request_params, method="GET", ) diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 728fd02f001..072343ae936 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -257,6 +257,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): """ Get the complete S3 URL for the file upload request """ + request_params: Final = {**litellm_params, **optional_params} bucket_name = litellm_params.get("s3_bucket_name") or os.getenv("AWS_S3_BUCKET_NAME") if not bucket_name: raise ValueError( @@ -265,7 +266,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name) s3_region_name: Final = litellm_params.get("s3_region_name") or optional_params.get("s3_region_name") - aws_region_name: Final = s3_region_name or self._get_aws_region_name(optional_params, model) + aws_region_name: Final = s3_region_name or self._get_aws_region_name(request_params, model) file_data: Final = data.get("file") purpose: Final = data.get("purpose") @@ -281,7 +282,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): # S3 endpoint URL format s3_endpoint_url: Final = ( - optional_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com" + request_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com" ).rstrip("/") return f"{s3_endpoint_url}/{bucket_name}/{encoded_object_name}" @@ -728,14 +729,16 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): if s3_region_name: optional_params = {**optional_params, "aws_region_name": s3_region_name} + request_params: Final = {**litellm_params, **optional_params} + # Sign the request and return a pre-signed request object signed_headers, signed_body = self._sign_s3_request( content=file_content, api_base=api_base, - optional_params=optional_params, + optional_params=request_params, s3_encryption_key_id=resolve_s3_encryption_key_id( litellm_params=litellm_params, - optional_params=optional_params, + optional_params=request_params, ), ) diff --git a/tests/litellm/llms/bedrock/test_bedrock_request_credentials.py b/tests/litellm/llms/bedrock/test_bedrock_request_credentials.py new file mode 100644 index 00000000000..72da2daf9b5 --- /dev/null +++ b/tests/litellm/llms/bedrock/test_bedrock_request_credentials.py @@ -0,0 +1,95 @@ +from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig +from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + +def test_bedrock_file_upload_signing_uses_deployment_credentials(monkeypatch): + config = BedrockFilesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, "" + + monkeypatch.setattr(config, "_sign_s3_request", capture_signing) + + result = config.transform_create_file_request( + model="", + create_file_data={ + "file": ( + "batch.jsonl", + b'{"custom_id":"req-1","body":{"model":"bedrock/model"}}\n', + "application/jsonl", + ), + "purpose": "batch", + }, + optional_params={}, + litellm_params={ + "s3_bucket_name": "deployment-bucket", + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + }, + ) + + assert "eu-west-1" in result["url"] + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" + assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" + assert captured["optional_params"]["aws_region_name"] == "eu-west-1" + + +def test_bedrock_batch_signing_uses_deployment_credentials(monkeypatch): + config = BedrockBatchesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, b"{}" + + monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) + + result = config.transform_create_batch_request( + model="us.anthropic.claude-haiku-4-5-20251001-v1:0", + create_batch_data={ + "input_file_id": "s3://deployment-bucket/input.jsonl", + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + }, + optional_params={}, + litellm_params={ + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + "aws_batch_role_arn": "arn:aws:iam::123456789012:role/bedrock-batch", + }, + ) + + assert result["url"].startswith("https://bedrock.eu-west-1.amazonaws.com/") + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" + assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" + assert captured["optional_params"]["aws_region_name"] == "eu-west-1" + + +def test_bedrock_batch_retrieval_signing_uses_deployment_credentials(monkeypatch): + config = BedrockBatchesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, b"" + + monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) + + result = config.transform_retrieve_batch_request( + batch_id="arn:aws:bedrock:eu-west-1:123456789012:model-invocation-job/job-1", + optional_params={}, + litellm_params={ + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + }, + ) + + assert result["url"].startswith("https://bedrock.eu-west-1.amazonaws.com/") + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" + assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" + assert captured["optional_params"]["aws_region_name"] == "eu-west-1" From f9b86b253a3fb87d003bb5ccc80c7d89aa91dd62 Mon Sep 17 00:00:00 2001 From: Harry Qian Date: Tue, 4 Aug 2026 17:14:26 +0800 Subject: [PATCH 005/139] fix(proxy): restore query-param validation under fastapi>=0.140.7 fastapi 0.140.7 removed get_flat_dependant(), which broke the import in management_v1/common.py and took down every /management/v1 route. Switch to get_flat_params() and filter to ParamTypes.query so unknown-query-param rejection keeps matching the old behavior. --- .../management_endpoints/management_v1/common.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/management_v1/common.py b/litellm/proxy/management_endpoints/management_v1/common.py index 8525d67a041..ec79820465a 100644 --- a/litellm/proxy/management_endpoints/management_v1/common.py +++ b/litellm/proxy/management_endpoints/management_v1/common.py @@ -4,7 +4,8 @@ from typing import Final from urllib.parse import urlencode from fastapi import Request -from fastapi.dependencies.utils import get_flat_dependant +from fastapi.dependencies.utils import get_flat_params +from fastapi.params import ParamTypes from fastapi.responses import JSONResponse from litellm.types.proxy.management_endpoints.management_v1 import ( @@ -42,7 +43,13 @@ def _declared_query_params(request: Request) -> frozenset[str]: dependant: Final = getattr(route, "dependant", None) if dependant is None: return frozenset() - return frozenset(field.alias for field in get_flat_dependant(dependant, skip_repeats=True).query_params) + # fastapi>=0.140.7 removed get_flat_dependant(); get_flat_params() returns the + # flattened (deduped) param list. Filter to query params to match the old behavior. + return frozenset( + field.alias + for field in get_flat_params(dependant) + if getattr(field.field_info, "in_", None) == ParamTypes.query + ) def escape_like(value: str) -> str: From da443d1266615507f52a101b461c80e0265069ae Mon Sep 17 00:00:00 2001 From: Harry Qian Date: Tue, 4 Aug 2026 18:21:22 +0800 Subject: [PATCH 006/139] test(proxy): lock in query-param validation across fastapi param types Guards _declared_query_params against a regression in the get_flat_params migration: the flatten step returns path, query, header and cookie params together, so a dropped ParamTypes.query filter would wrongly treat path or header names as declared query params and accept unknown ones. Removing the filter fails these tests. --- .../management_v1/test_common.py | 95 +++++++++++++++++++ 1 file changed, 95 insertions(+) create mode 100644 tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py new file mode 100644 index 00000000000..167a06ed551 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py @@ -0,0 +1,95 @@ +from typing import Annotated + +from fastapi import Depends, FastAPI, Header, Query, Request +from fastapi.testclient import TestClient + +from litellm.proxy.management_endpoints.management_v1.common import ( + ManagementProblem, + PROBLEM_CONTENT_TYPE, + _declared_query_params, + problem_response, + reject_unknown_query_params, +) + + +def _client() -> TestClient: + app = FastAPI() + + @app.exception_handler(ManagementProblem) + async def _handle(_request: Request, exc: ManagementProblem): + return problem_response(exc.problem) + + @app.get("/things/{thing_id}", dependencies=[Depends(reject_unknown_query_params)]) + def _handler( + thing_id: str, + request: Request, + status: Annotated[str | None, Query(alias="filter[status]")] = None, + page: Annotated[int, Query(ge=1)] = 1, + x_trace: Annotated[str | None, Header()] = None, + ) -> dict[str, bool]: + return {"ok": True} + + return TestClient(app, raise_server_exceptions=False) + + +def test_a_declared_query_param_is_accepted_by_its_alias(): + response = _client().get("/things/abc", params={"filter[status]": "active", "page": "2"}) + assert response.status_code == 200, response.text + + +def test_an_unknown_query_param_is_rejected_as_a_problem(): + response = _client().get("/things/abc", params={"bogus": "x"}) + assert response.status_code == 400 + assert response.headers["content-type"].startswith(PROBLEM_CONTENT_TYPE) + assert "bogus" in response.json()["detail"] + + +def test_a_path_param_name_is_not_a_declared_query_param(): + """The flatten step returns path+query+header together; only query names count as declared. + + If the ParamTypes.query filter were dropped, `thing_id` (a path param) would leak + into the declared set and this request would be wrongly accepted. + """ + response = _client().get("/things/abc", params={"thing_id": "x"}) + assert response.status_code == 400 + assert "thing_id" in response.json()["detail"] + + +def test_a_header_param_name_is_not_a_declared_query_param(): + response = _client().get("/things/abc", params={"x-trace": "x"}) + assert response.status_code == 400 + assert "x-trace" in response.json()["detail"] + + +def test_declared_query_params_isolates_query_aliases_from_other_param_types(): + captured: dict[str, frozenset[str]] = {} + app = FastAPI() + + @app.get("/things/{thing_id}") + def _handler( + thing_id: str, + request: Request, + status: Annotated[str | None, Query(alias="filter[status]")] = None, + page: Annotated[int, Query(ge=1)] = 1, + x_trace: Annotated[str | None, Header()] = None, + ) -> dict[str, bool]: + captured["declared"] = _declared_query_params(request) + return {"ok": True} + + TestClient(app).get("/things/abc") + assert captured["declared"] == frozenset({"filter[status]", "page"}) + + +def test_declared_query_params_is_empty_when_the_route_has_no_dependant(): + request = Request( + { + "type": "http", + "method": "GET", + "scheme": "http", + "root_path": "", + "path": "/things/abc", + "query_string": b"", + "headers": [(b"host", b"testserver")], + } + ) + assert _declared_query_params(request) == frozenset() From b11d342022e1dfe88ad05438ecab49524173028a Mon Sep 17 00:00:00 2001 From: daleselaji-dev <265319989+daleselaji-dev@users.noreply.github.com> Date: Fri, 7 Aug 2026 16:40:17 +0800 Subject: [PATCH 007/139] chore(bedrock): document mutable signing params --- litellm/llms/bedrock/batches/transformation.py | 15 ++++++++++++--- litellm/llms/bedrock/files/transformation.py | 10 ++++++++-- 2 files changed, 20 insertions(+), 5 deletions(-) diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index 8b437af8581..e92a2e7a00a 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -130,7 +130,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): Get the complete URL for Bedrock batch creation. Bedrock batch jobs are created via the model invocation job API. """ - request_params: Final = {**litellm_params, **optional_params} + request_params: Final = { + **litellm_params, + **optional_params, + } # mutable-ok: merged params are read by AWS helpers aws_region_name: Final = self._get_aws_region_name(request_params, model) # Bedrock model invocation job endpoint @@ -233,7 +236,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): # For Bedrock, we need to return a pre-signed request with AWS auth headers # Use common utility for AWS signing - request_params: Final = {**litellm_params, **optional_params} + request_params: Final = { + **litellm_params, + **optional_params, + } # mutable-ok: merged params are read by AWS helpers endpoint_url: Final = ( f"https://bedrock.{self._get_aws_region_name(request_params, model)}.amazonaws.com/model-invocation-job" ) @@ -389,7 +395,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): endpoint_url: Final = f"https://bedrock.{region}.amazonaws.com/model-invocation-job/{encoded_arn}" # Use common utility for AWS signing - request_params: Final = {**litellm_params, **optional_params} + request_params: Final = { + **litellm_params, + **optional_params, + } # mutable-ok: merged params are read by AWS helpers signed_headers, _ = self.common_utils.sign_aws_request( service_name="bedrock", data={}, # GET request has no body diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 072343ae936..cd4f52ddbfe 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -257,7 +257,10 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): """ Get the complete S3 URL for the file upload request """ - request_params: Final = {**litellm_params, **optional_params} + request_params: Final = { + **litellm_params, + **optional_params, + } # mutable-ok: merged params are read by AWS helpers bucket_name = litellm_params.get("s3_bucket_name") or os.getenv("AWS_S3_BUCKET_NAME") if not bucket_name: raise ValueError( @@ -729,7 +732,10 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): if s3_region_name: optional_params = {**optional_params, "aws_region_name": s3_region_name} - request_params: Final = {**litellm_params, **optional_params} + request_params: Final = { + **litellm_params, + **optional_params, + } # mutable-ok: merged params are read by AWS helpers # Sign the request and return a pre-signed request object signed_headers, signed_body = self._sign_s3_request( From 8f998a9ca477942815b525656bbc6f52a9fcecef Mon Sep 17 00:00:00 2001 From: daleselaji-dev <265319989+daleselaji-dev@users.noreply.github.com> Date: Fri, 7 Aug 2026 16:45:06 +0800 Subject: [PATCH 008/139] fix(bedrock): prevent caller AWS identity override --- .../llms/bedrock/batches/transformation.py | 21 ++- litellm/llms/bedrock/common_utils.py | 38 +++++ litellm/llms/bedrock/files/transformation.py | 12 +- .../test_bedrock_files_and_batches.py | 130 ++++++++++++++++++ .../test_bedrock_request_credentials.py | 95 ------------- 5 files changed, 179 insertions(+), 117 deletions(-) delete mode 100644 tests/litellm/llms/bedrock/test_bedrock_request_credentials.py diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index e92a2e7a00a..04f395f2bf1 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -28,7 +28,11 @@ from litellm.types.llms.openai import ( from litellm.types.utils import LiteLLMBatch, LlmProviders from ..base_aws_llm import BaseAWSLLM -from ..common_utils import CommonBatchFilesUtils, resolve_s3_encryption_key_id +from ..common_utils import ( + CommonBatchFilesUtils, + merge_bedrock_aws_request_params, + resolve_s3_encryption_key_id, +) # Bedrock batch input files are uploaded as # s3://bucket/litellm-bedrock-files-{model, ":" -> "-"}-{uuid4}.jsonl (see @@ -130,10 +134,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): Get the complete URL for Bedrock batch creation. Bedrock batch jobs are created via the model invocation job API. """ - request_params: Final = { - **litellm_params, - **optional_params, - } # mutable-ok: merged params are read by AWS helpers + request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params) aws_region_name: Final = self._get_aws_region_name(request_params, model) # Bedrock model invocation job endpoint @@ -236,10 +237,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): # For Bedrock, we need to return a pre-signed request with AWS auth headers # Use common utility for AWS signing - request_params: Final = { - **litellm_params, - **optional_params, - } # mutable-ok: merged params are read by AWS helpers + request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params) endpoint_url: Final = ( f"https://bedrock.{self._get_aws_region_name(request_params, model)}.amazonaws.com/model-invocation-job" ) @@ -395,10 +393,7 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): endpoint_url: Final = f"https://bedrock.{region}.amazonaws.com/model-invocation-job/{encoded_arn}" # Use common utility for AWS signing - request_params: Final = { - **litellm_params, - **optional_params, - } # mutable-ok: merged params are read by AWS helpers + request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params) signed_headers, _ = self.common_utils.sign_aws_request( service_name="bedrock", data={}, # GET request has no body diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index d18cb7d8734..57202f2d626 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -36,6 +36,44 @@ class BedrockError(BaseLLMException): pass +_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = ( + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + "aws_region_name", + "aws_session_name", + "aws_profile_name", + "aws_role_name", + "aws_web_identity_token", + "aws_sts_endpoint", + "aws_external_id", +) + + +def merge_bedrock_aws_request_params( + litellm_params: Mapping[str, Any], + optional_params: Mapping[str, Any], +) -> dict[str, Any]: + """Merge deployment and request parameters without allowing auth escalation. + + Deployment configuration is authoritative for AWS authentication. When a + deployment supplies static credentials, caller-supplied profile/role/token + selectors must not redirect signing to another identity available on the + server. Requests may still provide AWS credentials when the deployment has + no static credentials configured. + """ + request_params: Final = {**optional_params, **litellm_params} # mutable-ok: AWS helpers require a plain dict + has_static_deployment_credentials = all( + isinstance(litellm_params.get(key), str) and bool(litellm_params.get(key)) + for key in ("aws_access_key_id", "aws_secret_access_key", "aws_region_name") + ) + if has_static_deployment_credentials: + for key in _BEDROCK_AWS_AUTH_PARAMETER_KEYS: + if key not in litellm_params: + request_params.pop(key, None) + return request_params + + # Lazy import cache to avoid circular imports and performance impact _get_model_info = None diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index cd4f52ddbfe..7663105ba4a 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -46,7 +46,7 @@ from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums from litellm.utils import get_llm_provider from ..base_aws_llm import BaseAWSLLM -from ..common_utils import BedrockError, resolve_s3_encryption_key_id +from ..common_utils import BedrockError, merge_bedrock_aws_request_params, resolve_s3_encryption_key_id # litellm_params key used to hand the SigV4-signed GET headers from # `transform_file_content_request` to `validate_environment` (the only hook @@ -257,10 +257,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): """ Get the complete S3 URL for the file upload request """ - request_params: Final = { - **litellm_params, - **optional_params, - } # mutable-ok: merged params are read by AWS helpers + request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params) bucket_name = litellm_params.get("s3_bucket_name") or os.getenv("AWS_S3_BUCKET_NAME") if not bucket_name: raise ValueError( @@ -732,10 +729,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): if s3_region_name: optional_params = {**optional_params, "aws_region_name": s3_region_name} - request_params: Final = { - **litellm_params, - **optional_params, - } # mutable-ok: merged params are read by AWS helpers + request_params: Final = merge_bedrock_aws_request_params(litellm_params, optional_params) # Sign the request and return a pre-signed request object signed_headers, signed_body = self._sign_s3_request( diff --git a/tests/batches_tests/test_bedrock_files_and_batches.py b/tests/batches_tests/test_bedrock_files_and_batches.py index 431d5a2a60c..52e88937916 100644 --- a/tests/batches_tests/test_bedrock_files_and_batches.py +++ b/tests/batches_tests/test_bedrock_files_and_batches.py @@ -389,3 +389,133 @@ def test_bedrock_batch_with_encryption_key_in_post_request(): ) print("SUCCESS: s3_encryption_key_id properly included in AWS POST request") + + +def test_bedrock_file_upload_signing_uses_deployment_credentials(monkeypatch): + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + config = BedrockFilesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, "" + + monkeypatch.setattr(config, "_sign_s3_request", capture_signing) + + result = config.transform_create_file_request( + model="", + create_file_data={ + "file": ( + "batch.jsonl", + b'{"custom_id":"req-1","body":{"model":"bedrock/model"}}\n', + "application/jsonl", + ), + "purpose": "batch", + }, + optional_params={}, + litellm_params={ + "s3_bucket_name": "deployment-bucket", + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + }, + ) + + assert "eu-west-1" in result["url"] + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" + assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" + assert captured["optional_params"]["aws_region_name"] == "eu-west-1" + + +def test_bedrock_batch_signing_uses_deployment_credentials(monkeypatch): + from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig + + config = BedrockBatchesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, b"{}" + + monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) + + result = config.transform_create_batch_request( + model="us.anthropic.claude-haiku-4-5-20251001-v1:0", + create_batch_data={ + "input_file_id": "s3://deployment-bucket/input.jsonl", + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + }, + optional_params={}, + litellm_params={ + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + "aws_batch_role_arn": "arn:aws:iam::123456789012:role/bedrock-batch", + }, + ) + + assert result["url"].startswith("https://bedrock.eu-west-1.amazonaws.com/") + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" + assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" + assert captured["optional_params"]["aws_region_name"] == "eu-west-1" + + +def test_bedrock_batch_retrieval_signing_uses_deployment_credentials(monkeypatch): + from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig + + config = BedrockBatchesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, b"" + + monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) + + result = config.transform_retrieve_batch_request( + batch_id="arn:aws:bedrock:eu-west-1:123456789012:model-invocation-job/job-1", + optional_params={}, + litellm_params={ + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + }, + ) + + assert result["url"].startswith("https://bedrock.eu-west-1.amazonaws.com/") + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" + assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" + assert captured["optional_params"]["aws_region_name"] == "eu-west-1" + + +def test_bedrock_deployment_credentials_block_caller_profile_override(monkeypatch): + from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig + + config = BedrockBatchesConfig() + captured = {} + + def capture_signing(**kwargs): + captured.update(kwargs) + return {}, b"{}" + + monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) + + config.transform_create_batch_request( + model="us.anthropic.claude-haiku-4-5-20251001-v1:0", + create_batch_data={ + "input_file_id": "s3://deployment-bucket/input.jsonl", + "completion_window": "24h", + }, + optional_params={"aws_profile_name": "caller-controlled-profile"}, + litellm_params={ + "aws_access_key_id": "deployment-access-key", + "aws_secret_access_key": "deployment-secret", + "aws_region_name": "eu-west-1", + "aws_batch_role_arn": "arn:aws:iam::123456789012:role/bedrock-batch", + }, + ) + + assert "aws_profile_name" not in captured["optional_params"] + assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" diff --git a/tests/litellm/llms/bedrock/test_bedrock_request_credentials.py b/tests/litellm/llms/bedrock/test_bedrock_request_credentials.py deleted file mode 100644 index 72da2daf9b5..00000000000 --- a/tests/litellm/llms/bedrock/test_bedrock_request_credentials.py +++ /dev/null @@ -1,95 +0,0 @@ -from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig -from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - - -def test_bedrock_file_upload_signing_uses_deployment_credentials(monkeypatch): - config = BedrockFilesConfig() - captured = {} - - def capture_signing(**kwargs): - captured.update(kwargs) - return {}, "" - - monkeypatch.setattr(config, "_sign_s3_request", capture_signing) - - result = config.transform_create_file_request( - model="", - create_file_data={ - "file": ( - "batch.jsonl", - b'{"custom_id":"req-1","body":{"model":"bedrock/model"}}\n', - "application/jsonl", - ), - "purpose": "batch", - }, - optional_params={}, - litellm_params={ - "s3_bucket_name": "deployment-bucket", - "aws_access_key_id": "deployment-access-key", - "aws_secret_access_key": "deployment-secret", - "aws_region_name": "eu-west-1", - }, - ) - - assert "eu-west-1" in result["url"] - assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" - assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" - assert captured["optional_params"]["aws_region_name"] == "eu-west-1" - - -def test_bedrock_batch_signing_uses_deployment_credentials(monkeypatch): - config = BedrockBatchesConfig() - captured = {} - - def capture_signing(**kwargs): - captured.update(kwargs) - return {}, b"{}" - - monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) - - result = config.transform_create_batch_request( - model="us.anthropic.claude-haiku-4-5-20251001-v1:0", - create_batch_data={ - "input_file_id": "s3://deployment-bucket/input.jsonl", - "completion_window": "24h", - "endpoint": "/v1/chat/completions", - }, - optional_params={}, - litellm_params={ - "aws_access_key_id": "deployment-access-key", - "aws_secret_access_key": "deployment-secret", - "aws_region_name": "eu-west-1", - "aws_batch_role_arn": "arn:aws:iam::123456789012:role/bedrock-batch", - }, - ) - - assert result["url"].startswith("https://bedrock.eu-west-1.amazonaws.com/") - assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" - assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" - assert captured["optional_params"]["aws_region_name"] == "eu-west-1" - - -def test_bedrock_batch_retrieval_signing_uses_deployment_credentials(monkeypatch): - config = BedrockBatchesConfig() - captured = {} - - def capture_signing(**kwargs): - captured.update(kwargs) - return {}, b"" - - monkeypatch.setattr(config.common_utils, "sign_aws_request", capture_signing) - - result = config.transform_retrieve_batch_request( - batch_id="arn:aws:bedrock:eu-west-1:123456789012:model-invocation-job/job-1", - optional_params={}, - litellm_params={ - "aws_access_key_id": "deployment-access-key", - "aws_secret_access_key": "deployment-secret", - "aws_region_name": "eu-west-1", - }, - ) - - assert result["url"].startswith("https://bedrock.eu-west-1.amazonaws.com/") - assert captured["optional_params"]["aws_access_key_id"] == "deployment-access-key" - assert captured["optional_params"]["aws_secret_access_key"] == "deployment-secret" - assert captured["optional_params"]["aws_region_name"] == "eu-west-1" From 5883aa354d42a3225fef485034e86a52a275cdd7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 7 Aug 2026 04:45:39 -0700 Subject: [PATCH 009/139] fix(router): keep batch fallbacks inside the model group that owns the file A batch or fine-tuning job is created from a file the caller already uploaded, and that file only exists under the credentials of the deployment that stored it. When the router fell back to a different model group it handed that file id to a provider that has never seen it, so the caller got the second provider's complaint about the file id instead of the error that explains what was actually wrong with their request. run_async_fallback now skips fallback targets outside the original model group whenever the request carries input_file_id or training_file. Order-based fallbacks stay inside the group, so retrying across deployments still works. The same handler also crashed with "'NoneType' object has no attribute 'update'" whenever a fallback fired on a request with metadata set to None, which /v1/batches always does when the caller sends no metadata, turning the provider's 400 into a 500. Record the model group with a merge instead of setdefault, and write it to litellm_metadata on the endpoints that use it so the router's bookkeeping no longer lands in the metadata stored on the provider's batch. --- .../router_utils/fallback_event_handlers.py | 41 ++++- .../test_fallback_event_handlers.py | 141 ++++++++++++++++++ tests/test_litellm/test_router.py | 62 ++++++++ 3 files changed, 241 insertions(+), 3 deletions(-) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 1c6bb52ccb8..c4a84a1d61e 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -9,6 +9,7 @@ from litellm.router_utils.add_retry_fallback_headers import ( add_fallback_headers_to_response, get_fallback_error_info, ) +from litellm.router_utils.batch_utils import _get_router_metadata_variable_name from litellm.types.router import LiteLLMParamsTypedDict if TYPE_CHECKING: @@ -82,6 +83,28 @@ def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[li return fallback_model_group, generic_fallback_idx +PROVIDER_SCOPED_RESOURCE_KEYS: Final = ("input_file_id", "training_file") + + +def _get_fallback_target_model_group(fallback_entry: str | dict[str, object]) -> str | None: + if isinstance(fallback_entry, str): + return fallback_entry + target: Final = fallback_entry.get("model") + return target if isinstance(target, str) else None + + +def references_provider_scoped_resource(kwargs: dict[str, object]) -> bool: + """ + True when the request names a file that only exists under one provider's credentials. + + Batch and fine-tuning jobs are created from a file the caller already uploaded, and + that file lives in the account of the deployment that stored it. Handing the id to a + different model group can only fail, and the second provider's error replaces the + error the caller actually needs to see. + """ + return any(kwargs.get(key) for key in PROVIDER_SCOPED_RESOURCE_KEYS) + + async def run_async_fallback( *args: tuple[Any], litellm_router: LitellmRouter, @@ -120,10 +143,21 @@ async def run_async_fallback( error_from_fallbacks = original_exception fallback_errors = (get_fallback_error_info(original_exception),) + metadata_variable_name: Final = _get_router_metadata_variable_name( + function_name=getattr(kwargs.get("original_function"), "__name__", None) + ) + same_model_group_only: Final = references_provider_scoped_resource(kwargs) for mg in fallback_model_group: if mg == original_model_group: continue + if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group: + verbose_router_logger.info( + "Skipping fallback to model_group = %s: request is pinned to model_group = %s by its uploaded file", + mask_sensitive_structure(mg), + original_model_group, + ) + continue try: # LOGGING kwargs = litellm_router.log_retry(kwargs=kwargs, e=original_exception) @@ -132,9 +166,10 @@ async def run_async_fallback( kwargs["model"] = mg elif isinstance(mg, dict): kwargs.update(mg) - kwargs.setdefault("metadata", {}).update( - {"model_group": kwargs.get("model", None)} - ) # update model_group used, if fallbacks are done + kwargs[metadata_variable_name] = { + **(kwargs.get(metadata_variable_name) or {}), + "model_group": kwargs.get("model", None), + } # update model_group used, if fallbacks are done fallback_depth = fallback_depth + 1 kwargs["fallback_depth"] = fallback_depth kwargs["max_fallbacks"] = max_fallbacks diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index 98a34de295c..d93aa4ab023 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -142,6 +142,147 @@ async def test_run_async_fallback_skips_original_model_group(): assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1 +class AttemptRecordingRouter: + def __init__(self): + self.attempted_model_groups = [] + self.received_kwargs = None + + def log_retry(self, kwargs, e): + return kwargs + + async def async_function_with_fallbacks(self, *args, **kwargs): + self.attempted_model_groups.append(kwargs.get("model")) + self.received_kwargs = kwargs + return StreamingWrapper() + + +async def _acreate_batch(*args, **kwargs): + raise AssertionError("only used for its __name__") + + +@pytest.mark.asyncio +async def test_run_async_fallback_keeps_uploaded_file_requests_in_their_model_group(): + """An input_file_id only exists under the credentials of the group it was uploaded + to, so a cross-group fallback can only fail with the wrong provider's error.""" + router = AttemptRecordingRouter() + owning_provider_error = RuntimeError("openai connection error") + + with pytest.raises(RuntimeError, match="openai connection error"): + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=owning_provider_error, + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + input_file_id="file-owned-by-openai", + original_function=_acreate_batch, + ) + + assert router.attempted_model_groups == [] + + +@pytest.mark.asyncio +async def test_run_async_fallback_keeps_fine_tuning_requests_in_their_model_group(): + router = AttemptRecordingRouter() + + with pytest.raises(RuntimeError, match="openai connection error"): + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=RuntimeError("openai connection error"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + training_file="file-owned-by-openai", + ) + + assert router.attempted_model_groups == [] + + +@pytest.mark.asyncio +async def test_run_async_fallback_allows_same_model_group_retry_for_uploaded_file_requests(): + """Order-based fallbacks stay inside the owning group, so they must still run.""" + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=[{"model": "openai-group", "_target_order": 2}], + original_model_group="openai-group", + original_exception=RuntimeError("first deployment failed"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + input_file_id="file-owned-by-openai", + original_function=_acreate_batch, + ) + + assert router.attempted_model_groups == ["openai-group"] + + +@pytest.mark.asyncio +async def test_run_async_fallback_still_crosses_model_groups_without_an_uploaded_file(): + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=RuntimeError("openai connection error"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + ) + + assert router.attempted_model_groups == ["azure-group"] + + +@pytest.mark.asyncio +async def test_run_async_fallback_handles_explicitly_none_metadata(): + """/v1/batches always sets `metadata`, and sets it to None when the caller sent + none, so setdefault() on it hands back None instead of a dict.""" + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=RuntimeError("openai connection error"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + metadata=None, + ) + + assert router.received_kwargs["metadata"] == {"model_group": "azure-group"} + + +@pytest.mark.asyncio +async def test_run_async_fallback_records_batch_model_group_outside_provider_metadata(): + """`metadata` on a batch request is forwarded to the provider and stored on the + batch, so the router's own model_group belongs in litellm_metadata.""" + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=[{"model": "openai-group", "_target_order": 2}], + original_model_group="openai-group", + original_exception=RuntimeError("first deployment failed"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + input_file_id="file-owned-by-openai", + metadata={"caller": "nightly-job"}, + litellm_metadata={"model_group": "openai-group"}, + original_function=_acreate_batch, + ) + + assert router.received_kwargs["metadata"] == {"caller": "nightly-job"} + assert router.received_kwargs["litellm_metadata"]["model_group"] == "openai-group" + + def test_get_fallback_model_group_does_not_mutate_fallbacks(): """A string fallback must be resolved without mutating the caller's fallbacks list, which is the live router config shared across requests.""" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 4a3395a7d3f..b76e69bc978 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -6755,6 +6755,68 @@ async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error(): assert mock_create.call_args.kwargs["model"] == "owning-model" +@pytest.mark.asyncio +async def test_acreate_batch_surfaces_owning_provider_error_without_disable_fallbacks(): + """The router itself has to keep a batch inside the group that owns the input file: + the proxy only sets disable_fallbacks on the managed-files route, so the caller + otherwise gets the fallback provider's error for a file it never received.""" + from litellm.types.utils import LiteLLMBatch + + router = litellm.Router( + model_list=[ + { + "model_name": "owning-model", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-owning", + }, + }, + { + "model_name": "fallback-model", + "litellm_params": { + "model": "azure/gpt-4o-mini", + "api_key": "sk-fallback", + "api_base": "https://fallback.openai.azure.com", + "api_version": "2024-08-01-preview", + }, + }, + ], + fallbacks=[{"owning-model": ["fallback-model"]}], + num_retries=0, + ) + attempted_models = [] + + async def _acreate_batch(model, **kwargs): + attempted_models.append(model) + if model == "owning-model": + raise litellm.APIConnectionError( + message="Connection error - openai is unreachable", + model="openai/gpt-4o-mini", + llm_provider="openai", + ) + return LiteLLMBatch( + id="batch-created-on-the-wrong-provider", + completion_window="24h", + created_at=0, + endpoint="/v1/chat/completions", + input_file_id="file-owned-by-openai", + object="batch", + status="validating", + ) + + with patch.object(router, "_acreate_batch", _acreate_batch): + with pytest.raises(litellm.APIConnectionError, match="openai is unreachable"): + await router.acreate_batch( + model="owning-model", + input_file_id="file-owned-by-openai", + endpoint="/v1/chat/completions", + completion_window="24h", + metadata={"team": "batch-jobs"}, + ) + + assert attempted_models == ["owning-model"] + + @pytest.mark.asyncio async def test_acreate_batch_request_bedrock_tags_override_deployment_tags(): import httpx From d7bc63da5c5eb44daeaaa2a876f757257bb5d68a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 7 Aug 2026 05:25:48 -0700 Subject: [PATCH 010/139] style(router): drop the inline comment on the fallback metadata merge --- litellm/router_utils/fallback_event_handlers.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index c4a84a1d61e..00df20be845 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -169,7 +169,7 @@ async def run_async_fallback( kwargs[metadata_variable_name] = { **(kwargs.get(metadata_variable_name) or {}), "model_group": kwargs.get("model", None), - } # update model_group used, if fallbacks are done + } fallback_depth = fallback_depth + 1 kwargs["fallback_depth"] = fallback_depth kwargs["max_fallbacks"] = max_fallbacks From 855c49d0ef01a09161f8d2ec195be01447669f1a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 01:44:56 -0700 Subject: [PATCH 011/139] fix(proxy): skip prisma-dependent hooks when no database is attached --- .../storage_backend_service.py | 10 ++ litellm/proxy/utils.py | 5 + .../test_storage_backend_service.py | 127 ++++++++++++++++++ .../utils/proxy_logging/test_lifecycle.py | 94 +++++++++++-- 4 files changed, 228 insertions(+), 8 deletions(-) create mode 100644 tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py index 4c301c96f30..e766f335071 100644 --- a/litellm/proxy/openai_files_endpoints/storage_backend_service.py +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -68,6 +68,16 @@ class StorageBackendFileService: code=400, ) + if target_model_names: + managed_files_hook: Final = proxy_logging_obj.get_proxy_hook("managed_files") + if not isinstance(managed_files_hook, BaseFileEndpoints): + raise ProxyException( + message="Uploading with target_model_names requires a database-connected proxy, and this proxy has no database configured", + type="invalid_request_error", + param="target_model_names", + code=400, + ) + # Extract file information file_content: Final = file_data["content"] filename: Final = file_data.get("filename", "file") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e59c6adaf22..5f22ca021ac 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -544,6 +544,11 @@ class ProxyLogging: for hook in PROXY_HOOKS: proxy_hook = get_proxy_hook(hook) expected_args = inspect.getfullargspec(proxy_hook).args + if "prisma_client" in expected_args and prisma_client is None: + verbose_proxy_logger.debug( + "Skipping proxy hook %s: it requires a database and no prisma client is configured", hook + ) + continue passed_in_args: dict[str, Any] = {} if "internal_usage_cache" in expected_args: passed_in_args["internal_usage_cache"] = self.internal_usage_cache diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py new file mode 100644 index 00000000000..07a85a70815 --- /dev/null +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py @@ -0,0 +1,127 @@ +import pytest + +from litellm.llms.base_llm.files.transformation import BaseFileEndpoints +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.openai_files_endpoints import storage_backend_service +from litellm.proxy.openai_files_endpoints.storage_backend_service import ( + StorageBackendFileService, +) + + +class _RecordingStorageBackend: + def __init__(self): + self.upload_calls = [] + + async def upload_file(self, **kwargs): + self.upload_calls.append(kwargs) + return "https://storage.example/blob-1" + + +class _FakeManagedFilesHook(BaseFileEndpoints): + def __init__(self): + self.stored = [] + + async def acreate_file( + self, create_file_request, llm_router, target_model_names_list, litellm_parent_otel_span, user_api_key_dict + ): + raise NotImplementedError + + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router=None): + raise NotImplementedError + + async def afile_list(self, purpose, litellm_parent_otel_span, **data): + raise NotImplementedError + + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): + raise NotImplementedError + + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): + raise NotImplementedError + + async def store_unified_file_id(self, **kwargs): + self.stored.append(kwargs) + + +class _FakeProxyLogging: + def __init__(self, hook): + self._hook = hook + + def get_proxy_hook(self, hook_name): + return self._hook if hook_name == "managed_files" else None + + +def _file_data(): + return {"content": b"x", "filename": "input.jsonl", "content_type": "application/jsonl"} + + +@pytest.mark.asyncio +async def test_upload_with_target_model_names_but_no_hook_raises_before_uploading(monkeypatch): + backend = _RecordingStorageBackend() + monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend) + + with pytest.raises(ProxyException) as exc_info: + await StorageBackendFileService.upload_file_to_storage_backend( + file_data=_file_data(), + target_storage="azure_storage", + target_model_names=["gpt-x"], + purpose="batch", + proxy_logging_obj=_FakeProxyLogging(hook=None), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + snapshot = { + "code": exc_info.value.code, + "message_names_requirement": "requires a database-connected proxy" in exc_info.value.message, + "upload_calls": backend.upload_calls, + } + assert snapshot == {"code": "400", "message_names_requirement": True, "upload_calls": []} + + +@pytest.mark.asyncio +async def test_upload_without_target_model_names_skips_hook_requirement(monkeypatch): + backend = _RecordingStorageBackend() + monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend) + + file_object = await StorageBackendFileService.upload_file_to_storage_backend( + file_data=_file_data(), + target_storage="azure_storage", + target_model_names=[], + purpose="batch", + proxy_logging_obj=_FakeProxyLogging(hook=None), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + snapshot = { + "upload_count": len(backend.upload_calls), + "id_prefix": file_object.id.split("-")[0], + } + assert snapshot == {"upload_count": 1, "id_prefix": "file"} + + +@pytest.mark.asyncio +async def test_upload_with_target_model_names_and_hook_stores_unified_id(monkeypatch): + backend = _RecordingStorageBackend() + monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend) + hook = _FakeManagedFilesHook() + + file_object = await StorageBackendFileService.upload_file_to_storage_backend( + file_data=_file_data(), + target_storage="azure_storage", + target_model_names=["gpt-x"], + purpose="batch", + proxy_logging_obj=_FakeProxyLogging(hook=hook), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + snapshot = { + "upload_count": len(backend.upload_calls), + "store_count": len(hook.stored), + "stored_id_matches_response": hook.stored[0]["file_id"] == file_object.id, + "model_mappings": hook.stored[0]["model_mappings"], + } + assert snapshot == { + "upload_count": 1, + "store_count": 1, + "stored_id_matches_response": True, + "model_mappings": {"gpt-x": "https://storage.example/blob-1"}, + } diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py index e33da672599..cf906259246 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py @@ -8,7 +8,6 @@ because they are direct dependents on the lifecycle state. from __future__ import annotations -import asyncio from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock, patch @@ -17,7 +16,6 @@ import pytest import litellm from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.utils import ( - InternalUsageCache, ProxyLogging, ) @@ -102,9 +100,7 @@ def test_update_values_with_no_args_is_noop(proxy_logging): def test_update_values_invalid_type_for_alerting_raises(proxy_logging): - proxy_logging.slack_alerting_instance = MagicMock( - update_values=MagicMock(side_effect=TypeError("bad type")) - ) + proxy_logging.slack_alerting_instance = MagicMock(update_values=MagicMock(side_effect=TypeError("bad type"))) with pytest.raises(TypeError): proxy_logging.update_values(alerting={"not": "a list"}) # type: ignore[arg-type] @@ -190,6 +186,90 @@ def test_add_proxy_hooks_registers_callbacks(proxy_logging, monkeypatch): } +def _stub_hook_classes(): + class _PrismaFreeHook: + def __init__(self, internal_usage_cache): + self.internal_usage_cache = internal_usage_cache + + class _PrismaRequiringHook: + def __init__(self, internal_usage_cache, prisma_client): + self.internal_usage_cache = internal_usage_cache + self.prisma_client = prisma_client + + class _PrismaOnlyHook: + def __init__(self, prisma_client): + self.prisma_client = prisma_client + + return { + "cache_control_check": _PrismaFreeHook, + "needs_db_hook": _PrismaRequiringHook, + "db_only_hook": _PrismaOnlyHook, + } + + +def test_add_proxy_hooks_skips_prisma_requiring_hook_when_no_db(proxy_logging, monkeypatch): + hook_classes = _stub_hook_classes() + registered: List[Any] = [] + + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", list(hook_classes.keys())) + monkeypatch.setattr(utils_mod, "get_proxy_hook", hook_classes.__getitem__) + monkeypatch.setattr( + litellm.logging_callback_manager, + "add_litellm_callback", + lambda cb: registered.append(cb), + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + proxy_logging._add_proxy_hooks(llm_router=None) + + snapshot = { + "mapping_keys": list(proxy_logging.proxy_hook_mapping.keys()), + "registered_types": [type(r).__name__ for r in registered], + "needs_db_hook_lookup": proxy_logging.get_proxy_hook("needs_db_hook"), + "db_only_hook_lookup": proxy_logging.get_proxy_hook("db_only_hook"), + } + assert snapshot == { + "mapping_keys": ["cache_control_check"], + "registered_types": ["_PrismaFreeHook"], + "needs_db_hook_lookup": None, + "db_only_hook_lookup": None, + } + + +def test_add_proxy_hooks_registers_prisma_requiring_hook_with_db(proxy_logging, monkeypatch): + hook_classes = _stub_hook_classes() + registered: List[Any] = [] + fake_prisma = MagicMock() + + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", list(hook_classes.keys())) + monkeypatch.setattr(utils_mod, "get_proxy_hook", hook_classes.__getitem__) + monkeypatch.setattr( + litellm.logging_callback_manager, + "add_litellm_callback", + lambda cb: registered.append(cb), + ) + + with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): + proxy_logging._add_proxy_hooks(llm_router=None) + + snapshot = { + "mapping_keys": list(proxy_logging.proxy_hook_mapping.keys()), + "registered_count": len(registered), + "needs_db_hook_got_prisma": proxy_logging.proxy_hook_mapping["needs_db_hook"].prisma_client is fake_prisma, + "db_only_hook_got_prisma": proxy_logging.proxy_hook_mapping["db_only_hook"].prisma_client is fake_prisma, + } + assert snapshot == { + "mapping_keys": ["cache_control_check", "needs_db_hook", "db_only_hook"], + "registered_count": 3, + "needs_db_hook_got_prisma": True, + "db_only_hook_got_prisma": True, + } + + def test_add_proxy_hooks_unknown_hook_raises(proxy_logging, monkeypatch): from litellm.proxy import utils as utils_mod @@ -267,9 +347,7 @@ def test_init_litellm_callbacks_replaces_string_with_instance(proxy_logging, mon snapshot = { "replaced_first_item": litellm.callbacks[0] is sentinel_instance, "callbacks_grew_with_service": len(litellm.callbacks) >= 2, - "service_logging_appended": any( - "ServiceLogging" in type(c).__name__ for c in litellm.callbacks - ), + "service_logging_appended": any("ServiceLogging" in type(c).__name__ for c in litellm.callbacks), } assert snapshot == { "replaced_first_item": True, From 8812debeff43b077c64c2cbf2f2d2dd4e3187639 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 20:58:30 +0000 Subject: [PATCH 012/139] docs: allow functional comments as an exception in CLAUDE.md Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index f1bb46c1fd3..046932274cd 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,4 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt +Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. Exceptions are granted for comments that do something rather than document something for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From abbad8ad528fb6a36c037ad90e9a210fd3435753 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 21:04:43 +0000 Subject: [PATCH 013/139] docs: limit the comment exception to tool-read directives Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index 046932274cd..4901346bc6e 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,4 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. Exceptions are granted for comments that do something rather than document something for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed +Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The one exception is a comment a tool reads and acts on, as opposed to one documenting code for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed. Human-readable annotations like TODO, FIXME, and section headers don't qualify Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From 6e6e0d662b4db1d46552ca011b19b5775f0db2ad Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 21:07:40 +0000 Subject: [PATCH 014/139] docs: frame the comment rule around AI slop and allow TODO/FIXME Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index 4901346bc6e..b249e2d01fc 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,4 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The one exception is a comment a tool reads and acts on, as opposed to one documenting code for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed. Human-readable annotations like TODO, FIXME, and section headers don't qualify +Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The point of that rule is to keep out AI slop comments that just restate what the code already says, so comments carrying information the code can't are fine. That covers comments a tool reads and acts on, such as an entry in `.git-blame-ignore-revs` saying which commit is excluded from git blame, or a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a violation is truly unavoidable, and it covers a real TODO or FIXME flagging known unfinished work. Write those, and the reasons they require, wherever they're needed Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From 1f8964a4c4ab2d77eae31ee2a9bf8de2ae6c0f14 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 14:20:13 -0700 Subject: [PATCH 015/139] chore: handwrite the rule --- CLAUDE.md | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index b249e2d01fc..6990408e686 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,12 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The point of that rule is to keep out AI slop comments that just restate what the code already says, so comments carrying information the code can't are fine. That covers comments a tool reads and acts on, such as an entry in `.git-blame-ignore-revs` saying which commit is excluded from git blame, or a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a violation is truly unavoidable, and it covers a real TODO or FIXME flagging known unfinished work. Write those, and the reasons they require, wherever they're needed +Do not write comments unless they are: +- absolutely necessary to explain some very complex business logic +- used as an input for tools to read and act on. For example: + - entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame + - a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a truly unavoidable violation +- a TODO or FIXME + - Not great to have those, but if it's unavoidable, make sure to include a strong reason for why it's there or, better yet, link to a GitHub issue for the follow-up work + +Explanation: code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From 805fc497761a47c82ceb062b10945434e37cc53b Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 14:29:36 -0700 Subject: [PATCH 016/139] chore: mention AI slop reason --- CLAUDE.md | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 6990408e686..dcdfe2d15b9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,12 +1,12 @@ Do not write comments unless they are: -- absolutely necessary to explain some very complex business logic +- absolutely necessary to explain some very complex business logic (in which case, keep it concise and clear) - used as an input for tools to read and act on. For example: - entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame - a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a truly unavoidable violation - a TODO or FIXME - - Not great to have those, but if it's unavoidable, make sure to include a strong reason for why it's there or, better yet, link to a GitHub issue for the follow-up work + - Not great to have those, but if it's unavoidable, make sure to include a strong, concise reason for why it's there or, better yet, link to a GitHub issue for the follow-up work -Explanation: code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance +Explanation: The point of this rule is to keep out AI slop comments. AI writes way too many and way too verbose comments. Code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From efc4e6f28c0951dc20fc61e8cb9be53834fdb7b7 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Sat, 8 Aug 2026 16:01:47 -0700 Subject: [PATCH 017/139] fix(batches): keep batch state in sync on a poll without claiming attribution (#34456) A poll of a Vertex passthrough batch wrote nothing to the managed-object row, so status and file_object stayed frozen at the create-time snapshot and GET /v1/batches served a stale status and an empty output file id for the life of the batch. Only the create may claim a batch, but every observation of one may refresh its state. store_unified_object_id takes create_if_missing, which the poll clears: it refreshes status and file_object through update_many, and leaves a row that is absent absent rather than creating one owned by the observer, since created_by and team_id are written by whoever reaches the create branch. The update payload is now shared with the upsert so it cannot drift into writing api_key, request_tags, created_by or team_id. The passthrough identity re-assertion that was previously part of this PR ships separately in #36121, so this PR keeps only the batch attribution work. The creating key owns user_api_key_alias only when it actually has one. Guarding the overwrite on the presence of a key rather than on a resolved alias nulled the field out for every key generated without key_alias, and for any key rotated or deleted before its batch finished, losing the creating user's alias that the spend row previously carried. The guard now matches the team-alias line below it. --- .../proxy/common_utils/check_batch_cost.py | 77 +++++++-- .../proxy/hooks/managed_files.py | 49 +++++- .../migration.sql | 5 + .../litellm_proxy_extras/schema.prisma | 2 + .../proxy/hooks/proxy_track_cost_callback.py | 23 ++- .../vertex_passthrough_logging_handler.py | 76 ++++++++- litellm/proxy/schema.prisma | 2 + schema.prisma | 2 + .../proxy_unit_tests/test_check_batch_cost.py | 150 ++++++++++++++++ ..._batch_update_db_managed_output_file_id.py | 139 +++++++++++++++ .../proxy/test_managed_files_hook.py | 111 ++++++++++++ .../hooks/test_proxy_track_cost_callback.py | 161 ++++++++++++++++++ .../test_vertex_ai_batch_passthrough.py | 136 +++++++++++++++ 13 files changed, 906 insertions(+), 27 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 7acdd5dbdaf..dc8f17fb665 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -3,7 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t """ from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -43,11 +43,15 @@ class CheckBatchCost: # the guaranteed-failing primary query on every subsequent cycle. self._has_batch_processed_column: bool = True - async def _get_user_info(self, batch_id, user_id) -> dict: + async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> Dict[str, Any]: """ Look up user email and key alias by user_id for enriching the S3 callback metadata. Returns a dict with user_api_key_user_email and user_api_key_alias (both may be None). + Returns an empty dict when user_id is None: batches created by a team or service + account key carry no user id, and find_unique(where={"user_id": None}) raises. """ + if not user_id: + return {} try: user_row = await self.prisma_client.db.litellm_usertable.find_unique( where={"user_id": user_id} @@ -62,6 +66,66 @@ class CheckBatchCost: verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}") return {} + async def _get_key_alias(self, batch_id: str, api_key: str | None) -> str | None: + """Resolve the creating virtual key's alias from its hashed token.""" + if not api_key: + return None + try: + key_row = await self.prisma_client.db.litellm_verificationtoken.find_unique( + where={"token": api_key} + ) + return getattr(key_row, "key_alias", None) if key_row is not None else None + except Exception as e: + verbose_proxy_logger.error(f"CheckBatchCost: could not look up key alias for batch {batch_id}: {e}") + return None + + async def _get_team_alias(self, team_id: str | None) -> str | None: + """Resolve a team's alias from its id.""" + if not team_id: + return None + try: + team_row = await self.prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + return getattr(team_row, "team_alias", None) if team_row is not None else None + except Exception as e: + verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}") + return None + + async def _build_creator_attribution_metadata( + self, job: "LiteLLM_ManagedObjectTable", batch_id: str + ) -> Dict[str, Any]: + """ + Rebuild the spend-tracking metadata for the key, team, and tags that created the + batch so the batch-cost spend log is attributed the same way a non-batch request + is. Rows created before api_key and request_tags were persisted carry only + created_by and team_id, and fall back to those. A named creating key owns + user_api_key_alias; when it has no alias, or the key has since been rotated or + deleted, the field keeps the creating user's alias that _get_user_info filled in, + because a resolvable name is more useful on the spend row than a null. + """ + api_key = getattr(job, "api_key", None) + team_id = getattr(job, "team_id", None) + request_tags = getattr(job, "request_tags", None) + + metadata: Dict[str, Any] = { + "user_api_key_user_id": job.created_by, + "user_api_key": api_key, + "user_api_key_team_id": team_id, + **(await self._get_user_info(batch_id, job.created_by)), + } + + key_alias = await self._get_key_alias(batch_id, api_key) + if key_alias is not None: + metadata["user_api_key_alias"] = key_alias + team_alias = await self._get_team_alias(team_id) + if team_alias is not None: + metadata["user_api_key_team_alias"] = team_alias + if isinstance(request_tags, list) and request_tags: + metadata["tags"] = [tag for tag in request_tags if isinstance(tag, str)] + + return metadata + async def _cleanup_stale_managed_objects(self) -> None: """ Mark managed objects older than MANAGED_OBJECT_STALENESS_CUTOFF_DAYS days @@ -485,9 +549,6 @@ class CheckBatchCost: function_id=str(uuid.uuid4()), ) - creator_user_id = job.created_by - user_info = await self._get_user_info(batch_id, job.created_by) - logging_obj.update_environment_variables( litellm_params={ # set the user-agent header so that S3 callback consumers can easily identify CheckBatchCost callbacks @@ -496,11 +557,7 @@ class CheckBatchCost: "user-agent": CHECK_BATCH_COST_USER_AGENT, } }, - "metadata": { - "user_api_key_user_id": creator_user_id, - "user_api_key_team_id": getattr(job, "team_id", None), - **user_info, - }, + "metadata": await self._build_creator_attribution_metadata(job, batch_id), }, optional_params={}, ) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index f0914240f79..a3994dccfd6 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -163,6 +163,8 @@ class _ManagedObjectTableActions(Protocol): self, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]] ) -> "PrismaManagedObjectRow": ... + async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... + class _CursorPageArgs(TypedDict, total=False): cursor: Mapping[str, str] @@ -263,7 +265,24 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): model_object_id: str, file_purpose: Literal["batch", "fine-tune", "response"], user_api_key_dict: UserAPIKeyAuth, + request_tags: Sequence[str] | None = None, + persist_attribution: bool = False, + create_if_missing: bool = True, ) -> None: + """Persist a managed object row, caching it and upserting it in the DB. + + persist_attribution is set only by the batch create, which is the one caller + that can speak for the creator; it gates the api_key and request_tags columns + that CheckBatchCost bills against, so a later poll or retrieve of the same + batch cannot record itself as the paying key. Like created_by and team_id, + both are written only in the upsert create branch, never on update. + + create_if_missing is cleared by callers that observe a batch they did not + create, such as a poll. They still refresh status and file_object, but a + row absent from the table is left absent rather than created with the + observer as its creator, because created_by and team_id are written from + whoever calls the create branch. + """ verbose_logger.info(f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache") litellm_managed_object = LiteLLM_ManagedObjectTable( unified_object_id=unified_object_id, @@ -277,6 +296,29 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): litellm_parent_otel_span=litellm_parent_otel_span, ) + from prisma import Json + + api_key = user_api_key_dict.api_key or None + attribution_columns = ( + { + **({"api_key": api_key} if api_key is not None else {}), + **({"request_tags": Json(list(request_tags))} if request_tags else {}), + } + if persist_attribution + else {} + ) + # FIX: Update status and file_object on every operation to keep state in sync + update_columns: Final = { + "file_object": file_object.model_dump_json(), + "status": file_object.status, + "updated_by": user_api_key_dict.user_id, + } + if not create_if_missing: + await _managed_object_table(self.prisma_client).update_many( + where={"unified_object_id": unified_object_id}, + data=update_columns, + ) + return await _managed_object_table(self.prisma_client).upsert( where={"unified_object_id": unified_object_id}, data={ @@ -289,12 +331,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): "team_id": user_api_key_dict.team_id, "updated_by": user_api_key_dict.user_id, "status": file_object.status, + **attribution_columns, }, - "update": { - "file_object": file_object.model_dump_json(), - "status": file_object.status, - "updated_by": user_api_key_dict.user_id, - }, # FIX: Update status and file_object on every operation to keep state in sync + "update": update_columns, }, ) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql new file mode 100644 index 00000000000..79bc6b24de8 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql @@ -0,0 +1,5 @@ +-- Add api_key and request_tags columns to LiteLLM_ManagedObjectTable +-- Captured at batch-create time so CheckBatchCost can attribute batch-cost spend +-- back to the creating virtual key (and its tags) even when created_by is null. +ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "api_key" TEXT; +ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "request_tags" JSONB DEFAULT '[]'; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 9c871b65f40..cabddf6f1a1 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -985,6 +985,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t created_at DateTime @default(now()) created_by String? team_id String? + api_key String? + request_tags Json? @default("[]") updated_at DateTime @updatedAt updated_by String? diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 3346f9d7e3b..0e22b5324c1 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -48,6 +48,15 @@ _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset( } ) +# Both spellings, because call_type reaches the callback as str(...) of either the +# enum member or its value. +_CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset( + ( + CallTypes.aretrieve_batch.value, + str(CallTypes.aretrieve_batch), + ) +) + class _ProxyDBLogger(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -212,7 +221,10 @@ class _ProxyDBLogger(CustomLogger): # Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls). # Avoids a cache/DB lookup on every normal LLM request. if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): - metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) + metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original + metadata=metadata, + resolve_missing_key_identity=str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES, + ) _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata) user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None)) @@ -337,7 +349,7 @@ class _ProxyDBLogger(CustomLogger): spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) @staticmethod - async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: + async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict: """ Enriches failure spend log metadata by looking up the key object (and team object) from cache/DB when key fields are missing. @@ -349,6 +361,11 @@ class _ProxyDBLogger(CustomLogger): 2. Post-auth failures (provider errors, rate limits): key fields are populated but team_alias is missing because LiteLLM_VerificationTokenView SQL view doesn't include it. We look up the team object to fill in team_alias. + + Scenario 1 reads the key's identity as it stands right now, so it is only correct + for a log emitted within the request it describes. Callers that log after a delay, + against an identity captured earlier, pass resolve_missing_key_identity=False and + keep their own user_id, team_id and org_id. """ api_key_hash: Final = metadata.get("user_api_key") if not api_key_hash: @@ -361,7 +378,7 @@ class _ProxyDBLogger(CustomLogger): ) # Step 1: If key fields are missing, look up the full key object - if metadata.get("user_api_key_alias") is None: + if resolve_missing_key_identity and metadata.get("user_api_key_alias") is None: try: key_obj: Final = await get_key_object( hashed_token=api_key_hash, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 2e3f7bb9aa6..9c5b7dc563e 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -1,4 +1,6 @@ +import asyncio import re +from collections.abc import Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Final, cast from urllib.parse import urlparse @@ -39,6 +41,32 @@ else: EndpointType = Any +def _optional_str(value: object) -> str | None: + return value if isinstance(value, str) else None + + +def _optional_str_tuple(value: object) -> tuple[str, ...] | None: + if not isinstance(value, list): + return None + items: Final = cast(list[object], value) # cast-ok: isinstance-narrowed; element type unknown + return tuple(tag for tag in items if isinstance(tag, str)) + + +def _request_tags(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None: + """Tags for the batch-cost spend row: the request's own tags when it sent any, + otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a + tagged key does not put its tags in the top-level metadata "tags" on the + passthrough path) + """ + tags: Final = _optional_str_tuple(request_metadata.get("tags")) + if tags: + return tags + key_auth_metadata: Final = request_metadata.get("user_api_key_auth_metadata") + if isinstance(key_auth_metadata, dict): + return _optional_str_tuple(key_auth_metadata.get("tags")) + return None + + class VertexPassthroughLoggingHandler: @staticmethod def vertex_passthrough_handler( @@ -657,11 +685,13 @@ class VertexPassthroughLoggingHandler: # Store the managed object for cost tracking # This will be picked up by check_batch_cost polling mechanism + is_batch_create: Final = url_route.split("?")[0].rstrip("/").endswith("batchPredictionJobs") VertexPassthroughLoggingHandler._store_batch_managed_object( unified_object_id=unified_object_id, batch_object=litellm_batch_response, model_object_id=batch_id, logging_obj=logging_obj, + is_batch_create=is_batch_create, **kwargs, ) @@ -779,17 +809,45 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + @staticmethod + def _log_batch_registration_result( + finished: asyncio.Task, unified_object_id: str, model_object_id: str, is_batch_create: bool + ) -> None: + error: Final = finished.exception() if not finished.cancelled() else None + if finished.cancelled() or error is not None: + consequence: Final = ( + "its cost will not be tracked" if is_batch_create else "its status and output file may be stale" + ) + verbose_proxy_logger.error( + "Failed to store batch managed object with unified_object_id=%s, batch_id=%s; %s: %s", + unified_object_id, + model_object_id, + consequence, + error, + ) + return + verbose_proxy_logger.info( + "Stored batch managed object with unified_object_id=%s, batch_id=%s", + unified_object_id, + model_object_id, + ) + @staticmethod def _store_batch_managed_object( unified_object_id: str, batch_object: LiteLLMBatch, model_object_id: str, logging_obj: LiteLLMLoggingObj, + is_batch_create: bool, **kwargs, ) -> None: """ Store batch managed object for cost tracking. This will be picked up by the check_batch_cost polling mechanism. + + A poll refreshes the batch status and file object but neither creates the row + nor writes attribution, so the creating key and its tags are persisted from + the create alone. """ try: # Get the managed files hook from the logging object @@ -805,7 +863,7 @@ class VertexPassthroughLoggingHandler: user_api_key_dict: Final = UserAPIKeyAuth( user_id=_request_metadata.get("user_api_key_user_id", "default-user"), - api_key="", + api_key=_optional_str(_request_metadata.get("user_api_key")), team_id=_request_metadata.get("user_api_key_team_id"), team_alias=None, user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value @@ -827,9 +885,7 @@ class VertexPassthroughLoggingHandler: ) # Store the unified object for batch cost tracking - import asyncio - - asyncio.create_task( + task: Final = asyncio.create_task( managed_files_hook.store_unified_object_id( unified_object_id=unified_object_id, file_object=batch_object, @@ -837,13 +893,15 @@ class VertexPassthroughLoggingHandler: model_object_id=model_object_id, file_purpose="batch", user_api_key_dict=user_api_key_dict, + request_tags=_request_tags(_request_metadata), + persist_attribution=is_batch_create, + create_if_missing=is_batch_create, ) ) - - verbose_proxy_logger.info( - "Stored batch managed object with unified_object_id=%s, batch_id=%s", - unified_object_id, - model_object_id, + task.add_done_callback( + lambda finished: VertexPassthroughLoggingHandler._log_batch_registration_result( + finished, unified_object_id, model_object_id, is_batch_create + ) ) else: verbose_proxy_logger.warning( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 9c871b65f40..cabddf6f1a1 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -985,6 +985,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t created_at DateTime @default(now()) created_by String? team_id String? + api_key String? + request_tags Json? @default("[]") updated_at DateTime @updatedAt updated_by String? diff --git a/schema.prisma b/schema.prisma index 9c871b65f40..cabddf6f1a1 100644 --- a/schema.prisma +++ b/schema.prisma @@ -985,6 +985,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t created_at DateTime @default(now()) created_by String? team_id String? + api_key String? + request_tags Json? @default("[]") updated_at DateTime @updatedAt updated_by String? diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 800c97d7ba1..6d7ada17ec5 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -1641,3 +1641,153 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: decoded = _is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] +class TestBatchCostAttribution: + """CheckBatchCost rebuilds the creator's spend metadata from the managed-object row so + the batch-cost log is attributed like a non-batch request.""" + + def _instance(self, key_row=None, team_row=None, user_row=None): + from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost + + prisma = MagicMock() + prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_row) + prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + return CheckBatchCost( + proxy_logging_obj=MagicMock(), + prisma_client=prisma, + llm_router=MagicMock(), + ) + + def _job(self, **overrides): + from types import SimpleNamespace + + fields = { + "created_by": "alice", + "team_id": "team-alpha", + "api_key": "hash-alice", + "request_tags": ["env:prod"], + } + fields.update(overrides) + return SimpleNamespace(unified_object_id="uoi", **fields) + + @pytest.mark.asyncio + async def test_metadata_carries_key_team_and_tags(self): + """The spend row names the creating key, its team, both aliases, and the tags.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=SimpleNamespace(key_alias="prod-key"), + team_row=SimpleNamespace(team_alias="Team Alpha"), + user_row=SimpleNamespace(user_email="alice@example.com", user_alias=None), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key"] == "hash-alice" + assert metadata["user_api_key_user_id"] == "alice" + assert metadata["user_api_key_team_id"] == "team-alpha" + assert metadata["user_api_key_alias"] == "prod-key" + assert metadata["user_api_key_team_alias"] == "Team Alpha" + assert metadata["tags"] == ["env:prod"] + + @pytest.mark.asyncio + async def test_metadata_tolerates_legacy_row_without_columns(self): + """Rows created before the columns existed carry only created_by/team_id and must + still produce an attributed row rather than raising.""" + instance = self._instance() + job = self._job(api_key=None, request_tags=None) + + metadata = await instance._build_creator_attribution_metadata(job, "batch-1") + + assert metadata["user_api_key"] is None + assert metadata["user_api_key_user_id"] == "alice" + assert metadata["user_api_key_team_id"] == "team-alpha" + assert "tags" not in metadata + + @pytest.mark.asyncio + async def test_metadata_keeps_key_when_team_key_has_no_user(self): + """A team-scoped key carries no user id. The user lookup is skipped (prisma rejects + a None user_id) and the key hash still drives key-level attribution.""" + from types import SimpleNamespace + + instance = self._instance(key_row=SimpleNamespace(key_alias="svc-key")) + job = self._job(created_by=None) + + metadata = await instance._build_creator_attribution_metadata(job, "batch-1") + + assert metadata["user_api_key"] == "hash-alice" + assert metadata["user_api_key_user_id"] is None + assert metadata["user_api_key_alias"] == "svc-key" + instance.prisma_client.db.litellm_usertable.find_unique.assert_not_called() + + @pytest.mark.asyncio + async def test_metadata_drops_non_string_tags(self): + """Non-string tags are dropped so a malformed stored value cannot slip past the + tag-budget checks that consume this metadata.""" + instance = self._instance() + job = self._job(request_tags=["env:prod", 7, None, "team:ml"]) + + metadata = await instance._build_creator_attribution_metadata(job, "batch-1") + + assert metadata["tags"] == ["env:prod", "team:ml"] + + @pytest.mark.asyncio + async def test_key_alias_lookup_failure_does_not_break_attribution(self): + """An alias lookup failure must not lose the spend row; the key hash and team still + attribute it.""" + instance = self._instance() + instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + side_effect=Exception("db down") + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key"] == "hash-alice" + assert metadata.get("user_api_key_alias") is None + + @pytest.mark.asyncio + async def test_unnamed_key_keeps_the_creating_user_alias(self): + """Regression: a key generated without key_alias resolves to no alias, and the + overwrite must not null out the creating user's alias that _get_user_info supplied. + Most keys carry no alias, so this is the common batch, not an edge case.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=SimpleNamespace(key_alias=None), + user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key_alias"] == "Alice Chen" + assert metadata["user_api_key"] == "hash-alice" + + @pytest.mark.asyncio + async def test_rotated_key_keeps_the_creating_user_alias(self): + """Batches outlive keys. When the creating key has been rotated or deleted the + lookup returns no row, and the spend log keeps a resolvable name instead of null.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=None, + user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key_alias"] == "Alice Chen" + + @pytest.mark.asyncio + async def test_named_key_still_owns_the_alias(self): + """The fallback must not weaken the intended precedence: a key that has its own + alias still overrides the creating user's.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=SimpleNamespace(key_alias="prod-key"), + user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key_alias"] == "prod-key" diff --git a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py index 1b60c97b510..ebd33aa2e53 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -356,3 +356,142 @@ async def test_ensure_batch_response_returns_early_without_auth(): assert response.output_file_id == "file-raw-output" mock_managed_files.get_unified_output_file_id.assert_not_called() + + +def _in_memory_managed_files(): + """Build a real _PROXY_LiteLLMManagedFiles whose prisma upsert hits an in-memory row.""" + from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + + store: dict = {} + + async def _upsert(where, data): + key = where["unified_object_id"] + if key in store: + store[key].update(data["update"]) + else: + store[key] = dict(data["create"]) + + table = MagicMock() + table.upsert = AsyncMock(side_effect=_upsert) + prisma = MagicMock() + prisma.db.litellm_managedobjecttable = table + + cache = MagicMock() + cache.async_set_cache = AsyncMock() + + return ( + _PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma), + store, + ) + + +@pytest.mark.asyncio +async def test_store_unified_object_id_persists_key_and_tags_on_create(): + """Regression (spend loss): the batch create persists the creating key hash and tags so + CheckBatchCost can write an attributed spend row instead of a blank one the DB drops.""" + instance, store = _in_memory_managed_files() + creator = UserAPIKeyAuth(user_id="alice", team_id="team-alpha", api_key="hash-alice") + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="validating"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=["env:prod"], + persist_attribution=True, + ) + + row = store["unified-b"] + assert row["api_key"] == "hash-alice" + assert row["created_by"] == "alice" + assert row["team_id"] == "team-alpha" + assert row["request_tags"].data == ["env:prod"] + + +@pytest.mark.asyncio +async def test_store_unified_object_id_omits_key_and_tags_without_persist_attribution(): + """Regression (spend redirect): a caller that is not the batch create (a poll, or the + generic post-call hook on a retrieve) carries a real hashed key, but must never have it + recorded as the batch's paying key. created_by/team_id keep their existing behavior.""" + instance, store = _in_memory_managed_files() + poller = UserAPIKeyAuth(user_id="bob", team_id="team-bravo", api_key="hash-bob") + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="in_progress"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=poller, + request_tags=["env:dev"], + ) + + row = store["unified-b"] + assert "api_key" not in row + assert "request_tags" not in row + assert row["created_by"] == "bob" + + +@pytest.mark.asyncio +async def test_store_unified_object_id_attribution_columns_are_write_once(): + """Identity is written only in the upsert create branch, so a later store for the same + batch (a status update, a poll) can neither reassign the paying key nor clear it.""" + instance, store = _in_memory_managed_files() + creator = UserAPIKeyAuth(user_id="alice", team_id="team-alpha", api_key="hash-alice") + poller = UserAPIKeyAuth(user_id="bob", team_id="team-bravo", api_key="hash-bob") + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="validating"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=["env:prod"], + persist_attribution=True, + ) + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="completed"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=poller, + request_tags=["poller-tag"], + persist_attribution=True, + ) + + row = store["unified-b"] + assert row["api_key"] == "hash-alice" + assert row["created_by"] == "alice" + assert row["status"] == "completed" + + upsert_data = instance.prisma_client.db.litellm_managedobjecttable.upsert.call_args.kwargs["data"] + assert "api_key" not in upsert_data["update"] + assert "request_tags" not in upsert_data["update"] + + +@pytest.mark.asyncio +async def test_store_unified_object_id_omits_unset_columns(): + """A batch created with no tags (the common case) still registers: the optional columns + are omitted rather than passed as None, which prisma rejects for the Json column.""" + instance, store = _in_memory_managed_files() + creator = UserAPIKeyAuth(user_id="alice", team_id="team-alpha", api_key=None) + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="validating"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=None, + persist_attribution=True, + ) + + create_data = instance.prisma_client.db.litellm_managedobjecttable.upsert.call_args.kwargs["data"]["create"] + assert "api_key" not in create_data + assert "request_tags" not in create_data + assert "unified-b" in store diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 63646ae53f8..34cd0cabc2c 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -510,6 +510,117 @@ def _make_real_managed_files_instance(): ) +def _make_object_store_instance(): + """A real store_unified_object_id over an AsyncMock prisma client, so both the + upsert and the update-only write path can be asserted.""" + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + mock_cache = MagicMock() + mock_cache.async_set_cache = AsyncMock() + + mock_prisma = MagicMock() + mock_prisma.db.litellm_managedobjecttable.upsert = AsyncMock() + mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock() + + return ( + _PROXY_LiteLLMManagedFiles( + internal_usage_cache=mock_cache, + prisma_client=mock_prisma, + ), + mock_prisma, + ) + + +@pytest.mark.asyncio +async def test_poll_refreshes_batch_state_without_claiming_the_row(): + """Regression (stale batch state): a poll observes a batch it did not create, so it + must still refresh status and file_object -- otherwise GET /v1/batches serves the + create-time snapshot forever -- while writing none of the attribution columns and + never creating a row it would then own.""" + managed_files, mock_prisma = _make_object_store_instance() + poller = UserAPIKeyAuth( + api_key="sk-the-poller", user_id="bob", team_id="team-bravo", parent_otel_span=None + ) + + await managed_files.store_unified_object_id( + unified_object_id="uoi-1", + file_object=_make_batch_response(status="completed"), + litellm_parent_otel_span=None, + model_object_id="batch-123", + file_purpose="batch", + user_api_key_dict=poller, + request_tags=("poller:tag",), + persist_attribution=False, + create_if_missing=False, + ) + + # the row is refreshed in place, and cannot be conjured by a poll + mock_prisma.db.litellm_managedobjecttable.upsert.assert_not_awaited() + update_many = mock_prisma.db.litellm_managedobjecttable.update_many + update_many.assert_awaited_once() + call = update_many.await_args + assert call.kwargs["where"] == {"unified_object_id": "uoi-1"} + + written = call.kwargs["data"] + assert written["status"] == "completed" + assert json.loads(written["file_object"])["output_file_id"] == "file-output-abc" + # nothing the poller could be billed for + for owned in ("api_key", "request_tags", "created_by", "team_id"): + assert owned not in written + + +@pytest.mark.asyncio +async def test_create_still_upserts_and_claims_attribution(): + """The create is the one caller that can speak for the batch, so it keeps the upsert + (creating the row when absent) and writes the attribution columns.""" + managed_files, mock_prisma = _make_object_store_instance() + creator = UserAPIKeyAuth( + api_key="sk-the-creator", user_id="alice", team_id="team-alpha", parent_otel_span=None + ) + + await managed_files.store_unified_object_id( + unified_object_id="uoi-2", + file_object=_make_batch_response(status="validating"), + litellm_parent_otel_span=None, + model_object_id="batch-456", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=("env:prod",), + persist_attribution=True, + ) + + mock_prisma.db.litellm_managedobjecttable.update_many.assert_not_awaited() + upsert = mock_prisma.db.litellm_managedobjecttable.upsert + upsert.assert_awaited_once() + created = upsert.await_args.kwargs["data"]["create"] + # UserAPIKeyAuth hashes an sk- token on construction; the hash is what is billed + assert created["api_key"] == creator.api_key + assert created["api_key"] != "sk-the-creator" + assert created["created_by"] == "alice" + assert created["team_id"] == "team-alpha" + + +@pytest.mark.asyncio +async def test_default_callers_still_create_their_rows(): + """create_if_missing defaults to True, so the fine-tune, Responses and Anthropic + callers, none of which pass it, keep upserting exactly as before.""" + managed_files, mock_prisma = _make_object_store_instance() + + await managed_files.store_unified_object_id( + unified_object_id="uoi-3", + file_object=_make_batch_response(), + litellm_parent_otel_span=None, + model_object_id="ft-789", + file_purpose="fine-tune", + user_api_key_dict=_make_user_api_key_dict(), + ) + + mock_prisma.db.litellm_managedobjecttable.upsert.assert_awaited_once() + mock_prisma.db.litellm_managedobjecttable.update_many.assert_not_awaited() + + @pytest.mark.asyncio async def test_store_unified_file_id_is_idempotent_via_upsert(): """Regression test for the managed-batch retrieve 500 (UniqueViolationError on diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 69f04ce2bbe..2b162774aea 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -17,6 +17,7 @@ from litellm.proxy.hooks.proxy_track_cost_callback import ( _should_track_cost_callback, _update_database_and_spend_counters, ) +from litellm.types.utils import CallTypes @pytest.mark.asyncio @@ -783,6 +784,166 @@ async def test_enrich_failure_metadata_skips_when_no_api_key(): mock_get_key.assert_not_called() +@pytest.mark.asyncio +async def test_enrich_failure_metadata_keeps_captured_identity_when_not_resolving(): + """ + With resolve_missing_key_identity=False the key is not read, so a null user_id, + team_id and org_id captured earlier stay null instead of being refilled from the + key as it stands now. The team_alias lookup still runs off the captured team_id. + """ + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "alias-assigned-later" + mock_key_obj.user_id = "user-assigned-later" + mock_key_obj.team_id = "team-assigned-later" + mock_key_obj.org_id = "org-assigned-later" + + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "captured-team-alias" + + with ( + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ) as mock_get_key, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ), + ): + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": None, + "user_api_key_user_id": None, + "user_api_key_team_id": "captured-team-id", + "user_api_key_team_alias": None, + "user_api_key_org_id": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata, resolve_missing_key_identity=False + ) + + mock_get_key.assert_not_called() + assert result["user_api_key_user_id"] is None + assert result["user_api_key_team_id"] == "captured-team-id" + assert result["user_api_key_org_id"] is None + assert result["user_api_key_alias"] is None + assert result["user_api_key_team_alias"] == "captured-team-alias" + + +@pytest.mark.asyncio +async def test_enrich_failure_metadata_ignores_flag_when_alias_present(): + """ + A captured alias already closes the key lookup, so resolve_missing_key_identity + changes nothing for a key that has one; only the alias-less key depends on it. + """ + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "captured-team-alias" + + with ( + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + ) as mock_get_key, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ), + ): + for resolve in (True, False): + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": "captured-alias", + "user_api_key_user_id": None, + "user_api_key_team_id": "captured-team-id", + "user_api_key_team_alias": None, + "user_api_key_org_id": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata, resolve_missing_key_identity=resolve + ) + mock_get_key.assert_not_called() + assert result["user_api_key_user_id"] is None + assert result["user_api_key_alias"] == "captured-alias" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_type, expect_key_read", + [ + (CallTypes.aretrieve_batch.value, False), + (CallTypes.aretrieve_batch, False), + (CallTypes.acompletion.value, True), + ], +) +async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, expect_key_read): + """ + The batch cost row is logged long after the batch was created, so it keeps the + identity persisted at create time. Every other call type still backfills from + the key. + """ + logger = _ProxyDBLogger() + + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "alias-assigned-later" + mock_key_obj.user_id = "user-assigned-later" + mock_key_obj.team_id = "team-assigned-later" + mock_key_obj.org_id = "org-assigned-later" + + kwargs = { + "call_type": call_type, + "model": None, + "litellm_call_id": "test-call-id", + "stream": False, + "litellm_params": { + "metadata": { + "user_api_key": "hashed_key", + "user_api_key_alias": None, + "user_api_key_user_id": None, + "user_api_key_team_id": None, + "user_api_key_org_id": None, + } + }, + } + + with ( + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ) as mock_get_key, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=MagicMock(team_alias=None), + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert mock_get_key.called is expect_key_read + + written = kwargs["litellm_params"]["metadata"] + if expect_key_read: + assert written["user_api_key_user_id"] == "user-assigned-later" + assert written["user_api_key_team_id"] == "team-assigned-later" + else: + assert written["user_api_key_user_id"] is None + assert written["user_api_key_team_id"] is None + assert written["user_api_key_org_id"] is None + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): """ diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py index 52da7a4a81d..d53e6dedf0b 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py @@ -258,6 +258,7 @@ class TestVertexAIBatchPassthroughHandler: batch_object=batch_object, model_object_id=model_object_id, logging_obj=mock_logging_obj, + is_batch_create=True, user_api_key_dict={"user_id": "test-user"}, ) @@ -307,6 +308,7 @@ class TestVertexAIBatchPassthroughHandler: batch_object={"id": "b1", "object": "batch", "status": "validating"}, model_object_id="b1", logging_obj=mock_logging_obj, + is_batch_create=True, **kwargs, ) @@ -315,6 +317,140 @@ class TestVertexAIBatchPassthroughHandler: assert call_kwargs["user_api_key_dict"].user_id == expected_user_id assert call_kwargs["user_api_key_dict"].team_id == expected_team_id + def _store_with_metadata(self, mock_logging_obj, mock_managed_files_hook, metadata): + with ( + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_pl, + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.verbose_proxy_logger" + ), + ): + mock_pl.get_proxy_hook.return_value = mock_managed_files_hook + VertexPassthroughLoggingHandler._store_batch_managed_object( + unified_object_id="uoi", + batch_object={"id": "b1", "object": "batch", "status": "validating"}, + model_object_id="b1", + logging_obj=mock_logging_obj, + is_batch_create=True, + litellm_params={"metadata": metadata}, + ) + mock_managed_files_hook.store_unified_object_id.assert_called_once() + return mock_managed_files_hook.store_unified_object_id.call_args[1] + + def test_create_persists_key_hash_and_tags( + self, mock_logging_obj, mock_managed_files_hook + ): + """Regression (spend loss): the batch create must persist the creating key's hashed + token and its tags so CheckBatchCost can attribute the batch-cost spend row. Before + this fix the stored api_key was always "" and the row was dropped as unattributed.""" + call_kwargs = self._store_with_metadata( + mock_logging_obj, + mock_managed_files_hook, + { + "user_api_key": "hashed-key-a", + "user_api_key_user_id": "alice", + "user_api_key_team_id": "team-alpha", + "user_api_key_auth_metadata": {"tags": ["env:prod", 7, "team:ml"]}, + }, + ) + + assert call_kwargs["user_api_key_dict"].api_key == "hashed-key-a" + # non-string tags are dropped so downstream tag budgets cannot be bypassed + assert call_kwargs["request_tags"] == ("env:prod", "team:ml") + assert call_kwargs["persist_attribution"] is True + + @pytest.mark.parametrize( + "metadata, expected", + [ + # a request that sent its own tags (x-litellm-tags header or body metadata) + ({"tags": ["req:a", "req:b"]}, ("req:a", "req:b")), + # request tags win over the key's own tags + ( + {"tags": ["req:a"], "user_api_key_auth_metadata": {"tags": ["key:b"]}}, + ("req:a",), + ), + # no request tags: fall back to the tags the key itself carries + ({"user_api_key_auth_metadata": {"tags": ["key:b"]}}, ("key:b",)), + # neither: no tags on the spend row + ({}, None), + ], + ) + def test_request_tags_precedence( + self, mock_logging_obj, mock_managed_files_hook, metadata, expected + ): + """Request tags take precedence over the key's tags, and the key's tags are the + fallback because a tagged key does not put its tags in the top-level metadata.""" + call_kwargs = self._store_with_metadata( + mock_logging_obj, + mock_managed_files_hook, + {"user_api_key": "hashed-key-a", **metadata}, + ) + + assert call_kwargs["request_tags"] == expected + + @pytest.mark.parametrize( + "url_route, expected", + [ + ("/v1/projects/p/locations/us-central1/batchPredictionJobs", True), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs/", True), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs?alt=json", True), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs/123456", False), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs/123456?alt=json", False), + ], + ) + def test_batch_is_registered_from_the_create_route_only( + self, mock_logging_obj, url_route, expected + ): + """Only a POST to the collection route is the create, and only the create claims + attribution. Every id-scoped route is a poll or retrieve, which still reports the + batch so its status and file object stay in sync, but carries is_batch_create=False + so it neither claims the batch nor creates a row it would then own.""" + response = MagicMock() + response.status_code = 200 + response.json.return_value = { + "name": "projects/p/locations/us-central1/batchPredictionJobs/123456", + "model": "publishers/google/models/gemini-2.5-flash", + } + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.verbose_proxy_logger" + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.VertexPassthroughLoggingHandler._store_batch_managed_object" + ) as mock_store, + patch( + "litellm.llms.vertex_ai.batches.transformation.VertexAIBatchTransformation" + ) as mock_transformation, + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.VertexPassthroughLoggingHandler.get_actual_model_id_from_router", + return_value="gemini-2.5-flash", + ), + ): + mock_transformation.transform_vertex_ai_batch_response_to_openai_batch_response.return_value = { + "id": "123456", + "object": "batch", + "status": "validating", + "created_at": 1704067200, + "input_file_id": "gs://bucket/in.jsonl", + "completion_window": "24h", + } + mock_transformation._get_batch_id_from_vertex_ai_batch_response.return_value = "123456" + + VertexPassthroughLoggingHandler.batch_prediction_jobs_handler( + httpx_response=response, + logging_obj=mock_logging_obj, + url_route=url_route, + result="", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + ) + + # every route reports the batch; only the create claims it + mock_store.assert_called_once() + assert mock_store.call_args[1]["unified_object_id"] + assert mock_store.call_args[1]["is_batch_create"] is expected + def test_batch_cost_calculation_integration(self): """Single Vertex AI response → non-zero cost with correct token counts.""" from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage From 2112422c713cf3e497da3351c8317992f77d975f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 16:34:26 -0700 Subject: [PATCH 018/139] test(managed-files): read the scoped page id from the row's unified_file_id --- .../litellm_enterprise/proxy/hooks/test_managed_files.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 9e94b9f2a0f..d3efcf2e7a0 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -3119,8 +3119,9 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): ) managed_row = MagicMock() + managed_row.unified_file_id = "litellm_proxy:mine" managed_row.file_object = { - "id": "litellm_proxy:mine", + "id": "file-mine", "bytes": 100, "created_at": 1, "filename": "mine.jsonl", From 82662dc104db0b26e214c90346627801a7999da3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 17:28:43 -0700 Subject: [PATCH 019/139] fix(proxy): report has_more false on caller-scoped file list pages --- enterprise/litellm_enterprise/proxy/hooks/managed_files.py | 6 ++++-- .../litellm_enterprise/proxy/hooks/test_managed_files.py | 3 ++- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index a6ba1e0a791..37d267fcd6e 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1297,13 +1297,15 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): """Rebuild ``first_id`` / ``last_id`` from the caller-scoped page. The upstream cursors point at rows that were just filtered out, so - leaving them in place discloses other callers' file ids. + leaving them in place discloses other callers' file ids. ``has_more`` + is always cleared because ``after`` is never forwarded upstream, so + no further page is reachable through the proxy. """ if hasattr(response, "first_id"): response.first_id = data[0].id if data else None if hasattr(response, "last_id"): response.last_id = data[-1].id if data else None - if not data and hasattr(response, "has_more"): + if hasattr(response, "has_more"): response.has_more = False async def afile_retrieve( diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index d3efcf2e7a0..fde1feb80e2 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -3112,7 +3112,7 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): upstream_page = AsyncCursorPage[FileObject].construct( data=[_raw_file("file-someone-else"), _raw_file("file-mine")], - has_more=False, + has_more=True, first_id="file-someone-else", last_id="file-mine", object="list", @@ -3146,3 +3146,4 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): assert [file_object.id for file_object in response.data] == ["litellm_proxy:mine"] assert response.first_id == "litellm_proxy:mine" assert response.last_id == "litellm_proxy:mine" + assert response.has_more is False From 30c4898de9cea90e777a43a5260fe77011acdb5b Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 8 Aug 2026 20:16:22 -0700 Subject: [PATCH 020/139] fix(ui): hide admin-only Logs tabs from roles that cannot call their endpoints The Logs nav entry is open to internal users so they can read their own request logs, but the page rendered all four tabs unconditionally. Audit Logs calls GET /audit and Deleted Teams calls GET /v2/team/list?status=deleted, neither of which an internal user is permitted to call, so the page fired requests that came back 401. Gate both tabs on new viewAuditLogs / viewDeletedTeams capabilities, using the same CAPABILITY_ROLES map and useCan hook introduced for Tool Policies. Hiding a tab drops its panel from the tree entirely, so the request is never issued rather than issued and rejected. Selecting a tab also mapped index 0 to "request logs" and every other index to "audit logs", which activated the audit panel whenever a user opened Deleted Keys or Deleted Teams. Derive the active tab from the visible tab list instead, so the mapping survives tabs being filtered out. --- .../view_logs/index.integration.test.tsx | 104 ++++++++++++++++++ .../src/components/view_logs/index.test.tsx | 79 ++++++++++++- .../src/components/view_logs/index.tsx | 91 +++++++++------ .../src/utils/capabilities.test.ts | 13 +++ .../src/utils/capabilities.ts | 2 + 5 files changed, 254 insertions(+), 35 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx new file mode 100644 index 00000000000..b86ad015b91 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx @@ -0,0 +1,104 @@ +import { screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import SpendLogsTable from "./index"; +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; + +const { useAuthorizedMock } = vi.hoisted(() => ({ useAuthorizedMock: vi.fn() })); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: useAuthorizedMock, +})); + +vi.mock("./RequestLogsPanel", () => ({ + default: function RequestLogsPanelMock() { + return
; + }, +})); + +const fetchMock = vi.fn(); + +const jsonResponse = (body: unknown) => ({ + ok: true, + status: 200, + statusText: "OK", + json: async () => body, +}); + +const requestedUrls = () => fetchMock.mock.calls.map(([url]) => String(url)); + +const emptyAuditLogs = { audit_logs: [], total: 0, page: 1, page_size: 50, total_pages: 0 }; + +const defaultProps = { + accessToken: "sk-test", + token: "jwt-test", + userRole: "Admin", + userID: "user-1", + premiumUser: true, +}; + +const renderAs = (sessionRole: string) => { + useAuthorizedMock.mockReturnValue({ accessToken: "sk-test", userRole: sessionRole, premiumUser: true }); + return renderWithProviders(); +}; + +describe("SpendLogsTable network access by role", () => { + beforeEach(() => { + testQueryClient.clear(); + vi.clearAllMocks(); + fetchMock.mockImplementation(async (url: string) => { + if (String(url).includes("/audit")) { + return jsonResponse(emptyAuditLogs); + } + if (String(url).includes("/v2/team/list")) { + return jsonResponse({ teams: [] }); + } + return jsonResponse({ keys: [], total_count: 0 }); + }); + vi.stubGlobal("fetch", fetchMock); + }); + + it("fires neither the audit nor the deleted-teams request for an internal user", async () => { + const user = userEvent.setup(); + renderAs("Internal User"); + + // Liveness gate: the sibling Deleted Keys panel does reach the network, so a + // silent absence below means the gate worked, not that nothing rendered. + await waitFor(() => expect(requestedUrls().some((url) => url.includes("/key/list"))).toBe(true)); + + await user.click(screen.getByRole("tab", { name: "Deleted Keys" })); + await user.click(screen.getByRole("tab", { name: "Request Logs" })); + + expect(requestedUrls().filter((url) => url.includes("/audit"))).toEqual([]); + expect(requestedUrls().filter((url) => url.includes("/v2/team/list"))).toEqual([]); + }); + + it("fetches deleted teams and audit logs for an admin", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await waitFor(() => + expect(requestedUrls().some((url) => url.includes("/v2/team/list") && url.includes("status=deleted"))).toBe(true), + ); + + expect(requestedUrls().filter((url) => url.includes("/audit"))).toEqual([]); + + await user.click(screen.getByRole("tab", { name: "Audit Logs" })); + + await waitFor(() => expect(requestedUrls().some((url) => url.includes("/audit"))).toBe(true)); + }); + + it("leaves the audit request unsent when an admin selects a tab after Audit Logs", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await user.click(screen.getByRole("tab", { name: "Deleted Teams" })); + + expect(screen.getByRole("tab", { name: "Deleted Teams" })).toHaveAttribute("aria-selected", "true"); + expect(requestedUrls().filter((url) => url.includes("/audit"))).toEqual([]); + + await user.click(screen.getByRole("tab", { name: "Audit Logs" })); + + await waitFor(() => expect(requestedUrls().some((url) => url.includes("/audit"))).toBe(true)); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx index b2e77ec7fd5..785fa0cc6f8 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx @@ -1,9 +1,15 @@ import { screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import { describe, expect, it, vi } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import SpendLogsTable from "./index"; import { renderWithProviders } from "../../../tests/test-utils"; +const { useAuthorizedMock } = vi.hoisted(() => ({ useAuthorizedMock: vi.fn() })); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: useAuthorizedMock, +})); + vi.mock("./RequestLogsPanel", () => ({ default: function RequestLogsPanelMock({ isActive }: { isActive: boolean }) { return
{isActive ? "active" : "inactive"}
; @@ -36,9 +42,18 @@ const defaultProps = { premiumUser: false, }; +const renderAs = (sessionRole: string) => { + useAuthorizedMock.mockReturnValue({ userRole: sessionRole }); + return renderWithProviders(); +}; + describe("SpendLogsTable", () => { + beforeEach(() => { + useAuthorizedMock.mockReturnValue({ userRole: "Admin" }); + }); + it("renders the four log tabs", () => { - renderWithProviders(); + renderAs("Admin"); for (const label of ["Request Logs", "Audit Logs", "Deleted Keys", "Deleted Teams"]) { expect(screen.getByRole("tab", { name: label })).toBeInTheDocument(); @@ -47,7 +62,7 @@ describe("SpendLogsTable", () => { it("marks only the visible tab's panel active so background tabs do not query", async () => { const user = userEvent.setup(); - renderWithProviders(); + renderAs("Admin"); expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("active"); @@ -57,8 +72,64 @@ describe("SpendLogsTable", () => { expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); }); + describe("admin-only tabs", () => { + it.each(["Internal User", "Internal Viewer"])("hides Audit Logs and Deleted Teams from %s", (role) => { + renderAs(role); + + expect(screen.getByRole("tab", { name: "Request Logs" })).toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Deleted Keys" })).toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Audit Logs" })).not.toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Deleted Teams" })).not.toBeInTheDocument(); + }); + + it("never mounts the panels that call the admin-only endpoints for an internal user", () => { + renderAs("Internal User"); + + expect(screen.queryByTestId("audit-logs-panel")).not.toBeInTheDocument(); + expect(screen.queryByTestId("deleted-teams-page")).not.toBeInTheDocument(); + expect(screen.getByTestId("deleted-keys-page")).toBeInTheDocument(); + }); + }); + + describe("tab index mapping", () => { + it("activates the panel the admin selected, not the one at the old hardcoded index", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await user.click(screen.getByRole("tab", { name: "Deleted Keys" })); + + expect(screen.getByTestId("audit-logs-panel")).toHaveTextContent("inactive"); + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); + }); + + it("keeps the audit panel inert when an admin selects the last tab", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await user.click(screen.getByRole("tab", { name: "Deleted Teams" })); + + expect(screen.getByTestId("audit-logs-panel")).toHaveTextContent("inactive"); + expect(screen.getByTestId("deleted-teams-page")).toBeInTheDocument(); + }); + + it("selects the last visible tab for an internal user and returns to Request Logs", async () => { + const user = userEvent.setup(); + renderAs("Internal User"); + + await user.click(screen.getByRole("tab", { name: "Deleted Keys" })); + + expect(screen.getByTestId("deleted-keys-page")).toBeInTheDocument(); + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); + + await user.click(screen.getByRole("tab", { name: "Request Logs" })); + + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("active"); + }); + }); + describe("auth-not-ready guard", () => { it("shows a loading spinner when credentials are not yet resolved", () => { + useAuthorizedMock.mockReturnValue({ userRole: "Admin" }); renderWithProviders(); expect(document.querySelector(".ant-spin")).toBeInTheDocument(); @@ -66,7 +137,7 @@ describe("SpendLogsTable", () => { }); it("renders the tabs (no spinner) once all credentials are present", () => { - renderWithProviders(); + renderAs("Admin"); expect(document.querySelector(".ant-spin")).not.toBeInTheDocument(); expect(screen.getByRole("tab", { name: "Request Logs" })).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index 8e7423e3fae..7269564dcec 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -1,5 +1,6 @@ import { useState } from "react"; import { Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage"; import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage"; import AuditLogsPanel from "./AuditLogsPanel"; @@ -14,8 +15,22 @@ interface SpendLogsTableProps { premiumUser: boolean; } +type LogsTabId = "request logs" | "audit logs" | "deleted keys" | "deleted teams"; + +interface LogsTab { + id: LogsTabId; + label: string; +} + +const REQUEST_LOGS_TAB: LogsTab = { id: "request logs", label: "Request Logs" }; +const AUDIT_LOGS_TAB: LogsTab = { id: "audit logs", label: "Audit Logs" }; +const DELETED_KEYS_TAB: LogsTab = { id: "deleted keys", label: "Deleted Keys" }; +const DELETED_TEAMS_TAB: LogsTab = { id: "deleted teams", label: "Deleted Teams" }; + export default function SpendLogsTable({ accessToken, token, userRole, userID, premiumUser }: SpendLogsTableProps) { - const [activeTab, setActiveTab] = useState("request logs"); + const [activeTab, setActiveTab] = useState(REQUEST_LOGS_TAB.id); + const canViewAuditLogs = useCan("viewAuditLogs"); + const canViewDeletedTeams = useCan("viewDeletedTeams"); if (!accessToken || !token || !userRole || !userID) { return ( @@ -25,41 +40,55 @@ export default function SpendLogsTable({ accessToken, token, userRole, userID, p ); } + const tabs: LogsTab[] = [ + REQUEST_LOGS_TAB, + ...(canViewAuditLogs ? [AUDIT_LOGS_TAB] : []), + DELETED_KEYS_TAB, + ...(canViewDeletedTeams ? [DELETED_TEAMS_TAB] : []), + ]; + + const renderPanel = (tabId: LogsTabId) => { + switch (tabId) { + case "request logs": + return ( + + ); + case "audit logs": + return ( + + ); + case "deleted keys": + return ; + case "deleted teams": + return ; + } + }; + return (
- setActiveTab(index === 0 ? "request logs" : "audit logs")}> + setActiveTab(tabs[index].id)}> - Request Logs - Audit Logs - Deleted Keys - Deleted Teams + {tabs.map((tab) => ( + {tab.label} + ))} - - - - - - - - - - - - + {tabs.map((tab) => ( + {renderPanel(tab.id)} + ))}
diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index f48609b0b9d..611c9626065 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -18,6 +18,19 @@ describe("hasCapability", () => { ); }); +describe.each(["viewAuditLogs", "viewDeletedTeams"] as const)("hasCapability - %s", (capability) => { + it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])("should grant it to %s", (role) => { + expect(hasCapability(role, capability)).toBe(true); + }); + + it.each(["Internal User", "Internal Viewer", "App User", "Org Admin", "Unknown Role", "", null, undefined])( + "should deny it to %s", + (role) => { + expect(hasCapability(role, capability)).toBe(false); + }, + ); +}); + describe("rolesWithCapability", () => { it("should return a copy so callers cannot mutate the capability map", () => { const roles = rolesWithCapability("viewToolPolicies"); diff --git a/ui/litellm-dashboard/src/utils/capabilities.ts b/ui/litellm-dashboard/src/utils/capabilities.ts index 77ead2568fb..f0847cc3400 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.ts @@ -2,6 +2,8 @@ import { all_admin_roles } from "./roles"; const CAPABILITY_ROLES = { viewToolPolicies: all_admin_roles, + viewAuditLogs: all_admin_roles, + viewDeletedTeams: all_admin_roles, } as const satisfies Record; export type Capability = keyof typeof CAPABILITY_ROLES; From 6a540a1bf848129dc16228d1b23f92120d7a7f03 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 8 Aug 2026 20:17:12 -0700 Subject: [PATCH 021/139] fix(ui): gate organization and agent usage views behind capabilities The Usage page admits internal users because their own usage view works, but the entity breakdown selector inside it also offered Organization Usage, so picking it fired /organization/daily/activity and collected a 401. Neither that route nor /agent/daily/activity appears in any non-admin route list, so both are default-deny. The team breakdown leaked the second one too: it fetches agent activity unconditionally to fill its Top Agents card, which 401s for the same roles. Adds viewOrganizationUsage and viewAgentUsage to the existing capability map and points the selector option, the page section, and the fetch's enabled flag at the same capability, so a role that cannot call the endpoint never sees the breakdown and never issues the request. The team and tag breakdowns, which internal users can read, are untouched, and the default Usage view was already one of those. --- .../EntityUsage/EntityUsage.test.tsx | 43 ++++++++++++++++++- .../components/EntityUsage/EntityUsage.tsx | 34 +++++++++++---- .../components/UsagePageView.test.tsx | 25 +++++++++++ .../_components/components/UsagePageView.tsx | 9 ++-- .../UsageViewSelect/UsageViewSelect.test.tsx | 29 +++++++++++-- .../UsageViewSelect/UsageViewSelect.tsx | 20 +++++---- .../src/utils/capabilities.test.ts | 36 ++++++++++------ .../src/utils/capabilities.ts | 2 + 8 files changed, 160 insertions(+), 38 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 82ca66b10c0..11528117f1e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -1,4 +1,4 @@ -import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react"; +import { act, cleanup, fireEvent, render, screen, waitFor, within } from "@testing-library/react"; import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import * as networking from "@/components/networking"; import EntityUsage from "./EntityUsage"; @@ -856,6 +856,47 @@ describe("EntityUsage", () => { expect(logo.getAttribute("src")).toContain("openai_small"); }); + describe("capability gating", () => { + it.each([ + ["organization", () => mockOrganizationDailyActivityCall, "Organization Spend Overview"], + ["agent", () => mockAgentDailyActivityCall, "Agent Spend Overview"], + ] as const)("fetches %s activity for an admin but not for an internal user", async (entityType, call, heading) => { + render(); + await waitFor(() => { + expect(call()).toHaveBeenCalled(); + }); + + cleanup(); + call().mockClear(); + + render(); + expect(await screen.findByText(heading)).toBeInTheDocument(); + expect(call()).not.toHaveBeenCalled(); + }); + + it("keeps the team breakdown but drops its agent sub-fetch for an internal user", async () => { + render(); + + await waitFor(() => { + expect(mockTeamDailyActivityCall).toHaveBeenCalled(); + }); + expect(screen.getByText("Team Spend Overview")).toBeInTheDocument(); + + expect(mockAgentDailyActivityCall).not.toHaveBeenCalled(); + expect(screen.queryByText("Agent Activity")).not.toBeInTheDocument(); + expect(screen.queryByText("Top Agents Driving Spend")).not.toBeInTheDocument(); + }); + + it("keeps the tag breakdown for an internal user", async () => { + render(); + + await waitFor(() => { + expect(mockTagDailyActivityCall).toHaveBeenCalled(); + }); + expect(screen.getByText("Tag Spend Overview")).toBeInTheDocument(); + }); + }); + it("renders a letter avatar instead of an img for an unknown provider slug", async () => { const spendDataUnknownProvider = { ...mockSpendData, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 4d44791d1a9..5a0f2abf15b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -2,6 +2,7 @@ import useTeams from "@/app/(dashboard)/hooks/useTeams"; import { BarChart, DonutChart } from "@/components/shared/charts"; import { MoneyCell } from "@/components/shared/table_cells"; import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { hasCapability, type Capability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { Card, @@ -108,7 +109,19 @@ const ENTITY_FETCH_FNS: Record Promise> = { user: userDailyActivityCall, }; -const EntityUsage: React.FC = ({ accessToken, entityType, entityId, entityList, dateValue }) => { +const ENTITY_CAPABILITIES: Partial> = { + organization: "viewOrganizationUsage", + agent: "viewAgentUsage", +}; + +const EntityUsage: React.FC = ({ + accessToken, + entityType, + entityId, + entityList, + userRole, + dateValue, +}) => { const { teams } = useTeams(); const [selectedTags, setSelectedTags] = useState([]); const [modelViewType, setModelViewType] = useState("groups"); @@ -125,7 +138,11 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti }, [entityType, selectedTags]); const fetchFn = ENTITY_FETCH_FNS[entityType]; - const enabled = !!accessToken && !!startTime && !!endTime; + const entityCapability = ENTITY_CAPABILITIES[entityType]; + const canViewEntity = entityCapability === undefined || hasCapability(userRole, entityCapability); + const showAgentBreakdown = entityType === "team" && hasCapability(userRole, "viewAgentUsage"); + const hasRequestWindow = !!accessToken && !!startTime && !!endTime; + const enabled = hasRequestWindow && canViewEntity; const { data: spendDataRaw, @@ -150,7 +167,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti } = usePaginatedDailyActivity({ fetchFn: agentDailyActivityCall, args: [accessToken, startTime, endTime, null], - enabled: enabled && entityType === "team", + enabled: enabled && showAgentBreakdown, }); const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData; @@ -158,7 +175,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti const modelBreakdownKey = modelViewType === "groups" ? "model_groups" : "models"; const modelMetrics = processActivityData(spendData, modelBreakdownKey, teams || []); const keyMetrics = processActivityData(spendData, "api_keys", teams || []); - const agentMetrics = entityType === "team" ? processActivityData(agentSpendData, "entities", teams || []) : {}; + const agentMetrics = showAgentBreakdown ? processActivityData(agentSpendData, "entities", teams || []) : {}; const getTopModels = () => { const modelSpend: { [key: string]: any } = {}; @@ -621,8 +638,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti - {/* Top Agents - only for team entity type */} - {entityType === "team" && ( + {showAgentBreakdown && ( Top Agents Driving Spend @@ -708,7 +724,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti ), }, - ...(entityType === "team" + ...(showAgentBreakdown ? [{ key: "agents", label: "Agent Activity", content: }] : []), { @@ -757,7 +773,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti } /> )} - {agentIsFetchingMore && entityType === "team" && ( + {agentIsFetchingMore && showAgentBreakdown && ( = ({ accessToken, entityType, enti } /> )} - {agentCancelled && entityType === "team" && ( + {agentCancelled && showAgentBreakdown && ( { userId: "user-123", userEmail: "test@example.com", userRole: "Internal User", + userRoleLabel: "Internal User", + isViewOnly: false, premiumUser: true, disabledPersonalKeyCreation: false, showSSOBanner: false, @@ -861,6 +863,29 @@ describe("UsagePage", () => { }); }); + // The select hides both views from a non-admin, so this drives the section + // gate directly through the mocked select, which always offers every option. + it.each(["organization", "agent"])("should not render the %s usage view for an internal user", async (usageView) => { + mockUseAuthorized.mockReturnValue(nonAdminSession); + + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + const usageSelect = screen.getByTestId("usage-view-select"); + act(() => { + fireEvent.change(usageSelect, { target: { value: "team" } }); + }); + expect(screen.getAllByText("Entity Usage").length).toBeGreaterThan(0); + + act(() => { + fireEvent.change(usageSelect, { target: { value: usageView } }); + }); + expect(screen.queryByText("Entity Usage")).not.toBeInTheDocument(); + }); + describe("admin user selector", () => { it("should render user selector for admin users in global view", async () => { renderWithProviders(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index c3645d6371e..494df313ac0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -33,6 +33,7 @@ import { useCustomers } from "@/app/(dashboard)/hooks/customers/useCustomers"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser"; import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers"; +import { hasCapability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { all_admin_roles, internalUserRoles } from "@/utils/roles"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; @@ -109,6 +110,8 @@ const UsagePage: React.FC = ({ teams, organizations }) => { const { data: currentUser } = useCurrentUser(); const isAdmin = all_admin_roles.includes(userRole || ""); const canViewTagUsage = isAdmin || internalUserRoles.includes(userRole || ""); + const canViewOrganizationUsage = hasCapability(userRole, "viewOrganizationUsage"); + const canViewAgentUsage = hasCapability(userRole, "viewAgentUsage"); // Debounced search for user selector const [userSearchInput, setUserSearchInput] = useState(""); @@ -513,7 +516,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { setUsageView(value)} - isAdmin={isAdmin} + userRole={userRole} canViewTagUsage={canViewTagUsage} /> @@ -950,7 +953,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { )} {/* Organization Usage Panel */} - {usageView === "organization" && ( + {usageView === "organization" && canViewOrganizationUsage && ( = ({ teams, organizations }) => { /> )} - {usageView === "agent" && ( + {usageView === "agent" && canViewAgentUsage && ( { }); it("should render", () => { - render(); + render(); expect(screen.getByText("Usage View")).toBeInTheDocument(); expect(screen.getByText("Select the usage data you want to view")).toBeInTheDocument(); expect(screen.getByRole("combobox")).toBeInTheDocument(); + expect(screen.getByRole("option", { name: "Your Usage" })).toBeInTheDocument(); }); it("should call onChange when value changes", () => { - render(); + render(); const select = screen.getByRole("combobox"); act(() => { @@ -109,14 +110,34 @@ describe("UsageViewSelect", () => { }); it("should show Tag Usage for non-admin users with tag usage permission", () => { - render(); + render(); expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument(); }); it("should hide Tag Usage for non-admin users without tag usage permission", () => { - render(); + render(); expect(screen.queryByRole("option", { name: "Tag Usage" })).not.toBeInTheDocument(); }); + + it.each(["Organization Usage", "Agent Usage (A2A)"])("should show %s to an admin", (optionName) => { + render(); + + expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument(); + }); + + // Neither /organization/daily/activity nor /agent/daily/activity admits an + // internal user, so the option that fires them must not be selectable. + it.each(["Organization Usage", "Agent Usage (A2A)"])("should hide %s from an internal user", (optionName) => { + render(); + + expect(screen.queryByRole("option", { name: optionName })).not.toBeInTheDocument(); + }); + + it.each(["Team Usage", "Tag Usage"])("should keep %s available to an internal user", (optionName) => { + render(); + + expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx index 94b483cb539..54c1d5ab7cc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx @@ -11,6 +11,8 @@ import { } from "@ant-design/icons"; import { Badge, Select } from "antd"; import React from "react"; +import { hasCapability, type Capability } from "@/utils/capabilities"; +import { all_admin_roles } from "@/utils/roles"; export type UsageOption = | "global" | "my-usage" @@ -24,7 +26,7 @@ export type UsageOption = export interface UsageViewSelectProps { value: UsageOption; onChange: (value: UsageOption) => void; - isAdmin: boolean; + userRole: string | null; canViewTagUsage?: boolean; title?: string; description?: string; @@ -35,6 +37,7 @@ interface OptionConfig { label: string; description: string; icon: React.ReactNode; + capability?: Capability; adminOnly?: boolean; showForAdmin?: string; showForNonAdmin?: string; @@ -63,12 +66,9 @@ const OPTIONS: OptionConfig[] = [ { value: "organization", label: "Organization Usage", - showForAdmin: "Organization Usage", - showForNonAdmin: "Your Organization Usage", - description: "View organization-level usage", - descriptionForAdmin: "View usage across all organizations", - descriptionForNonAdmin: "View your organization's usage", + description: "View usage across all organizations", icon: , + capability: "viewOrganizationUsage", }, { value: "team", @@ -95,7 +95,7 @@ const OPTIONS: OptionConfig[] = [ label: "Agent Usage (A2A)", description: "View usage by AI agents", icon: , - adminOnly: true, + capability: "viewAgentUsage", }, { value: "user", @@ -115,14 +115,18 @@ const OPTIONS: OptionConfig[] = [ export const UsageViewSelect: React.FC = ({ value, onChange, - isAdmin, + userRole, canViewTagUsage = false, title = "Usage View", description = "Select the usage data you want to view", "data-id": dataId, }) => { + const isAdmin = all_admin_roles.includes(userRole ?? ""); const getFilteredOptions = () => { return OPTIONS.filter((option) => { + if (option.capability) { + return hasCapability(userRole, option.capability); + } if (option.value === "tag" && canViewTagUsage) { return true; } diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index f48609b0b9d..3f5a8ac81fb 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -1,21 +1,31 @@ import { describe, expect, it } from "vitest"; -import { hasCapability, rolesWithCapability } from "./capabilities"; +import { hasCapability, rolesWithCapability, type Capability } from "./capabilities"; + +const ADMIN_ROLES = ["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"]; +const NON_ADMIN_ROLES = [ + "Internal User", + "Internal Viewer", + "App User", + "Org Admin", + "Unknown Role", + "", + null, + undefined, +]; + +const ADMIN_ONLY_CAPABILITIES: Capability[] = ["viewToolPolicies", "viewOrganizationUsage", "viewAgentUsage"]; describe("hasCapability", () => { - it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])( - "should grant viewToolPolicies to %s", - (role) => { - expect(hasCapability(role, "viewToolPolicies")).toBe(true); - }, - ); + describe.each(ADMIN_ONLY_CAPABILITIES)("%s", (capability) => { + it.each(ADMIN_ROLES)("should grant it to %s", (role) => { + expect(hasCapability(role, capability)).toBe(true); + }); - it.each(["Internal User", "Internal Viewer", "App User", "Org Admin", "Unknown Role", "", null, undefined])( - "should deny viewToolPolicies to %s", - (role) => { - expect(hasCapability(role, "viewToolPolicies")).toBe(false); - }, - ); + it.each(NON_ADMIN_ROLES)("should deny it to %s", (role) => { + expect(hasCapability(role, capability)).toBe(false); + }); + }); }); describe("rolesWithCapability", () => { diff --git a/ui/litellm-dashboard/src/utils/capabilities.ts b/ui/litellm-dashboard/src/utils/capabilities.ts index 77ead2568fb..c4d878b81a5 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.ts @@ -2,6 +2,8 @@ import { all_admin_roles } from "./roles"; const CAPABILITY_ROLES = { viewToolPolicies: all_admin_roles, + viewOrganizationUsage: all_admin_roles, + viewAgentUsage: all_admin_roles, } as const satisfies Record; export type Capability = keyof typeof CAPABILITY_ROLES; From 2502ee4a2ace88ac10dd8ed30e2a03fff075ce26 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 8 Aug 2026 20:34:56 -0700 Subject: [PATCH 022/139] fix(ui): gate policy and prompt lookups on an admin capability /policies/list and /prompts/list are default-deny for internal_user, but the Virtual Keys create/edit flow, the Teams forms and the Playground called them on mount, so every internal user landing on the dashboard fired two requests that 401. Add viewPolicies and viewPrompts to the capability map and use them to gate the nav entry, the form field and the fetch together, following the pattern from the Tool Policies migration. Non-admins now see no policy or prompt selector at all rather than an empty dropdown. --- .../playground/components/chat_ui/ChatUI.tsx | 56 +++---- .../components/complianceUI/ComplianceUI.tsx | 48 +++--- .../src/components/Teams.test.tsx | 55 ++++++- ui/litellm-dashboard/src/components/Teams.tsx | 68 ++++---- .../src/components/leftnav.test.tsx | 22 +++ .../src/components/leftnav.tsx | 10 +- .../organisms/create_key_button.test.tsx | 49 +++++- .../organisms/create_key_button.tsx | 145 +++++++++--------- .../policies/PolicySelector.test.tsx | 20 +++ .../components/policies/PolicySelector.tsx | 10 +- .../src/components/team/TeamInfo.test.tsx | 46 ++++++ .../src/components/team/TeamInfo.tsx | 56 +++---- .../templates/key_edit_view.test.tsx | 50 +++++- .../components/templates/key_edit_view.tsx | 87 ++++++----- .../src/utils/capabilities.test.ts | 28 ++++ .../src/utils/capabilities.ts | 2 + 16 files changed, 533 insertions(+), 219 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx index 57ff7906eda..0241ef8a77e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx @@ -25,6 +25,7 @@ import React, { useEffect, useRef, useState } from "react"; import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; import { v4 as uuidv4 } from "uuid"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import GuardrailSelector from "@/components/guardrails/GuardrailSelector"; import PolicySelector from "@/components/policies/PolicySelector"; import MCPToolArgumentsForm, { MCPToolArgumentsFormRef } from "@/components/mcp_tools/MCPToolArgumentsForm"; @@ -106,6 +107,7 @@ const ChatUI: React.FC = ({ simplified = false, fixedModel, }) => { + const canViewPolicies = useCan("viewPolicies"); const [mcpServers, setMCPServers] = useState([]); const [mcpToolsets, setMCPToolsets] = useState([]); const [isToolsetsInfoModalVisible, setIsToolsetsInfoModalVisible] = useState(false); @@ -1652,32 +1654,34 @@ const ChatUI: React.FC = ({ />
-
- - Policies - - Select policy/policies to apply to this LLM API call. Policies define which guardrails are - applied based on conditions. You can set up your policies{" "} - - here - - . - - } - > - - - - -
+ {canViewPolicies && ( +
+ + Policies + + Select policy/policies to apply to this LLM API call. Policies define which guardrails are + applied based on conditions. You can set up your policies{" "} + + here + + . + + } + > + + + + +
+ )} {/* Code Interpreter Toggle - Only for Responses endpoint */} {endpointType === EndpointType.RESPONSES && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx index 39346105f2a..c3b417987e6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx @@ -6,6 +6,7 @@ import { type ComplianceFramework, type CompliancePrompt, } from "@/data/compliancePrompts"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import { getGuardrailsList, testPoliciesAndGuardrails } from "@/components/networking"; import PolicySelector, { getPolicyOptionEntries } from "@/components/policies/PolicySelector"; import { Policy } from "@/components/policies/types"; @@ -123,6 +124,7 @@ export default function ComplianceUI({ fixedModel, proxySettings, }: ComplianceUIProps) { + const canViewPolicies = useCan("viewPolicies"); const frameworks = getFrameworks(); const [policyValueToLabel, setPolicyValueToLabel] = useState>(new Map()); @@ -701,29 +703,37 @@ export default function ComplianceUI({

Test Configuration

-

Select policies, guardrails, or both to test against.

+

+ {canViewPolicies + ? "Select policies, guardrails, or both to test against." + : "Select guardrails to test against."} +

-
- - {accessToken && ( - - )} -
+ {canViewPolicies && ( + <> +
+ + {accessToken && ( + + )} +
-
-
- or -
-
+
+
+ or +
+
+ + )}