diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 294b9d464f5..d022af689d3 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -1087,9 +1087,11 @@ class PrometheusLogger(CustomLogger): ), client_ip=standard_logging_payload["metadata"].get("requester_ip_address"), user_agent=standard_logging_payload["metadata"].get("user_agent"), - stream=str(standard_logging_payload.get("stream")) - if litellm.prometheus_emit_stream_label - else None, + stream=( + str(standard_logging_payload.get("stream")) + if litellm.prometheus_emit_stream_label + else None + ), ) if ( @@ -1755,9 +1757,11 @@ class PrometheusLogger(CustomLogger): client_ip=_metadata.get("requester_ip_address"), user_agent=_metadata.get("user_agent"), model_id=model_id, - stream=str(request_data.get("stream")) - if litellm.prometheus_emit_stream_label - else None, + stream=( + str(request_data.get("stream")) + if litellm.prometheus_emit_stream_label + else None + ), ) _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( @@ -2081,9 +2085,9 @@ class PrometheusLogger(CustomLogger): ): try: verbose_logger.debug("setting remaining tokens requests metric") - standard_logging_payload: Optional[ - StandardLoggingPayload - ] = request_kwargs.get("standard_logging_object") + standard_logging_payload: Optional[StandardLoggingPayload] = ( + request_kwargs.get("standard_logging_object") + ) if standard_logging_payload is None: return @@ -2716,9 +2720,7 @@ class PrometheusLogger(CustomLogger): ) return - async def fetch_keys( - page_size: int, page: int - ) -> Tuple[ + async def fetch_keys(page_size: int, page: int) -> Tuple[ List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]], Optional[int], ]: @@ -2789,9 +2791,7 @@ class PrometheusLogger(CustomLogger): ) return - async def fetch_orgs( - page_size: int, page: int - ) -> Tuple[list, Optional[int]]: + async def fetch_orgs(page_size: int, page: int) -> Tuple[list, Optional[int]]: skip = (page - 1) * page_size orgs = await prisma_client.db.litellm_organizationtable.find_many( skip=skip, @@ -2911,9 +2911,11 @@ class PrometheusLogger(CustomLogger): org_alias=org.organization_alias or "", spend=org.spend or 0.0, max_budget=budget_table.max_budget if budget_table else None, - budget_reset_at=getattr(budget_table, "budget_reset_at", None) - if budget_table - else None, + budget_reset_at=( + getattr(budget_table, "budget_reset_at", None) + if budget_table + else None + ), ) async def _set_team_budget_metrics_after_api_request( @@ -3395,10 +3397,10 @@ class PrometheusLogger(CustomLogger): from litellm.constants import PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES from litellm.integrations.custom_logger import CustomLogger - prometheus_loggers: List[ - CustomLogger - ] = litellm.logging_callback_manager.get_custom_loggers_for_type( - callback_type=PrometheusLogger + prometheus_loggers: List[CustomLogger] = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + callback_type=PrometheusLogger + ) ) # we need to get the initialized prometheus logger instance(s) and call logger.initialize_remaining_budget_metrics() on them verbose_logger.debug("found %s prometheus loggers", len(prometheus_loggers)) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 90b501a9cbb..e01f2994810 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -160,7 +160,8 @@ class AmazonConverseConfig(BaseConfig): if isinstance(content, list): has_guarded_text = any( - isinstance(item, dict) and item.get("type") == "guarded_text" for item in content + isinstance(item, dict) and item.get("type") == "guarded_text" + for item in content ) if has_guarded_text: continue # Skip this message if it already has guarded_text @@ -321,9 +322,13 @@ class AmazonConverseConfig(BaseConfig): # Check if the model is a Nova 2 model (matches nova-2-lite, nova-2-pro, etc.) # Also check for nova-2/ spec prefix for imported models - return model_without_region.startswith("amazon.nova-2-") or model_without_region.startswith("nova-2/") + return model_without_region.startswith( + "amazon.nova-2-" + ) or model_without_region.startswith("nova-2/") - def _map_web_search_options(self, web_search_options: dict, model: str) -> Optional[BedrockToolBlock]: + def _map_web_search_options( + self, web_search_options: dict, model: str + ) -> Optional[BedrockToolBlock]: """ Map web_search_options to Nova grounding systemTool. @@ -352,7 +357,9 @@ class AmazonConverseConfig(BaseConfig): # (unlike Anthropic), so we just enable grounding with no options return BedrockToolBlock(systemTool={"name": "nova_grounding"}) - def _transform_reasoning_effort_to_reasoning_config(self, reasoning_effort: str) -> dict: + def _transform_reasoning_effort_to_reasoning_config( + self, reasoning_effort: str + ) -> dict: """ Transform reasoning_effort parameter to Nova 2 reasoningConfig structure. @@ -397,7 +404,9 @@ class AmazonConverseConfig(BaseConfig): } } - def _handle_reasoning_effort_parameter(self, model: str, reasoning_effort: str, optional_params: dict) -> None: + def _handle_reasoning_effort_parameter( + self, model: str, reasoning_effort: str, optional_params: dict + ) -> None: """ Handle the reasoning_effort parameter based on the model type. @@ -434,7 +443,9 @@ class AmazonConverseConfig(BaseConfig): optional_params["reasoning_effort"] = reasoning_effort elif self._is_nova_2_model(model): # Nova 2 models: transform to reasoningConfig - reasoning_config = self._transform_reasoning_effort_to_reasoning_config(reasoning_effort) + reasoning_config = self._transform_reasoning_effort_to_reasoning_config( + reasoning_effort + ) optional_params.update(reasoning_config) else: # Anthropic and other models: convert to thinking parameter @@ -478,7 +489,9 @@ class AmazonConverseConfig(BaseConfig): "parallel_tool_calls", ] - if "arn" in model: # we can't infer the model from the arn, so just add all params + if ( + "arn" in model + ): # we can't infer the model from the arn, so just add all params supported_params.append("tools") supported_params.append("tool_choice") supported_params.append("thinking") @@ -500,7 +513,9 @@ class AmazonConverseConfig(BaseConfig): or base_model.startswith("meta.llama3-3") or base_model.startswith("meta.llama4") or base_model.startswith("amazon.nova") - or supports_function_calling(model=model, custom_llm_provider=self.custom_llm_provider) + or supports_function_calling( + model=model, custom_llm_provider=self.custom_llm_provider + ) ): supported_params.append("tools") @@ -510,7 +525,9 @@ class AmazonConverseConfig(BaseConfig): if litellm.utils.supports_tool_choice( model=model, custom_llm_provider=self.custom_llm_provider - ) or litellm.utils.supports_tool_choice(model=base_model, custom_llm_provider=self.custom_llm_provider): + ) or litellm.utils.supports_tool_choice( + model=base_model, custom_llm_provider=self.custom_llm_provider + ): # only anthropic and mistral support tool choice config. otherwise (E.g. cohere) will fail the call - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html supported_params.append("tool_choice") @@ -529,7 +546,9 @@ class AmazonConverseConfig(BaseConfig): model=model, custom_llm_provider=self.custom_llm_provider, ) - or supports_reasoning(model=base_model, custom_llm_provider=self.custom_llm_provider) + or supports_reasoning( + model=base_model, custom_llm_provider=self.custom_llm_provider + ) ): supported_params.append("thinking") supported_params.append("reasoning_effort") @@ -554,7 +573,9 @@ class AmazonConverseConfig(BaseConfig): return ToolChoiceValuesBlock(auto={}) elif isinstance(tool_choice, dict): # only supported for anthropic + mistral models - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_ToolChoice.html - specific_tool = SpecificToolChoiceBlock(name=tool_choice.get("function", {}).get("name", "")) + specific_tool = SpecificToolChoiceBlock( + name=tool_choice.get("function", {}).get("name", "") + ) return ToolChoiceValuesBlock(tool=specific_tool) else: raise litellm.utils.UnsupportedParamsError( @@ -574,9 +595,15 @@ class AmazonConverseConfig(BaseConfig): return ["mp4", "mov", "mkv", "webm", "flv", "mpeg", "mpg", "wmv", "3gp"] def get_all_supported_content_types(self) -> List[str]: - return self.get_supported_image_types() + self.get_supported_document_types() + self.get_supported_video_types() + return ( + self.get_supported_image_types() + + self.get_supported_document_types() + + self.get_supported_video_types() + ) - def is_computer_use_tool_used(self, tools: Optional[List[OpenAIChatCompletionToolParam]], model: str) -> bool: + def is_computer_use_tool_used( + self, tools: Optional[List[OpenAIChatCompletionToolParam]], model: str + ) -> bool: """Check if computer use tools are being used in the request.""" if tools is None: return False @@ -589,7 +616,9 @@ class AmazonConverseConfig(BaseConfig): return True return False - def _transform_computer_use_tools(self, computer_use_tools: List[OpenAIChatCompletionToolParam]) -> List[dict]: + def _transform_computer_use_tools( + self, computer_use_tools: List[OpenAIChatCompletionToolParam] + ) -> List[dict]: """Transform computer use tools to Bedrock format.""" transformed_tools: List[dict] = [] @@ -631,7 +660,9 @@ class AmazonConverseConfig(BaseConfig): def _separate_computer_use_tools( self, tools: List[OpenAIChatCompletionToolParam], model: str - ) -> Tuple[List[OpenAIChatCompletionToolParam], List[OpenAIChatCompletionToolParam]]: + ) -> Tuple[ + List[OpenAIChatCompletionToolParam], List[OpenAIChatCompletionToolParam] + ]: """ Separate computer use tools from regular function tools. @@ -740,18 +771,25 @@ class AmazonConverseConfig(BaseConfig): # Recurse into nested schemas if "properties" in result and isinstance(result["properties"], dict): result["properties"] = { - k: AmazonConverseConfig._add_additional_properties_to_schema(v) for k, v in result["properties"].items() + k: AmazonConverseConfig._add_additional_properties_to_schema(v) + for k, v in result["properties"].items() } if "items" in result and isinstance(result["items"], dict): - result["items"] = AmazonConverseConfig._add_additional_properties_to_schema(result["items"]) + result["items"] = AmazonConverseConfig._add_additional_properties_to_schema( + result["items"] + ) for defs_key in ("$defs", "definitions"): if defs_key in result and isinstance(result[defs_key], dict): result[defs_key] = { - k: AmazonConverseConfig._add_additional_properties_to_schema(v) for k, v in result[defs_key].items() + k: AmazonConverseConfig._add_additional_properties_to_schema(v) + for k, v in result[defs_key].items() } for key in ("anyOf", "allOf", "oneOf"): if key in result and isinstance(result[key], list): - result[key] = [AmazonConverseConfig._add_additional_properties_to_schema(item) for item in result[key]] + result[key] = [ + AmazonConverseConfig._add_additional_properties_to_schema(item) + for item in result[key] + ] return result @@ -781,7 +819,9 @@ class AmazonConverseConfig(BaseConfig): } """ if json_schema is not None: - json_schema = AmazonConverseConfig._add_additional_properties_to_schema(json_schema) + json_schema = AmazonConverseConfig._add_additional_properties_to_schema( + json_schema + ) schema_str = json.dumps(json_schema) if json_schema is not None else "{}" json_schema_def: JsonSchemaDefinition = {"schema": schema_str} if name is not None: @@ -803,9 +843,14 @@ class AmazonConverseConfig(BaseConfig): non_default_params: dict, optional_params: dict, ): - optional_params = self._add_tools_to_optional_params(optional_params=optional_params, tools=tools) + optional_params = self._add_tools_to_optional_params( + optional_params=optional_params, tools=tools + ) - if "meta.llama3-3-70b-instruct-v1:0" in model and non_default_params.get("stream", False) is True: + if ( + "meta.llama3-3-70b-instruct-v1:0" in model + and non_default_params.get("stream", False) is True + ): optional_params["fake_stream"] = True def map_openai_params( @@ -944,7 +989,10 @@ class AmazonConverseConfig(BaseConfig): if "type" in value and value["type"] == "text": return optional_params - if self._supports_native_structured_outputs(model, self.custom_llm_provider) and json_schema is not None: + if ( + self._supports_native_structured_outputs(model, self.custom_llm_provider) + and json_schema is not None + ): # Use Bedrock's native structured outputs API (outputConfig.textFormat) # No synthetic tool injection, no fake_stream needed. # Requires an explicit schema — json_object with no schema falls through @@ -962,10 +1010,14 @@ class AmazonConverseConfig(BaseConfig): json_schema=json_schema, description=description, ) - optional_params = self._add_tools_to_optional_params(optional_params=optional_params, tools=[_tool]) + optional_params = self._add_tools_to_optional_params( + optional_params=optional_params, tools=[_tool] + ) if ( - litellm.utils.supports_tool_choice(model=model, custom_llm_provider=self.custom_llm_provider) + litellm.utils.supports_tool_choice( + model=model, custom_llm_provider=self.custom_llm_provider + ) and not is_thinking_enabled ): optional_params["tool_choice"] = ToolChoiceValuesBlock( @@ -977,7 +1029,9 @@ class AmazonConverseConfig(BaseConfig): optional_params["json_mode"] = True return optional_params - def update_optional_params_with_thinking_tokens(self, non_default_params: dict, optional_params: dict): + def update_optional_params_with_thinking_tokens( + self, non_default_params: dict, optional_params: dict + ): """ Handles scenario where max tokens is not specified. For anthropic models (anthropic api/bedrock/vertex ai), this requires having the max tokens being set and being greater than the thinking token budget. @@ -995,9 +1049,13 @@ class AmazonConverseConfig(BaseConfig): is_thinking_enabled = self.is_thinking_enabled(optional_params) is_max_tokens_in_request = self.is_max_tokens_in_request(non_default_params) if is_thinking_enabled and not is_max_tokens_in_request: - thinking_token_budget = cast(dict, optional_params["thinking"]).get("budget_tokens", None) + thinking_token_budget = cast(dict, optional_params["thinking"]).get( + "budget_tokens", None + ) if thinking_token_budget is not None: - optional_params["maxTokens"] = thinking_token_budget + DEFAULT_MAX_TOKENS + optional_params["maxTokens"] = ( + thinking_token_budget + DEFAULT_MAX_TOKENS + ) @overload def _get_cache_point_block( @@ -1063,15 +1121,23 @@ class AmazonConverseConfig(BaseConfig): if message["role"] == "system": system_prompt_indices.append(idx) if isinstance(message["content"], str) and message["content"]: - system_content_blocks.append(SystemContentBlock(text=message["content"])) - cache_block = self._get_cache_point_block(message, block_type="system", model=model) + system_content_blocks.append( + SystemContentBlock(text=message["content"]) + ) + cache_block = self._get_cache_point_block( + message, block_type="system", model=model + ) if cache_block: system_content_blocks.append(cache_block) elif isinstance(message["content"], list): for m in message["content"]: if m.get("type") == "text" and m.get("text"): - system_content_blocks.append(SystemContentBlock(text=m["text"])) - cache_block = self._get_cache_point_block(m, block_type="system", model=model) + system_content_blocks.append( + SystemContentBlock(text=m["text"]) + ) + cache_block = self._get_cache_point_block( + m, block_type="system", model=model + ) if cache_block: system_content_blocks.append(cache_block) if len(system_prompt_indices) > 0: @@ -1109,10 +1175,16 @@ class AmazonConverseConfig(BaseConfig): # Exceptions should not be stored in optional_params (this is a defensive fix) cleaned_params = filter_exceptions_from_params(optional_params) inference_params = safe_deep_copy(cleaned_params) - supported_converse_params = list(AmazonConverseConfig.__annotations__.keys()) + ["top_k"] + supported_converse_params = list( + AmazonConverseConfig.__annotations__.keys() + ) + ["top_k"] supported_tool_call_params = ["tools", "tool_choice"] supported_config_params = list(self.get_config_blocks().keys()) - total_supported_params = supported_converse_params + supported_tool_call_params + supported_config_params + total_supported_params = ( + supported_converse_params + + supported_tool_call_params + + supported_config_params + ) inference_params.pop("json_mode", None) # used for handling json_schema # Anthropic-only key. Bedrock expects `outputConfig` (camelCase) and # will reject `output_config` if it leaks through pass-through routes. @@ -1123,15 +1195,25 @@ class AmazonConverseConfig(BaseConfig): if request_metadata is not None: self._validate_request_metadata(request_metadata) - output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None) - inference_params.pop("output_config", None) # Bedrock Converse doesn't support it + output_config: Optional[OutputConfigBlock] = inference_params.pop( + "outputConfig", None + ) + inference_params.pop( + "output_config", None + ) # Bedrock Converse doesn't support it # keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params' - additional_request_params = {k: v for k, v in inference_params.items() if k not in total_supported_params} - inference_params = {k: v for k, v in inference_params.items() if k in total_supported_params} + additional_request_params = { + k: v for k, v in inference_params.items() if k not in total_supported_params + } + inference_params = { + k: v for k, v in inference_params.items() if k in total_supported_params + } # Handle parallel_tool_calls configuration - parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None) + parallel_tool_use_config = additional_request_params.pop( + "_parallel_tool_use_config", None + ) if parallel_tool_use_config is not None and is_claude_4_5_on_bedrock(model): for key, value in parallel_tool_use_config.items(): if ( @@ -1146,7 +1228,9 @@ class AmazonConverseConfig(BaseConfig): additional_request_params.pop("parallel_tool_calls", None) # Only set the topK value in for models that support it - additional_request_params.update(self._handle_top_k_value(model, inference_params)) + additional_request_params.update( + self._handle_top_k_value(model, inference_params) + ) # Filter out internal/MCP-related parameters that shouldn't be sent to the API # These are LiteLLM internal parameters, not API parameters @@ -1155,7 +1239,9 @@ class AmazonConverseConfig(BaseConfig): # Filter out non-serializable objects (exceptions, callables, logging objects, etc.) # from additional_request_params to prevent JSON serialization errors # This filters: Exception objects, callable objects (functions), Logging objects, etc. - additional_request_params = filter_exceptions_from_params(additional_request_params) + additional_request_params = filter_exceptions_from_params( + additional_request_params + ) return ( inference_params, @@ -1202,7 +1288,9 @@ class AmazonConverseConfig(BaseConfig): # Only separate tools if computer use tools are actually present if filtered_tools and self.is_computer_use_tool_used(filtered_tools, model): # Separate computer use tools from regular function tools - computer_use_tools, regular_tools = self._separate_computer_use_tools(filtered_tools, model) + computer_use_tools, regular_tools = self._separate_computer_use_tools( + filtered_tools, model + ) # Process regular function tools using existing logic bedrock_tools = _bedrock_tools_pt(regular_tools) @@ -1263,7 +1351,9 @@ class AmazonConverseConfig(BaseConfig): anthropic_beta_list.append(computer_use_header) # Transform computer use tools to proper Bedrock format - transformed_computer_tools = self._transform_computer_use_tools(computer_use_tools) + transformed_computer_tools = self._transform_computer_use_tools( + computer_use_tools + ) additional_request_params["tools"] = transformed_computer_tools else: # No computer use tools, process all tools as regular tools @@ -1292,9 +1382,15 @@ class AmazonConverseConfig(BaseConfig): """ Bedrock doesn't support tool calling without `tools=` param specified. """ - if "tools" not in optional_params and messages is not None and has_tool_call_blocks(messages): + if ( + "tools" not in optional_params + and messages is not None + and has_tool_call_blocks(messages) + ): if litellm.modify_params: - optional_params["tools"] = add_dummy_tool(custom_llm_provider="bedrock_converse") + optional_params["tools"] = add_dummy_tool( + custom_llm_provider="bedrock_converse" + ) else: raise litellm.UnsupportedParamsError( message="Bedrock doesn't support tool calling without `tools=` param specified. Pass `tools=` param OR set `litellm.modify_params = True` // `litellm_settings::modify_params: True` to add dummy tool to the request.", @@ -1348,7 +1444,9 @@ class AmazonConverseConfig(BaseConfig): bedrock_tool_config: Optional[ToolConfigBlock] = None if len(bedrock_tools) > 0: - tool_choice_values: ToolChoiceValuesBlock = inference_params.pop("tool_choice", None) + tool_choice_values: ToolChoiceValuesBlock = inference_params.pop( + "tool_choice", None + ) bedrock_tool_config = ToolConfigBlock( tools=bedrock_tools, ) @@ -1358,7 +1456,9 @@ class AmazonConverseConfig(BaseConfig): data: CommonRequestObject = { "additionalModelRequestFields": additional_request_params, "system": system_content_blocks, - "inferenceConfig": self._transform_inference_params(inference_params=inference_params), + "inferenceConfig": self._transform_inference_params( + inference_params=inference_params + ), } # Handle all config blocks @@ -1388,10 +1488,14 @@ class AmazonConverseConfig(BaseConfig): litellm_params: dict, headers: Optional[dict] = None, ) -> RequestObject: - messages, system_content_blocks = self._transform_system_message(messages, model=model) + messages, system_content_blocks = self._transform_system_message( + messages, model=model + ) # Convert last user message to guarded_text if guardrailConfig is present - messages = self._convert_consecutive_user_messages_to_guarded_text(messages, optional_params) + messages = self._convert_consecutive_user_messages_to_guarded_text( + messages, optional_params + ) ## TRANSFORMATION ## _data: CommonRequestObject = self._transform_request_helper( @@ -1402,11 +1506,13 @@ class AmazonConverseConfig(BaseConfig): headers=headers, ) - bedrock_messages = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( - messages=messages, - model=model, - llm_provider="bedrock_converse", - user_continue_message=litellm_params.pop("user_continue_message", None), + bedrock_messages = ( + await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model=model, + llm_provider="bedrock_converse", + user_continue_message=litellm_params.pop("user_continue_message", None), + ) ) data: RequestObject = {"messages": bedrock_messages, **_data} @@ -1440,10 +1546,14 @@ class AmazonConverseConfig(BaseConfig): litellm_params: dict, headers: Optional[dict] = None, ) -> RequestObject: - messages, system_content_blocks = self._transform_system_message(messages, model=model) + messages, system_content_blocks = self._transform_system_message( + messages, model=model + ) # Convert last user message to guarded_text if guardrailConfig is present - messages = self._convert_consecutive_user_messages_to_guarded_text(messages, optional_params) + messages = self._convert_consecutive_user_messages_to_guarded_text( + messages, optional_params + ) _data: CommonRequestObject = self._transform_request_helper( model=model, @@ -1492,7 +1602,9 @@ class AmazonConverseConfig(BaseConfig): encoding=encoding, ) - def _transform_reasoning_content(self, reasoning_content_blocks: List[BedrockConverseReasoningContentBlock]) -> str: + def _transform_reasoning_content( + self, reasoning_content_blocks: List[BedrockConverseReasoningContentBlock] + ) -> str: """ Extract the reasoning text from the reasoning content blocks @@ -1508,7 +1620,9 @@ class AmazonConverseConfig(BaseConfig): self, thinking_blocks: List[BedrockConverseReasoningContentBlock] ) -> List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]: """Return a consistent format for thinking blocks between Anthropic and Bedrock.""" - thinking_blocks_list: List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]] = [] + thinking_blocks_list: List[ + Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock] + ] = [] for block in thinking_blocks: if "reasoningText" in block: _thinking_block = ChatCompletionThinkingBlock(type="thinking") @@ -1544,11 +1658,21 @@ class AmazonConverseConfig(BaseConfig): cache_creation_input_tokens = usage["cacheWriteInputTokens"] input_tokens += cache_creation_input_tokens - prompt_tokens_details = PromptTokensDetailsWrapper(cached_tokens=cache_read_input_tokens) - reasoning_tokens = token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0 + prompt_tokens_details = PromptTokensDetailsWrapper( + cached_tokens=cache_read_input_tokens + ) + reasoning_tokens = ( + token_counter(text=reasoning_content, count_response_tokens=True) + if reasoning_content + else 0 + ) completion_tokens_details = CompletionTokensDetailsWrapper( reasoning_tokens=reasoning_tokens, - text_tokens=(output_tokens - reasoning_tokens if reasoning_tokens > 0 else output_tokens), + text_tokens=( + output_tokens - reasoning_tokens + if reasoning_tokens > 0 + else output_tokens + ), ) openai_usage = Usage( prompt_tokens=input_tokens, @@ -1563,7 +1687,9 @@ class AmazonConverseConfig(BaseConfig): def get_tool_call_names( self, - tools: Optional[Union[List[ToolBlock], List[OpenAIChatCompletionToolParam]]] = None, + tools: Optional[ + Union[List[ToolBlock], List[OpenAIChatCompletionToolParam]] + ] = None, ) -> List[str]: if tools is None: return [] @@ -1602,8 +1728,13 @@ class AmazonConverseConfig(BaseConfig): try: tool_call_names = self.get_tool_call_names(tools) json_content = json.loads(message.content) - if json_content.get("type") == "function" and json_content.get("name") in tool_call_names: - tool_calls = [ChatCompletionMessageToolCall(function=Function(**json_content))] + if ( + json_content.get("type") == "function" + and json_content.get("name") in tool_call_names + ): + tool_calls = [ + ChatCompletionMessageToolCall(function=Function(**json_content)) + ] message.tool_calls = tool_calls message.content = None @@ -1613,9 +1744,7 @@ class AmazonConverseConfig(BaseConfig): return message, returned_finish_reason - def _translate_message_content( - self, content_blocks: List[ContentBlock] - ) -> Tuple[ + def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[ str, List[ChatCompletionToolCallChunk], Optional[List[BedrockConverseReasoningContentBlock]], @@ -1632,7 +1761,9 @@ class AmazonConverseConfig(BaseConfig): """ content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = None + reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( + None + ) citationsContentBlocks: Optional[List[CitationsContentBlock]] = None for idx, content in enumerate(content_blocks): """ @@ -1649,7 +1780,9 @@ class AmazonConverseConfig(BaseConfig): if "toolUse" in content: ## check tool name was formatted by litellm _response_tool_name = content["toolUse"]["name"] - response_tool_name = get_bedrock_tool_name(response_tool_name=_response_tool_name) + response_tool_name = get_bedrock_tool_name( + response_tool_name=_response_tool_name + ) _function_chunk = ChatCompletionToolCallFunctionChunk( name=response_tool_name, arguments=json.dumps(content["toolUse"]["input"]), @@ -1700,7 +1833,11 @@ class AmazonConverseConfig(BaseConfig): """ try: response_data = json.loads(json_str) - if isinstance(response_data, dict) and "properties" in response_data and len(response_data) == 1: + if ( + isinstance(response_data, dict) + and "properties" in response_data + and len(response_data) == 1 + ): response_data = response_data["properties"] return json.dumps(response_data) except json.JSONDecodeError: @@ -1724,7 +1861,11 @@ class AmazonConverseConfig(BaseConfig): if not json_mode or not tools: return tools if tools else None - json_tool_indices = [i for i, t in enumerate(tools) if t["function"].get("name") == RESPONSE_FORMAT_TOOL_NAME] + json_tool_indices = [ + i + for i, t in enumerate(tools) + if t["function"].get("name") == RESPONSE_FORMAT_TOOL_NAME + ] if not json_tool_indices: # No json_tool_call found, return tools unchanged @@ -1732,10 +1873,14 @@ class AmazonConverseConfig(BaseConfig): if len(json_tool_indices) == len(tools): # All tools are json_tool_call — convert first one to content - verbose_logger.debug("Processing JSON tool call response for response_format") + verbose_logger.debug( + "Processing JSON tool call response for response_format" + ) json_mode_content_str: Optional[str] = tools[0]["function"].get("arguments") if json_mode_content_str is not None: - json_mode_content_str = AmazonConverseConfig._unwrap_bedrock_properties(json_mode_content_str) + json_mode_content_str = AmazonConverseConfig._unwrap_bedrock_properties( + json_mode_content_str + ) chat_completion_message["content"] = json_mode_content_str return None @@ -1745,9 +1890,13 @@ class AmazonConverseConfig(BaseConfig): first_idx = json_tool_indices[0] json_mode_args = tools[first_idx]["function"].get("arguments") if json_mode_args is not None: - json_mode_args = AmazonConverseConfig._unwrap_bedrock_properties(json_mode_args) + json_mode_args = AmazonConverseConfig._unwrap_bedrock_properties( + json_mode_args + ) existing = chat_completion_message.get("content") or "" - chat_completion_message["content"] = existing + json_mode_args if existing else json_mode_args + chat_completion_message["content"] = ( + existing + json_mode_args if existing else json_mode_args + ) real_tools = [t for i, t in enumerate(tools) if i not in json_tool_indices] return real_tools if real_tools else None @@ -1825,7 +1974,9 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = None + reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( + None + ) citationsContentBlocks: Optional[List[CitationsContentBlock]] = None if message is not None: @@ -1844,11 +1995,17 @@ class AmazonConverseConfig(BaseConfig): provider_specific_fields["citationsContent"] = citationsContentBlocks if provider_specific_fields: - chat_completion_message["provider_specific_fields"] = provider_specific_fields + chat_completion_message["provider_specific_fields"] = ( + provider_specific_fields + ) if reasoningContentBlocks is not None: - chat_completion_message["reasoning_content"] = self._transform_reasoning_content(reasoningContentBlocks) - chat_completion_message["thinking_blocks"] = self._transform_thinking_blocks(reasoningContentBlocks) + chat_completion_message["reasoning_content"] = ( + self._transform_reasoning_content(reasoningContentBlocks) + ) + chat_completion_message["thinking_blocks"] = ( + self._transform_thinking_blocks(reasoningContentBlocks) + ) chat_completion_message["content"] = content_str filtered_tools = self._filter_json_mode_tools( json_mode=json_mode, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 0efc3dab638..ed600b26e83 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -8,6 +8,7 @@ Run checks for: 2. If user is in budget 3. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget """ + import asyncio import re import time @@ -414,9 +415,9 @@ async def common_checks( # noqa: PLR0915 model=_model, team_object=team_object, llm_router=llm_router, - team_model_aliases=valid_token.team_model_aliases - if valid_token - else None, + team_model_aliases=( + valid_token.team_model_aliases if valid_token else None + ), ): raise ProxyException( message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", @@ -3078,10 +3079,7 @@ async def _team_max_budget_check( BudgetExceededError if the team is over it's max budget. Triggers a budget alert if the team is over it's max budget. """ - if ( - team_object is not None - and team_object.max_budget is not None - ): + if team_object is not None and team_object.max_budget is not None: from litellm.proxy.proxy_server import get_current_spend # Read spend from cross-pod counter (Redis-first) or cached object (fallback) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 5db919659ed..bcfaed24398 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -66,10 +66,8 @@ class ResetBudgetJob: try: from litellm.proxy.proxy_server import spend_counter_cache - memberships = ( - await self.prisma_client.db.litellm_teammembership.find_many( - where={"budget_id": {"in": budget_ids}} - ) + memberships = await self.prisma_client.db.litellm_teammembership.find_many( + where={"budget_id": {"in": budget_ids}} ) for m in memberships: counter_key = f"spend:team_member:{m.user_id}:{m.team_id}" @@ -574,7 +572,11 @@ class ResetBudgetJob: counter_key = None if item_type == "key" and hasattr(item, "token") and item.token is not None: counter_key = f"spend:key:{item.token}" - elif item_type == "team" and hasattr(item, "team_id") and item.team_id is not None: + elif ( + item_type == "team" + and hasattr(item, "team_id") + and item.team_id is not None + ): counter_key = f"spend:team:{item.team_id}" if counter_key is not None: diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index a5a62813b9f..02d7b9baf4a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -722,9 +722,9 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: import ast try: - generic_user_role_mappings_data: Dict[ - LitellmUserRoles, List[str] - ] = ast.literal_eval(generic_role_mappings) + generic_user_role_mappings_data: Dict[LitellmUserRoles, List[str]] = ( + ast.literal_eval(generic_role_mappings) + ) if isinstance(generic_user_role_mappings_data, dict): from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings @@ -827,7 +827,9 @@ async def get_generic_sso_response( ], # sso specific jwt handler - used for restricted sso group access control generic_client_id: str, redirect_url: str, -) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload) +) -> Tuple[ + Union[OpenID, dict], Optional[dict], Optional[dict] +]: # (result, received_response, access_token_payload) # make generic sso provider from fastapi_sso.sso.base import DiscoveryDocument from fastapi_sso.sso.generic import create_provider @@ -879,9 +881,9 @@ async def get_generic_sso_response( verbose_proxy_logger.debug("calling generic_sso.verify_and_process") additional_generic_sso_headers_dict = _parse_generic_sso_headers() - code_verifier: Optional[ - str - ] = None # assigned inside try; initialized for type tracking + code_verifier: Optional[str] = ( + None # assigned inside try; initialized for type tracking + ) access_token_payload: Optional[dict] = None # decoded JWT access token claims try: @@ -1231,9 +1233,11 @@ async def _sync_user_role_from_jwt_role_map( user_info.user_role = mapped_role.value await user_api_key_cache.async_set_cache( key=user_info.user_id, - value=user_info.model_dump() - if hasattr(user_info, "model_dump") - else dict(user_info), + value=( + user_info.model_dump() + if hasattr(user_info, "model_dump") + else dict(user_info) + ), ) @@ -1259,9 +1263,9 @@ def apply_user_info_values_to_sso_user_defined_values( else: # SSO didn't provide a valid role, fall back to DB role or default if user_info is None or user_info.user_role is None: - user_defined_values[ - "user_role" - ] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value + user_defined_values["user_role"] = ( + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value + ) verbose_proxy_logger.debug( "No SSO or DB role found, using default: INTERNAL_USER_VIEW_ONLY" ) @@ -1369,12 +1373,14 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa: ) elif generic_client_id is not None: - result, received_response, access_token_payload = await get_generic_sso_response( - request=request, - jwt_handler=jwt_handler, - generic_client_id=generic_client_id, - redirect_url=redirect_url, - sso_jwt_handler=sso_jwt_handler, + result, received_response, access_token_payload = ( + await get_generic_sso_response( + request=request, + jwt_handler=jwt_handler, + generic_client_id=generic_client_id, + redirect_url=redirect_url, + sso_jwt_handler=sso_jwt_handler, + ) ) if result is None: @@ -1697,9 +1703,9 @@ async def insert_sso_user( if user_defined_values.get("max_budget") is None: user_defined_values["max_budget"] = litellm.max_internal_user_budget if user_defined_values.get("budget_duration") is None: - user_defined_values[ - "budget_duration" - ] = litellm.internal_user_budget_duration + user_defined_values["budget_duration"] = ( + litellm.internal_user_budget_duration + ) if user_defined_values["user_role"] is None: user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY @@ -3342,9 +3348,9 @@ class MicrosoftSSOHandler: # if user is trying to get the raw sso response for debugging, return the raw sso response if return_raw_sso_response: - original_msft_result[ - MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY - ] = user_team_ids + original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = ( + user_team_ids + ) original_msft_result["app_roles"] = app_roles return original_msft_result or {} @@ -3463,9 +3469,9 @@ class MicrosoftSSOHandler: # Fetch user membership from Microsoft Graph API all_group_ids = [] - next_link: Optional[ - str - ] = MicrosoftSSOHandler.graph_api_user_groups_endpoint + next_link: Optional[str] = ( + MicrosoftSSOHandler.graph_api_user_groups_endpoint + ) auth_headers = {"Authorization": f"Bearer {access_token}"} page_count = 0 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ab73d44acca..317b8f76179 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -639,9 +639,9 @@ except ImportError: server_root_path = get_server_root_path() _license_check = LicenseCheck() premium_user: bool = _license_check.is_premium() -premium_user_data: Optional[ - "EnterpriseLicenseData" -] = _license_check.airgapped_license_data +premium_user_data: Optional["EnterpriseLicenseData"] = ( + _license_check.airgapped_license_data +) global_max_parallel_request_retries_env: Optional[str] = os.getenv( "LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES" ) @@ -1524,9 +1524,9 @@ master_key: Optional[str] = None config_agents: Optional[List[AgentConfig]] = None otel_logging = False prisma_client: Optional[PrismaClient] = None -shared_aiohttp_session: Optional[ - "ClientSession" -] = None # Global shared session for connection reuse +shared_aiohttp_session: Optional["ClientSession"] = ( + None # Global shared session for connection reuse +) user_api_key_cache = DualCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) @@ -1537,13 +1537,13 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter( dual_cache=user_api_key_cache ) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) -redis_usage_cache: Optional[ - RedisCache -] = None # redis cache used for tracking spend, tpm/rpm limits +redis_usage_cache: Optional[RedisCache] = ( + None # redis cache used for tracking spend, tpm/rpm limits +) polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False -native_background_mode: List[ - str -] = [] # Models that should use native provider background mode instead of polling +native_background_mode: List[str] = ( + [] +) # Models that should use native provider background mode instead of polling polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache user_custom_auth = None user_custom_key_generate = None @@ -1714,9 +1714,7 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float: # 1. Try Redis first (cross-pod authoritative) if spend_counter_cache.redis_cache is not None: try: - val = await spend_counter_cache.redis_cache.async_get_cache( - key=counter_key - ) + val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key) if val is not None: return float(val) except Exception as e: @@ -1820,9 +1818,7 @@ async def _init_and_increment_spend_counter( key=counter_key, value=base_spend ) - await spend_counter_cache.async_increment_cache( - key=counter_key, value=increment - ) + await spend_counter_cache.async_increment_cache(key=counter_key, value=increment) async def update_cache( # noqa: PLR0915 @@ -2031,9 +2027,9 @@ async def update_cache( # noqa: PLR0915 _id = "team_id:{}".format(team_id) try: # Fetch the existing cost for the given user - existing_spend_obj: Optional[ - LiteLLM_TeamTable - ] = await user_api_key_cache.async_get_cache(key=_id) + existing_spend_obj: Optional[LiteLLM_TeamTable] = ( + await user_api_key_cache.async_get_cache(key=_id) + ) if existing_spend_obj is None: # do nothing if team not in api key cache return @@ -2154,11 +2150,9 @@ def run_ollama_serve(): with open(os.devnull, "w") as devnull: subprocess.Popen(command, stdout=devnull, stderr=devnull) except Exception as e: - verbose_proxy_logger.debug( - f""" + verbose_proxy_logger.debug(f""" LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve` - """ - ) + """) def _get_process_rss_mb() -> Optional[float]: @@ -5206,10 +5200,10 @@ class ProxyConfig: ) try: - guardrails_in_db: List[ - Guardrail - ] = await GuardrailRegistry.get_all_guardrails_from_db( - prisma_client=prisma_client + guardrails_in_db: List[Guardrail] = ( + await GuardrailRegistry.get_all_guardrails_from_db( + prisma_client=prisma_client + ) ) verbose_proxy_logger.debug( "guardrails from the DB %s", str(guardrails_in_db) @@ -5591,9 +5585,9 @@ async def initialize( # noqa: PLR0915 user_api_base = api_base dynamic_config[user_model]["api_base"] = api_base if api_version: - os.environ[ - "AZURE_API_VERSION" - ] = api_version # set this for azure - litellm can read this from the env + os.environ["AZURE_API_VERSION"] = ( + api_version # set this for azure - litellm can read this from the env + ) if max_tokens: # model-specific param dynamic_config[user_model]["max_tokens"] = max_tokens if temperature: # model-specific param @@ -5930,9 +5924,9 @@ class ProxyStartupEvent: """ from litellm.secret_managers.main import str_to_bool - _use_redis_transaction_buffer: Optional[ - Union[bool, str] - ] = general_settings.get("use_redis_transaction_buffer", False) + _use_redis_transaction_buffer: Optional[Union[bool, str]] = ( + general_settings.get("use_redis_transaction_buffer", False) + ) if isinstance(_use_redis_transaction_buffer, str): _use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer) @@ -11318,14 +11312,11 @@ async def login_v2(request: Request): # noqa: PLR0915 litellm_dashboard_ui += "/ui/" litellm_dashboard_ui += "?login=success" - # Token is included in the response body so the UI can set a JS-accessible - # cookie even when a reverse proxy (e.g. nginx-ingress) adds HttpOnly to the - # server-set cookie, which would otherwise cause an infinite login redirect. json_response = JSONResponse( - content={"redirect_url": litellm_dashboard_ui, "token": jwt_token}, + content={"redirect_url": litellm_dashboard_ui}, status_code=status.HTTP_200_OK, ) - json_response.set_cookie(key="token", value=jwt_token) + json_response.set_cookie(key="token", value=jwt_token, httponly=True) return json_response except Exception as e: verbose_proxy_logger.exception( @@ -12530,9 +12521,9 @@ async def get_config_list( hasattr(sub_field_info, "description") and sub_field_info.description is not None ): - nested_fields[ - idx - ].field_description = sub_field_info.description + nested_fields[idx].field_description = ( + sub_field_info.description + ) idx += 1 _stored_in_db = None diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index ca70b4bfa8e..74c7f9bca5f 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -28,7 +28,7 @@ def test_get_api_key(): assert get_api_key( custom_litellm_key_header=None, api_key=bearer_token, - AZURE_AI_API_KEY_header=None, + azure_api_key_header=None, anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, @@ -59,7 +59,7 @@ def test_get_api_key_with_custom_litellm_key_header( assert get_api_key( custom_litellm_key_header=custom_litellm_key_header, api_key=None, - AZURE_AI_API_KEY_header=None, + azure_api_key_header=None, anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, @@ -371,7 +371,7 @@ async def test_proxy_admin_expired_key_from_cache(): await _user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", # Add Bearer prefix - AZURE_AI_API_KEY_header="", + azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None, @@ -842,7 +842,7 @@ async def test_user_api_key_auth_builder_no_blocking_calls(): await _user_api_key_auth_builder( request=request, api_key=f"Bearer {api_key}", - AZURE_AI_API_KEY_header="", + azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, azure_apim_header=None,