Merge pull request #14982 from BerriAI/litellm_ci_cd_linting_fixes_09_29_2025_p2

Litellm ci cd linting fixes 09 29 2025 p2
This commit is contained in:
Krish Dholakia 2025-09-27 14:41:39 -07:00 • committed by GitHub
commit a7470b3291
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
38 changed files with 646 additions and 563 deletions

View file

@ -61,7 +61,7 @@ jobs:
- name: Run MyPy type checking
run: |
cd litellm
poetry run mypy . --ignore-missing-imports --disable-error-code=var-annotated
poetry run mypy .
cd ..
- name: Check for circular imports

View file

@ -2262,9 +2262,12 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]:
keys_parts = key.split(".")
# Traverse through the dictionary using the parts
value = metadata
value: Any = metadata
for part in keys_parts:
value = value.get(part, None) # Get the value, return None if not found
if isinstance(value, dict):
value = value.get(part, None) # Get the value, return None if not found
else:
value = None
if value is None:
break

View file

@ -377,7 +377,9 @@ public_model_groups: Optional[List[str]] = None
public_model_groups_links: Dict[str, str] = {}
#### REQUEST PRIORITIZATION ######
priority_reservation: Optional[Dict[str, float]] = None
priority_reservation_settings: "PriorityReservationSettings" = PriorityReservationSettings()
priority_reservation_settings: "PriorityReservationSettings" = (
PriorityReservationSettings()
)
######## Networking Settings ########
@ -443,7 +445,7 @@ def identify(event_details):
####### ADDITIONAL PARAMS ################### configurable params if you use proxy models like Helicone, map spend to org id, etc.
api_base: Optional[str] = None
headers = None
api_version = None
api_version: Optional[str] = None
organization = None
project = None
config_path = None
@ -494,7 +496,7 @@ azure_ai_models: Set = set()
jina_ai_models: Set = set()
voyage_models: Set = set()
infinity_models: Set = set()
heroku_models: Set = set()
heroku_models: Set = set()
databricks_models: Set = set()
cloudflare_models: Set = set()
codestral_models: Set = set()
@ -1357,6 +1359,7 @@ from .passthrough import allm_passthrough_route, llm_passthrough_route
### GLOBAL CONFIG ###
global_bitbucket_config: Optional[Dict[str, Any]] = None
def set_global_bitbucket_config(config: Dict[str, Any]) -> None:
"""Set global BitBucket configuration for prompt management."""
global global_bitbucket_config

View file

@ -301,9 +301,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.litellm_trace_id: str = litellm_trace_id or str(uuid.uuid4())
self.function_id = function_id
self.streaming_chunks: List[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: List[
Any
] = [] # for generating complete stream response
self.sync_streaming_chunks: List[Any] = (
[]
) # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
# Initialize dynamic callbacks
@ -344,7 +344,7 @@ class Logging(LiteLLMLoggingBaseClass):
litellm_params = scrub_sensitive_keys_in_metadata(litellm_params)
self.litellm_params = litellm_params
# Initialize cost breakdown field
self.cost_breakdown: Optional[CostBreakdown] = None
@ -676,9 +676,9 @@ class Logging(LiteLLMLoggingBaseClass):
if anthropic_cache_control_logger := AnthropicCacheControlHook.get_custom_logger_for_anthropic_cache_control_hook(
non_default_params
):
self.model_call_details[
"prompt_integration"
] = anthropic_cache_control_logger.__class__.__name__
self.model_call_details["prompt_integration"] = (
anthropic_cache_control_logger.__class__.__name__
)
return anthropic_cache_control_logger
#########################################################
@ -690,9 +690,9 @@ class Logging(LiteLLMLoggingBaseClass):
internal_usage_cache=None,
llm_router=None,
)
self.model_call_details[
"prompt_integration"
] = vector_store_custom_logger.__class__.__name__
self.model_call_details["prompt_integration"] = (
vector_store_custom_logger.__class__.__name__
)
return vector_store_custom_logger
return None
@ -744,9 +744,9 @@ class Logging(LiteLLMLoggingBaseClass):
model
): # if model name was changes pre-call, overwrite the initial model call name with the new one
self.model_call_details["model"] = model
self.model_call_details["litellm_params"][
"api_base"
] = self._get_masked_api_base(additional_args.get("api_base", ""))
self.model_call_details["litellm_params"]["api_base"] = (
self._get_masked_api_base(additional_args.get("api_base", ""))
)
def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915
# Log the exact input to the LLM API
@ -775,10 +775,10 @@ class Logging(LiteLLMLoggingBaseClass):
try:
# [Non-blocking Extra Debug Information in metadata]
if turn_off_message_logging is True:
_metadata[
"raw_request"
] = "redacted by litellm. \
_metadata["raw_request"] = (
"redacted by litellm. \
'litellm.turn_off_message_logging=True'"
)
else:
curl_command = self._get_request_curl_command(
api_base=additional_args.get("api_base", ""),
@ -789,32 +789,32 @@ class Logging(LiteLLMLoggingBaseClass):
_metadata["raw_request"] = str(curl_command)
# split up, so it's easier to parse in the UI
self.model_call_details[
"raw_request_typed_dict"
] = RawRequestTypedDict(
raw_request_api_base=str(
additional_args.get("api_base") or ""
),
raw_request_body=self._get_raw_request_body(
additional_args.get("complete_input_dict", {})
),
raw_request_headers=self._get_masked_headers(
additional_args.get("headers", {}) or {},
ignore_sensitive_headers=True,
),
error=None,
self.model_call_details["raw_request_typed_dict"] = (
RawRequestTypedDict(
raw_request_api_base=str(
additional_args.get("api_base") or ""
),
raw_request_body=self._get_raw_request_body(
additional_args.get("complete_input_dict", {})
),
raw_request_headers=self._get_masked_headers(
additional_args.get("headers", {}) or {},
ignore_sensitive_headers=True,
),
error=None,
)
)
except Exception as e:
self.model_call_details[
"raw_request_typed_dict"
] = RawRequestTypedDict(
error=str(e),
self.model_call_details["raw_request_typed_dict"] = (
RawRequestTypedDict(
error=str(e),
)
)
_metadata[
"raw_request"
] = "Unable to Log \
_metadata["raw_request"] = (
"Unable to Log \
raw request: {}".format(
str(e)
str(e)
)
)
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
try:
@ -1115,13 +1115,13 @@ class Logging(LiteLLMLoggingBaseClass):
for callback in callbacks:
try:
if isinstance(callback, CustomLogger):
response: Optional[
MCPPostCallResponseObject
] = await callback.async_post_mcp_tool_call_hook(
kwargs=kwargs,
response_obj=post_mcp_tool_call_response_obj,
start_time=start_time,
end_time=end_time,
response: Optional[MCPPostCallResponseObject] = (
await callback.async_post_mcp_tool_call_hook(
kwargs=kwargs,
response_obj=post_mcp_tool_call_response_obj,
start_time=start_time,
end_time=end_time,
)
)
######################################################################
# if any of the callbacks modify the response, use the modified response
@ -1168,19 +1168,19 @@ class Logging(LiteLLMLoggingBaseClass):
) -> None:
"""
Helper method to store cost breakdown in the logging object.
Args:
input_cost: Cost of input/prompt tokens
output_cost: Cost of output/completion tokens
output_cost: Cost of output/completion tokens
cost_for_built_in_tools_cost_usd_dollar: Cost of built-in tools
total_cost: Total cost of request
"""
self.cost_breakdown = CostBreakdown(
input_cost=input_cost,
output_cost=output_cost,
total_cost=total_cost,
tool_usage_cost=cost_for_built_in_tools_cost_usd_dollar
tool_usage_cost=cost_for_built_in_tools_cost_usd_dollar,
)
verbose_logger.debug(
f"Cost breakdown set - input: {input_cost}, output: {output_cost}, cost_for_built_in_tools_cost_usd_dollar: {cost_for_built_in_tools_cost_usd_dollar}, total: {total_cost}"
@ -1259,9 +1259,11 @@ class Logging(LiteLLMLoggingBaseClass):
"standard_built_in_tools_params": self.standard_built_in_tools_params,
"router_model_id": router_model_id,
"litellm_logging_obj": self,
"service_tier": self.optional_params.get("service_tier")
if self.optional_params
else None,
"service_tier": (
self.optional_params.get("service_tier")
if self.optional_params
else None
),
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(
@ -1271,9 +1273,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
self.model_call_details[
"response_cost_failure_debug_information"
] = debug_info
self.model_call_details["response_cost_failure_debug_information"] = (
debug_info
)
return None
try:
@ -1298,9 +1300,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
self.model_call_details[
"response_cost_failure_debug_information"
] = debug_info
self.model_call_details["response_cost_failure_debug_information"] = (
debug_info
)
return None
@ -1444,9 +1446,9 @@ class Logging(LiteLLMLoggingBaseClass):
end_time = datetime.datetime.now()
if self.completion_start_time is None:
self.completion_start_time = end_time
self.model_call_details[
"completion_start_time"
] = self.completion_start_time
self.model_call_details["completion_start_time"] = (
self.completion_start_time
)
self.model_call_details["log_event_type"] = "successful_api_call"
self.model_call_details["end_time"] = end_time
self.model_call_details["cache_hit"] = cache_hit
@ -1499,39 +1501,39 @@ class Logging(LiteLLMLoggingBaseClass):
"response_cost"
]
else:
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(result=logging_result)
self.model_call_details["response_cost"] = (
self._response_cost_calculator(result=logging_result)
)
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=logging_result,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=logging_result,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
elif isinstance(result, dict) or isinstance(result, list):
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=result,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=result,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
elif standard_logging_object is not None:
self.model_call_details[
"standard_logging_object"
] = standard_logging_object
self.model_call_details["standard_logging_object"] = (
standard_logging_object
)
else: # streaming chunks + image gen.
self.model_call_details["response_cost"] = None
@ -1682,23 +1684,23 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
"Logging Details LiteLLM-Success Call streaming complete"
)
self.model_call_details[
"complete_streaming_response"
] = complete_streaming_response
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(result=complete_streaming_response)
self.model_call_details["complete_streaming_response"] = (
complete_streaming_response
)
self.model_call_details["response_cost"] = (
self._response_cost_calculator(result=complete_streaming_response)
)
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=complete_streaming_response,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=complete_streaming_response,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=self.dynamic_success_callbacks,
@ -2026,10 +2028,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
self.model_call_details[
"complete_response"
] = self.model_call_details.get(
"complete_streaming_response", {}
self.model_call_details["complete_response"] = (
self.model_call_details.get(
"complete_streaming_response", {}
)
)
result = self.model_call_details["complete_response"]
openMeterLogger.log_success_event(
@ -2068,10 +2070,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
self.model_call_details[
"complete_response"
] = self.model_call_details.get(
"complete_streaming_response", {}
self.model_call_details["complete_response"] = (
self.model_call_details.get(
"complete_streaming_response", {}
)
)
result = self.model_call_details["complete_response"]
@ -2209,9 +2211,9 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
print_verbose("Async success callbacks: Got a complete streaming response")
self.model_call_details[
"async_complete_streaming_response"
] = complete_streaming_response
self.model_call_details["async_complete_streaming_response"] = (
complete_streaming_response
)
try:
if self.model_call_details.get("cache_hit", False) is True:
@ -2222,10 +2224,10 @@ class Logging(LiteLLMLoggingBaseClass):
model_call_details=self.model_call_details
)
# base_model defaults to None if not set on model_info
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(
result=complete_streaming_response
self.model_call_details["response_cost"] = (
self._response_cost_calculator(
result=complete_streaming_response
)
)
verbose_logger.debug(
@ -2238,16 +2240,16 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = None
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=complete_streaming_response,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj=complete_streaming_response,
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="success",
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
callbacks = self.get_combined_callback_list(
dynamic_success_callbacks=self.dynamic_async_success_callbacks,
@ -2460,18 +2462,18 @@ class Logging(LiteLLMLoggingBaseClass):
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="failure",
error_str=str(exception),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="failure",
error_str=str(exception),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
return start_time, end_time
@ -2979,14 +2981,17 @@ class Logging(LiteLLMLoggingBaseClass):
- For Non-streaming responses, we need to transform the response to a ModelResponse object.
- For streaming responses, anthropic_messages handler calls success_handler with a assembled ModelResponse.
"""
import httpx
if self.stream and isinstance(result, ModelResponse):
return result
elif isinstance(result, ModelResponse):
return result
if "httpx_response" in self.model_call_details:
httpx_response = self.model_call_details.get("httpx_response", None)
if httpx_response and isinstance(httpx_response, httpx.Response):
result = litellm.AnthropicConfig().transform_response(
raw_response=self.model_call_details.get("httpx_response", None),
raw_response=httpx_response,
model_response=litellm.ModelResponse(),
model=self.model,
messages=[],
@ -3355,9 +3360,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
endpoint=arize_config.endpoint,
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = f"space_id={arize_config.space_key},api_key={arize_config.api_key}"
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
f"space_id={arize_config.space_key},api_key={arize_config.api_key}"
)
for callback in _in_memory_loggers:
if (
isinstance(callback, ArizeLogger)
@ -3381,9 +3386,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
# auth can be disabled on local deployments of arize phoenix
if arize_phoenix_config.otlp_auth_headers is not None:
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = arize_phoenix_config.otlp_auth_headers
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
arize_phoenix_config.otlp_auth_headers
)
for callback in _in_memory_loggers:
if (
@ -3515,9 +3520,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
exporter="otlp_http",
endpoint="https://langtrace.ai/api/trace",
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
f"api_key={os.getenv('LANGTRACE_API_KEY')}"
)
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetry)
@ -4197,10 +4202,10 @@ class StandardLoggingPayloadSetup:
for key in StandardLoggingHiddenParams.__annotations__.keys():
if key in hidden_params:
if key == "additional_headers":
clean_hidden_params[
"additional_headers"
] = StandardLoggingPayloadSetup.get_additional_headers(
hidden_params[key]
clean_hidden_params["additional_headers"] = (
StandardLoggingPayloadSetup.get_additional_headers(
hidden_params[key]
)
)
else:
clean_hidden_params[key] = hidden_params[key] # type: ignore
@ -4252,9 +4257,9 @@ class StandardLoggingPayloadSetup:
if (
custom_logger
and hasattr(custom_logger, "s3_path")
and custom_logger.s3_path
and getattr(custom_logger, "s3_path")
):
s3_path = custom_logger.s3_path
s3_path = getattr(custom_logger, "s3_path")
except Exception:
# If any error occurs in getting the logger instance, use default empty s3_path
pass
@ -4704,9 +4709,9 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]):
):
for k, v in metadata["user_api_key_metadata"].items():
if k == "logging": # prevent logging user logging keys
cleaned_user_api_key_metadata[
k
] = "scrubbed_by_litellm_for_sensitive_keys"
cleaned_user_api_key_metadata[k] = (
"scrubbed_by_litellm_for_sensitive_keys"
)
else:
cleaned_user_api_key_metadata[k] = v

View file

@ -2,11 +2,11 @@ import asyncio
import json
import time
import traceback
from litellm._uuid import uuid
from typing import Dict, Iterable, List, Literal, Optional, Tuple, Union
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_extract_reasoning_content,
@ -31,6 +31,7 @@ from litellm.types.utils import Logprobs as TextCompletionLogprobs
from litellm.types.utils import (
Message,
ModelResponse,
ModelResponseStream,
RerankResponse,
StreamingChoices,
TextChoices,
@ -108,12 +109,12 @@ async def convert_to_streaming_response_async(response_object: Optional[dict] =
if response_object is None:
raise Exception("Error in response object format")
model_response_object = ModelResponse(stream=True)
model_response_object = ModelResponseStream()
if model_response_object is None:
raise Exception("Error in response creating model response object")
choice_list = []
choice_list: List[StreamingChoices] = []
for idx, choice in enumerate(response_object["choices"]):
if (
@ -182,8 +183,8 @@ def convert_to_streaming_response(response_object: Optional[dict] = None):
if response_object is None:
raise Exception("Error in response object format")
model_response_object = ModelResponse(stream=True)
choice_list = []
model_response_object = ModelResponseStream()
choice_list: List[StreamingChoices] = []
for idx, choice in enumerate(response_object["choices"]):
delta = Delta(**choice["message"])
finish_reason = choice.get("finish_reason", None)
@ -460,7 +461,7 @@ def convert_to_model_response_object( # noqa: PLR0915
if stream is True:
# for returning cached responses, we need to yield a generator
return convert_to_streaming_response(response_object=response_object)
choice_list = []
choice_list: List[Choices] = []
assert response_object["choices"] is not None and isinstance(
response_object["choices"], Iterable
@ -564,7 +565,7 @@ def convert_to_model_response_object( # noqa: PLR0915
provider_specific_fields=provider_specific_fields,
)
choice_list.append(choice)
model_response_object.choices = choice_list
model_response_object.choices = choice_list # type: ignore
if "usage" in response_object and response_object["usage"] is not None:
usage_object = litellm.Usage(**response_object["usage"])

View file

@ -5,7 +5,6 @@ import json
import threading
import time
import traceback
from litellm._uuid import uuid
from typing import Any, Callable, Dict, List, Optional, Union, cast
import httpx
@ -13,6 +12,7 @@ from pydantic import BaseModel
import litellm
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.model_response_utils import (
is_model_response_stream_empty,
)
@ -1024,7 +1024,7 @@ class CustomStreamWrapper:
return
def chunk_creator(self, chunk: Any): # type: ignore # noqa: PLR0915
if hasattr(chunk, 'id'):
if hasattr(chunk, "id"):
self.response_id = chunk.id
model_response = self.model_response_creator()
response_obj: Dict[str, Any] = {}
@ -1365,12 +1365,13 @@ class CustomStreamWrapper:
f"model_response finish reason 3: {self.received_finish_reason}; response_obj={response_obj}"
)
## FUNCTION CALL PARSING
original_chunk = (
response_obj.get("original_chunk") if response_obj is not None else None
)
if (
response_obj is not None
and response_obj.get("original_chunk", None) is not None
original_chunk is not None
): # function / tool calling branch - only set for openai/azure compatible endpoints
# enter this branch when no content has been passed in response
original_chunk = response_obj.get("original_chunk", None)
if hasattr(original_chunk, "id"):
model_response = self.set_model_id(
original_chunk.id, model_response

View file

@ -55,9 +55,9 @@ class AnthropicTextConfig(BaseConfig):
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
"""
max_tokens_to_sample: Optional[
int
] = litellm.max_tokens # anthropic requires a default
max_tokens_to_sample: Optional[int] = (
litellm.max_tokens
) # anthropic requires a default
stop_sequences: Optional[list] = None
temperature: Optional[int] = None
top_p: Optional[int] = None
@ -291,7 +291,7 @@ class AnthropicTextCompletionResponseIterator(BaseModelResponseIterator):
_chunk_text = chunk.get("completion", None)
if _chunk_text is not None and isinstance(_chunk_text, str):
text = _chunk_text
finish_reason = chunk.get("stop_reason", None)
finish_reason = chunk.get("stop_reason") or ""
if finish_reason is not None:
is_finished = True
returned_chunk = GenericStreamingChunk(

View file

@ -49,7 +49,7 @@ def get_cost_for_anthropic_web_search(
## Get the cost per web search request
search_context_pricing: SearchContextCostPerQuery = (
model_info.get("search_context_cost_per_query", {}) or {}
model_info.get("search_context_cost_per_query") or SearchContextCostPerQuery()
)
cost_per_web_search_request = search_context_pricing.get(
"search_context_size_medium", 0.0

View file

@ -182,12 +182,12 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
model: str,
messages: list,
model_response: ModelResponse,
api_key: str,
api_key: Optional[str],
api_base: str,
api_version: str,
api_type: str,
azure_ad_token: str,
azure_ad_token_provider: Callable,
azure_ad_token: Optional[str],
azure_ad_token_provider: Optional[Callable],
dynamic_params: bool,
print_verbose: Callable,
timeout: Union[float, httpx.Timeout],
@ -372,7 +372,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
async def acompletion(
self,
api_key: str,
api_key: Optional[str],
api_version: str,
model: str,
api_base: str,
@ -477,7 +477,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
self,
logging_obj,
api_base: str,
api_key: str,
api_key: Optional[str],
api_version: str,
dynamic_params: bool,
data: dict,
@ -555,7 +555,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
self,
logging_obj: LiteLLMLoggingObj,
api_base: str,
api_key: str,
api_key: Optional[str],
api_version: str,
dynamic_params: bool,
data: dict,

View file

@ -162,8 +162,8 @@ def get_azure_ad_token_from_username_password(
def get_azure_ad_token_from_oidc(
azure_ad_token: str,
azure_client_id: Optional[str],
azure_tenant_id: Optional[str],
azure_client_id: Optional[str] = None,
azure_tenant_id: Optional[str] = None,
scope: Optional[str] = None,
) -> str:
"""

View file

@ -30,11 +30,11 @@ class AzureTextCompletion(BaseAzureLLM):
model: str,
messages: list,
model_response: ModelResponse,
api_key: str,
api_key: Optional[str],
api_base: str,
api_version: str,
api_type: str,
azure_ad_token: str,
azure_ad_token: Optional[str],
azure_ad_token_provider: Optional[Callable],
print_verbose: Callable,
timeout,
@ -59,7 +59,7 @@ class AzureTextCompletion(BaseAzureLLM):
### CHECK IF CLOUDFLARE AI GATEWAY ###
### if so - set the model as part of the base url
if "gateway.ai.cloudflare.com" in api_base:
if api_base is not None and "gateway.ai.cloudflare.com" in api_base:
## build base url - assume api base includes resource name
client = self._init_azure_client_for_cloudflare_ai_gateway(
api_key=api_key,
@ -196,7 +196,7 @@ class AzureTextCompletion(BaseAzureLLM):
async def acompletion(
self,
api_key: str,
api_key: Optional[str],
api_version: str,
model: str,
api_base: str,
@ -263,7 +263,7 @@ class AzureTextCompletion(BaseAzureLLM):
self,
logging_obj,
api_base: str,
api_key: str,
api_key: Optional[str],
api_version: str,
data: dict,
model: str,
@ -320,7 +320,7 @@ class AzureTextCompletion(BaseAzureLLM):
self,
logging_obj,
api_base: str,
api_key: str,
api_key: Optional[str],
api_version: str,
data: dict,
model: str,

View file

@ -3,14 +3,15 @@ Transformation for Bedrock Invoke Agent
https://docs.aws.amazon.com/bedrock/latest/APIReference/API_agent-runtime_InvokeAgent.html
"""
import base64
import json
from litellm._uuid import uuid
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
)
@ -22,6 +23,11 @@ from litellm.types.llms.bedrock_invoke_agents import (
InvokeAgentEvent,
InvokeAgentEventHeaders,
InvokeAgentEventList,
InvokeAgentMetadata,
InvokeAgentModelInvocationInput,
InvokeAgentModelInvocationOutput,
InvokeAgentOrchestrationTrace,
InvokeAgentPreProcessingTrace,
InvokeAgentTrace,
InvokeAgentTracePayload,
InvokeAgentUsage,
@ -389,15 +395,22 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
self, trace_data: InvokeAgentTrace, usage_info: InvokeAgentUsage
) -> None:
"""Extract usage information from preprocessing trace."""
pre_processing = trace_data.get("preProcessingTrace", {})
pre_processing: Optional[InvokeAgentPreProcessingTrace] = trace_data.get(
"preProcessingTrace"
)
if not pre_processing:
return
model_output = pre_processing.get("modelInvocationOutput", {})
model_output: Optional[InvokeAgentModelInvocationOutput] = (
pre_processing.get("modelInvocationOutput")
or InvokeAgentModelInvocationOutput()
)
if not model_output:
return
metadata = model_output.get("metadata", {})
metadata: Optional[InvokeAgentMetadata] = (
model_output.get("metadata") or InvokeAgentMetadata()
)
if not metadata:
return
@ -412,11 +425,16 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
self, trace_data: InvokeAgentTrace
) -> Optional[str]:
"""Extract model information from orchestration trace."""
orchestration_trace = trace_data.get("orchestrationTrace", {})
orchestration_trace: Optional[InvokeAgentOrchestrationTrace] = trace_data.get(
"orchestrationTrace"
)
if not orchestration_trace:
return None
model_invocation = orchestration_trace.get("modelInvocationInput", {})
model_invocation: Optional[InvokeAgentModelInvocationInput] = (
orchestration_trace.get("modelInvocationInput")
or InvokeAgentModelInvocationInput()
)
if not model_invocation:
return None

View file

@ -7,7 +7,6 @@ import json
import time
import types
import urllib.parse
from litellm._uuid import uuid
from functools import partial
from typing import (
Any,
@ -26,6 +25,7 @@ import httpx # type: ignore
import litellm
from litellm import verbose_logger
from litellm._uuid import uuid
from litellm.caching.caching import InMemoryCache
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.litellm_logging import Logging
@ -498,9 +498,9 @@ class BedrockLLM(BaseAWSLLM):
content=None,
)
model_response.choices[0].message = _message # type: ignore
model_response._hidden_params[
"original_response"
] = outputText # allow user to access raw anthropic tool calling response
model_response._hidden_params["original_response"] = (
outputText # allow user to access raw anthropic tool calling response
)
if (
_is_function_call is True
and stream is not None
@ -808,9 +808,9 @@ class BedrockLLM(BaseAWSLLM):
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
if stream is True:
inference_params[
"stream"
] = True # cohere requires stream = True in inference params
inference_params["stream"] = (
True # cohere requires stream = True in inference params
)
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "anthropic":
if model.startswith("anthropic.claude-3"):
@ -1352,9 +1352,11 @@ class AWSEventStreamDecoder:
"name": None,
"arguments": delta_obj["toolUse"]["input"],
},
"index": self.tool_calls_index
if self.tool_calls_index is not None
else index,
"index": (
self.tool_calls_index
if self.tool_calls_index is not None
else index
),
}
elif "reasoningContent" in delta_obj:
provider_specific_fields = {
@ -1384,9 +1386,11 @@ class AWSEventStreamDecoder:
"name": None,
"arguments": "{}",
},
"index": self.tool_calls_index
if self.tool_calls_index is not None
else index,
"index": (
self.tool_calls_index
if self.tool_calls_index is not None
else index
),
}
elif "stopReason" in chunk_data:
finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop"))
@ -1448,7 +1452,7 @@ class AWSEventStreamDecoder:
######### /bedrock/invoke nova mappings ###############
elif "contentBlockDelta" in chunk_data:
# when using /bedrock/invoke/nova, the chunk_data is nested under "contentBlockDelta"
_chunk_data = chunk_data.get("contentBlockDelta", None)
_chunk_data = chunk_data.get("contentBlockDelta", {})
return self.converse_chunk_parser(chunk_data=_chunk_data)
######## bedrock.mistral mappings ###############
elif "outputs" in chunk_data:

View file

@ -89,6 +89,7 @@ from litellm.utils import (
if TYPE_CHECKING:
from aiohttp import ClientSession
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
@ -281,7 +282,7 @@ class BaseLLMHTTPHandler:
self,
model: str,
messages: list,
api_base: str,
api_base: Optional[str],
custom_llm_provider: str,
model_response: ModelResponse,
encoding,
@ -750,7 +751,7 @@ class BaseLLMHTTPHandler:
model_response: EmbeddingResponse,
api_key: Optional[str] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
aembedding: bool = False,
aembedding: Optional[bool] = False,
headers: Optional[Dict[str, Any]] = None,
) -> EmbeddingResponse:
provider_config = ProviderConfigManager.get_provider_embedding_config(
@ -3100,7 +3101,10 @@ class BaseLLMHTTPHandler:
_is_async: bool = False,
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
"""
Handles image edit requests.
@ -3290,7 +3294,10 @@ class BaseLLMHTTPHandler:
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
"""
Handles image generation requests.
When _is_async=True, returns a coroutine instead of making the call directly.

View file

@ -1,6 +1,7 @@
"""
Transformation for Calling Google models in their native format.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union, cast
import httpx
@ -25,27 +26,29 @@ else:
GenerateContentContentListUnionDict = Any
GenerateContentResponse = Any
ToolConfigDict = Any
from ..common_utils import get_api_key_from_env
class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"""
Configuration for calling Google models in their native format.
"""
##############################
# Constants
##############################
XGOOGLE_API_KEY = "x-goog-api-key"
##############################
@property
def custom_llm_provider(self) -> Literal["gemini", "vertex_ai"]:
return "gemini"
def __init__(self):
super().__init__()
VertexLLM.__init__(self)
def get_supported_generate_content_optional_params(self, model: str) -> List[str]:
"""
Get the list of supported Google GenAI parameters for the model.
@ -58,7 +61,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"""
return [
"http_options",
"system_instruction",
"system_instruction",
"temperature",
"top_p",
"top_k",
@ -84,10 +87,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"speech_config",
"audio_timestamp",
"automatic_function_calling",
"thinking_config"
"thinking_config",
]
def map_generate_content_optional_params(
self,
generate_content_config_dict: GenerateContentConfigDict,
@ -103,26 +105,29 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Returns:
Mapped parameters for the provider
"""
from litellm.types.google_genai.main import GenerateContentConfigDict
_generate_content_config_dict = GenerateContentConfigDict()
supported_google_genai_params = self.get_supported_generate_content_optional_params(model)
_generate_content_config_dict: Dict[str, Any] = {}
supported_google_genai_params = (
self.get_supported_generate_content_optional_params(model)
)
for param, value in generate_content_config_dict.items():
if param in supported_google_genai_params:
_generate_content_config_dict[param] = value
return dict(_generate_content_config_dict)
return _generate_content_config_dict
def validate_environment(
self,
self,
api_key: Optional[str],
headers: Optional[dict],
model: str,
litellm_params: Optional[Union[GenericLiteLLMParams, dict]]
litellm_params: Optional[Union[GenericLiteLLMParams, dict]],
) -> dict:
default_headers = {
"Content-Type": "application/json",
}
# Use the passed api_key first, then fall back to litellm_params and environment
gemini_api_key = api_key or self._get_google_ai_studio_api_key(dict(litellm_params or {}))
gemini_api_key = api_key or self._get_google_ai_studio_api_key(
dict(litellm_params or {})
)
if gemini_api_key is not None:
default_headers[self.XGOOGLE_API_KEY] = gemini_api_key
if headers is not None:
@ -137,14 +142,14 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
or get_api_key_from_env()
or litellm.api_key
)
def _get_common_auth_components(
self,
litellm_params: dict,
) -> Tuple[Any, Optional[str], Optional[str]]:
"""
Get common authentication components used by both sync and async methods.
Returns:
Tuple of (vertex_credentials, vertex_project, vertex_location)
"""
@ -152,7 +157,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
vertex_project = self.get_vertex_ai_project(litellm_params)
vertex_location = self.get_vertex_ai_location(litellm_params)
return vertex_credentials, vertex_project, vertex_location
def _build_final_headers_and_url(
self,
model: str,
@ -168,7 +173,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Build final headers and API URL from auth components.
"""
gemini_api_key = self._get_google_ai_studio_api_key(litellm_params)
auth_header, api_base = self._get_token_and_url(
model=model,
gemini_api_key=gemini_api_key,
@ -201,7 +206,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
"""
Sync version of get_auth_token_and_url.
"""
vertex_credentials, vertex_project, vertex_location = self._get_common_auth_components(litellm_params)
vertex_credentials, vertex_project, vertex_location = (
self._get_common_auth_components(litellm_params)
)
_auth_header, vertex_project = self._ensure_access_token(
credentials=vertex_credentials,
@ -238,7 +245,9 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Returns:
Tuple of headers and API base
"""
vertex_credentials, vertex_project, vertex_location = self._get_common_auth_components(litellm_params)
vertex_credentials, vertex_project, vertex_location = (
self._get_common_auth_components(litellm_params)
)
_auth_header, vertex_project = await self._ensure_access_token_async(
credentials=vertex_credentials,
@ -256,7 +265,6 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
api_base=api_base,
litellm_params=litellm_params,
)
def transform_generate_content_request(
self,
@ -269,6 +277,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
GenerateContentConfigDict,
GenerateContentRequestDict,
)
typed_generate_content_request = GenerateContentRequestDict(
model=model,
contents=contents,
@ -279,7 +288,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
request_dict = cast(dict, typed_generate_content_request)
return request_dict
def transform_generate_content_response(
self,
model: str,
@ -297,6 +306,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
Transformed response data
"""
from litellm.types.google_genai.main import GenerateContentResponse
try:
response = raw_response.json()
except Exception as e:
@ -305,7 +315,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM):
status_code=raw_response.status_code,
headers=raw_response.headers,
)
logging_obj.model_call_details["httpx_response"] = raw_response
return GenerateContentResponse(**response)
return GenerateContentResponse(**response)

View file

@ -40,17 +40,17 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
Reference: https://huggingface.github.io/text-generation-inference/#/Text%20Generation%20Inference/compat_generate
"""
hf_task: Optional[
hf_tasks
] = None # litellm-specific param, used to know the api spec to use when calling huggingface api
hf_task: Optional[hf_tasks] = (
None # litellm-specific param, used to know the api spec to use when calling huggingface api
)
best_of: Optional[int] = None
decoder_input_details: Optional[bool] = None
details: Optional[bool] = True # enables returning logprobs + best of
max_new_tokens: Optional[int] = None
repetition_penalty: Optional[float] = None
return_full_text: Optional[
bool
] = False # by default don't return the input as part of the output
return_full_text: Optional[bool] = (
False # by default don't return the input as part of the output
)
seed: Optional[int] = None
temperature: Optional[float] = None
top_k: Optional[int] = None
@ -120,9 +120,9 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
optional_params["top_p"] = value
if param == "n":
optional_params["best_of"] = value
optional_params[
"do_sample"
] = True # Need to sample if you want best of for hf inference endpoints
optional_params["do_sample"] = (
True # Need to sample if you want best of for hf inference endpoints
)
if param == "stream":
optional_params["stream"] = value
if param == "stop":
@ -268,7 +268,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
# check if the model has a registered custom prompt
model_prompt_details = litellm.custom_prompt_dict[model]
prompt = custom_prompt(
role_dict=model_prompt_details.get("roles", None),
role_dict=model_prompt_details.get("roles") or {},
initial_prompt_value=model_prompt_details.get(
"initial_prompt_value", ""
),
@ -363,9 +363,9 @@ class HuggingFaceEmbeddingConfig(BaseConfig):
"content-type": "application/json",
}
if api_key is not None:
default_headers[
"Authorization"
] = f"Bearer {api_key}" # Huggingface Inference Endpoint default is to accept bearer tokens
default_headers["Authorization"] = (
f"Bearer {api_key}" # Huggingface Inference Endpoint default is to accept bearer tokens
)
headers = {**headers, **default_headers}
return headers

View file

@ -1,5 +1,5 @@
"""
Support for gpt model family
Support for gpt model family
"""
from typing import List, Optional, Union
@ -87,7 +87,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig):
## RESPONSE OBJECT
if response_object is None or model_response_object is None:
raise ValueError("Error in response object format")
choice_list = []
choice_list: List[Choices] = []
for idx, choice in enumerate(response_object["choices"]):
message = Message(
content=choice["text"],
@ -100,7 +100,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig):
logprobs=choice.get("logprobs", None),
)
choice_list.append(choice)
model_response_object.choices = choice_list
model_response_object.choices = choice_list # type: ignore
if "usage" in response_object:
setattr(model_response_object, "usage", response_object["usage"])
@ -111,9 +111,9 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig):
if "model" in response_object:
model_response_object.model = response_object["model"]
model_response_object._hidden_params[
"original_response"
] = response_object # track original response, if users make a litellm.text_completion() request, we can return the original response
model_response_object._hidden_params["original_response"] = (
response_object # track original response, if users make a litellm.text_completion() request, we can return the original response
)
return model_response_object
except Exception as e:
raise e

View file

@ -64,9 +64,9 @@ class VertexFineTuningAPI(VertexLLM):
)
if create_fine_tuning_job_data.validation_file:
supervised_tuning_spec[
"validation_dataset"
] = create_fine_tuning_job_data.validation_file
supervised_tuning_spec["validation_dataset"] = (
create_fine_tuning_job_data.validation_file
)
_vertex_hyperparameters = (
self._transform_openai_hyperparameters_to_vertex_hyperparameters(
@ -140,7 +140,9 @@ class VertexFineTuningAPI(VertexLLM):
fine_tuned_model=response.get("tunedModelDisplayName", ""),
finished_at=None,
hyperparameters=self._translate_vertex_response_hyperparameters(
vertex_hyper_parameters=_supervisedTuningSpec.get("hyperParameters", {})
vertex_hyper_parameters=_supervisedTuningSpec.get(
"hyperParameters", FineTuneHyperparameters()
)
or {}
),
model=response.get("baseModel", "") or "",
@ -343,9 +345,9 @@ class VertexFineTuningAPI(VertexLLM):
elif "cachedContents" in request_route:
_model = request_data.get("model")
if _model is not None and "/publishers/google/models/" not in _model:
request_data[
"model"
] = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{_model}"
request_data["model"] = (
f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{_model}"
)
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1beta1/projects/{vertex_project}/locations/{vertex_location}{request_route}"
else:

View file

@ -43,7 +43,7 @@ class GoogleBatchEmbeddings(VertexLLM):
vertex_project=None,
vertex_location=None,
vertex_credentials=None,
aembedding=False,
aembedding: Optional[bool] = False,
timeout=300,
client=None,
) -> EmbeddingResponse:

View file

@ -1,7 +1,8 @@
"""
Transformation for Calling Google models in their native format.
"""
from typing import Dict, Literal, Optional, Union
from typing import Any, Dict, Literal, Optional, Union
from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig
from litellm.types.router import GenericLiteLLMParams
@ -58,22 +59,21 @@ class VertexAIGoogleGenAIConfig(GoogleGenAIConfig):
Returns:
Mapped parameters for the provider
"""
from litellm.types.google_genai.main import GenerateContentConfigDict
_generate_content_config_dict = GenerateContentConfigDict()
_generate_content_config_dict: Dict = {}
for param, value in generate_content_config_dict.items():
camel_case_key = self._camel_to_snake(param)
_generate_content_config_dict[camel_case_key] = value
return dict(_generate_content_config_dict)
return _generate_content_config_dict
def transform_generate_content_request(
self,
model: str,
contents: any,
tools: Optional[any],
contents: Any,
tools: Optional[Any],
generate_content_config_dict: Dict,
system_instruction: Optional[any] = None,
system_instruction: Optional[Any] = None,
) -> dict:
"""
Transform the generate content request for Vertex AI.

View file

@ -46,7 +46,7 @@ class VertexMultimodalEmbedding(VertexLLM):
vertex_project=None,
vertex_location=None,
vertex_credentials=None,
aembedding=False,
aembedding: Optional[bool] = False,
timeout=300,
client=None,
) -> EmbeddingResponse:

View file

@ -36,7 +36,7 @@ class VertexEmbedding(VertexBase):
timeout: Optional[Union[float, httpx.Timeout]],
api_key: Optional[str] = None,
encoding=None,
aembedding=False,
aembedding: Optional[bool] = False,
api_base: Optional[str] = None,
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
vertex_project: Optional[str] = None,
@ -86,8 +86,10 @@ class VertexEmbedding(VertexBase):
mode="embedding",
)
headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers)
vertex_request: VertexEmbeddingRequest = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input=input, optional_params=optional_params, model=model
vertex_request: VertexEmbeddingRequest = (
litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input=input, optional_params=optional_params, model=model
)
)
_client_params = {}
@ -176,8 +178,10 @@ class VertexEmbedding(VertexBase):
mode="embedding",
)
headers = self.set_headers(auth_header=auth_header, extra_headers=extra_headers)
vertex_request: VertexEmbeddingRequest = litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input=input, optional_params=optional_params, model=model
vertex_request: VertexEmbeddingRequest = (
litellm.vertexAITextEmbeddingConfig.transform_openai_request_to_vertex_embedding_request(
input=input, optional_params=optional_params, model=model
)
)
_async_client_params = {}

View file

@ -21,7 +21,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler):
*,
model: str,
messages: list,
api_base: str,
api_base: Optional[str],
custom_llm_provider: str,
custom_prompt_dict: dict,
model_response: ModelResponse,
@ -70,7 +70,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler):
)
return super().completion(
model=watsonx_auth_payload.get("model_id", None),
model=watsonx_auth_payload.get("model_id") or "",
messages=messages,
api_base=api_base,
custom_llm_provider=custom_llm_provider,

View file

@ -17,12 +17,12 @@ import random
import sys
import time
import traceback
from litellm._uuid import uuid
from concurrent import futures
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from copy import deepcopy
from functools import partial
from typing import (
TYPE_CHECKING,
Any,
Callable,
Coroutine,
@ -36,9 +36,10 @@ from typing import (
Union,
cast,
get_args,
TYPE_CHECKING,
)
from litellm._uuid import uuid
if TYPE_CHECKING:
from aiohttp import ClientSession
@ -721,12 +722,15 @@ async def _sleep_for_timeout_async(timeout: Union[float, str, httpx.Timeout]):
await asyncio.sleep(timeout.connect)
MOCK_RESPONSE_TYPE = Union[str, Exception, dict]
def mock_completion(
model: str,
messages: List,
stream: Optional[bool] = False,
n: Optional[int] = None,
mock_response: Union[str, Exception, dict] = "This is a mock request",
mock_response: Optional[MOCK_RESPONSE_TYPE] = "This is a mock request",
mock_tool_calls: Optional[List] = None,
mock_timeout: Optional[bool] = False,
logging=None,
@ -1007,7 +1011,7 @@ def completion( # type: ignore # noqa: PLR0915
######### unpacking kwargs #####################
args = locals()
api_base = kwargs.get("api_base", None)
mock_response = kwargs.get("mock_response", None)
mock_response: Optional[MOCK_RESPONSE_TYPE] = kwargs.get("mock_response", None)
mock_tool_calls = kwargs.get("mock_tool_calls", None)
mock_timeout = cast(Optional[bool], kwargs.get("mock_timeout", None))
force_timeout = kwargs.get("force_timeout", 600) ## deprecated
@ -1114,7 +1118,7 @@ def completion( # type: ignore # noqa: PLR0915
api_base = base_url
if num_retries is not None:
max_retries = num_retries
logging = litellm_logging_obj
logging: Logging = cast(Logging, litellm_logging_obj)
fallbacks = fallbacks or litellm.model_fallbacks
if fallbacks is not None:
return completion_with_fallbacks(**args)
@ -1427,7 +1431,7 @@ def completion( # type: ignore # noqa: PLR0915
api_version = (
api_version
or litellm.api_version
or get_secret("AZURE_API_VERSION")
or get_secret_str("AZURE_API_VERSION")
or litellm.AZURE_DEFAULT_API_VERSION
)
@ -1435,13 +1439,13 @@ def completion( # type: ignore # noqa: PLR0915
api_key
or litellm.api_key
or litellm.azure_key
or get_secret("AZURE_OPENAI_API_KEY")
or get_secret("AZURE_API_KEY")
or get_secret_str("AZURE_OPENAI_API_KEY")
or get_secret_str("AZURE_API_KEY")
)
azure_ad_token = optional_params.get("extra_body", {}).pop(
"azure_ad_token", None
) or get_secret("AZURE_AD_TOKEN")
) or get_secret_str("AZURE_AD_TOKEN")
azure_ad_token_provider = litellm_params.get(
"azure_ad_token_provider", None
@ -1529,25 +1533,32 @@ def completion( # type: ignore # noqa: PLR0915
)
elif custom_llm_provider == "azure_text":
# azure configs
api_type = get_secret("AZURE_API_TYPE") or "azure"
api_type = get_secret_str("AZURE_API_TYPE") or "azure"
api_base = api_base or litellm.api_base or get_secret("AZURE_API_BASE")
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
if api_base is None:
raise ValueError(
"api_base is required for Azure OpenAI LLM provider. Either set it dynamically or set the AZURE_API_BASE environment variable."
)
api_version = (
api_version or litellm.api_version or get_secret("AZURE_API_VERSION")
api_version
or litellm.api_version
or get_secret_str("AZURE_API_VERSION")
)
api_key = (
api_key
or litellm.api_key
or litellm.azure_key
or get_secret("AZURE_OPENAI_API_KEY")
or get_secret("AZURE_API_KEY")
or get_secret_str("AZURE_OPENAI_API_KEY")
or get_secret_str("AZURE_API_KEY")
)
azure_ad_token = optional_params.get("extra_body", {}).pop(
"azure_ad_token", None
) or get_secret("AZURE_AD_TOKEN")
) or get_secret_str("AZURE_AD_TOKEN")
azure_ad_token_provider = litellm_params.get(
"azure_ad_token_provider", None
@ -1573,7 +1584,7 @@ def completion( # type: ignore # noqa: PLR0915
headers=headers,
api_key=api_key,
api_base=api_base,
api_version=api_version,
api_version=cast(str, api_version),
api_type=api_type,
azure_ad_token=azure_ad_token,
azure_ad_token_provider=azure_ad_token_provider,
@ -2545,15 +2556,10 @@ def completion( # type: ignore # noqa: PLR0915
)
elif custom_llm_provider == "compactifai":
api_key = (
api_key
or get_secret_str("COMPACTIFAI_API_KEY")
or litellm.api_key
api_key or get_secret_str("COMPACTIFAI_API_KEY") or litellm.api_key
)
api_base = (
api_base
or "https://api.compactif.ai/v1"
)
api_base = api_base or "https://api.compactif.ai/v1"
## COMPLETION CALL
response = base_llm_http_handler.completion(
@ -2860,7 +2866,7 @@ def completion( # type: ignore # noqa: PLR0915
logging_obj=logging,
acompletion=acompletion,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
custom_llm_provider=custom_llm_provider, # type: ignore
client=client,
api_base=api_base,
extra_headers=extra_headers,
@ -2929,7 +2935,7 @@ def completion( # type: ignore # noqa: PLR0915
logging_obj=logging,
acompletion=acompletion,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
custom_llm_provider=custom_llm_provider, # type: ignore
client=client,
api_base=api_base,
extra_headers=extra_headers,
@ -3935,7 +3941,7 @@ def embedding( # noqa: PLR0915
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
mock_response: Optional[List[float]] = kwargs.get("mock_response", None) # type: ignore
azure_ad_token_provider = kwargs.get("azure_ad_token_provider", None)
aembedding = kwargs.get("aembedding", None)
aembedding: Optional[bool] = kwargs.get("aembedding", None)
extra_headers = kwargs.get("extra_headers", None)
headers = kwargs.get("headers", None)
### CUSTOM MODEL COST ###
@ -5615,7 +5621,7 @@ def speech( # noqa: PLR0915
if max_retries is None:
max_retries = litellm.num_retries or openai.DEFAULT_MAX_RETRIES
litellm_params_dict = get_litellm_params(**kwargs)
logging_obj = kwargs.get("litellm_logging_obj", None)
logging_obj: Logging = cast(Logging, kwargs.get("litellm_logging_obj"))
logging_obj.update_environment_variables(
model=model,
user=user,

View file

@ -1,6 +1,7 @@
"""
Cost calculator for MCP tools.
"""
from typing import TYPE_CHECKING, Any, Optional, cast
from litellm.types.mcp import MCPServerCostInfo
@ -13,11 +14,12 @@ if TYPE_CHECKING:
else:
LitellmLoggingObject = Any
class MCPCostCalculator:
@staticmethod
def calculate_mcp_tool_call_cost(
litellm_logging_obj: Optional[LitellmLoggingObject],
) -> float:
) -> float:
"""
Calculate the cost of an MCP tool call.
@ -25,29 +27,43 @@ class MCPCostCalculator:
"""
if litellm_logging_obj is None:
return 0.0
#########################################################
# Get the response cost from logging object model_call_details
# This is set when a user modifies the response in a post_mcp_tool_call_hook
#########################################################
response_cost = litellm_logging_obj.model_call_details.get("response_cost", None)
response_cost = litellm_logging_obj.model_call_details.get(
"response_cost", None
)
if response_cost is not None:
return response_cost
#########################################################
# Unpack the mcp_tool_call_metadata
#########################################################
mcp_tool_call_metadata: StandardLoggingMCPToolCall = cast(StandardLoggingMCPToolCall, litellm_logging_obj.model_call_details.get("mcp_tool_call_metadata", {})) or {}
mcp_server_cost_info_raw = mcp_tool_call_metadata.get("mcp_server_cost_info", {}) or {}
mcp_server_cost_info: MCPServerCostInfo = cast(MCPServerCostInfo, mcp_server_cost_info_raw)
mcp_tool_call_metadata: StandardLoggingMCPToolCall = (
cast(
StandardLoggingMCPToolCall,
litellm_logging_obj.model_call_details.get(
"mcp_tool_call_metadata", {}
),
)
or {}
)
mcp_server_cost_info: MCPServerCostInfo = (
mcp_tool_call_metadata.get("mcp_server_cost_info") or MCPServerCostInfo()
)
#########################################################
# User defined cost per query
#########################################################
default_cost_per_query = mcp_server_cost_info.get("default_cost_per_query", None)
tool_name_to_cost_per_query: dict = mcp_server_cost_info.get("tool_name_to_cost_per_query", {}) or {}
default_cost_per_query = mcp_server_cost_info.get(
"default_cost_per_query", None
)
tool_name_to_cost_per_query: dict = (
mcp_server_cost_info.get("tool_name_to_cost_per_query", {}) or {}
)
tool_name = mcp_tool_call_metadata.get("name", "")
#########################################################
# 1. If tool_name is in tool_name_to_cost_per_query, use the cost per query
# 2. If tool_name is not in tool_name_to_cost_per_query, use the default cost per query

View file

@ -1,16 +1,7 @@
import enum
import json
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
List,
Literal,
Optional,
Union,
)
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union
import httpx
from pydantic import (
@ -26,11 +17,7 @@ from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import AllMessageValues, OpenAIFileObject
from litellm.types.mcp import (
MCPAuthType,
MCPTransport,
MCPTransportType,
)
from litellm.types.mcp import MCPAuthType, MCPTransport, MCPTransportType
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.router import RouterErrors, UpdateRouterConfig
from litellm.types.secret_managers.main import KeyManagementSystem
@ -404,16 +391,16 @@ class LiteLLMRoutes(enum.Enum):
]
key_management_routes = [
KeyManagementRoutes.KEY_GENERATE,
KeyManagementRoutes.KEY_UPDATE,
KeyManagementRoutes.KEY_DELETE,
KeyManagementRoutes.KEY_INFO,
KeyManagementRoutes.KEY_REGENERATE,
KeyManagementRoutes.KEY_GENERATE_SERVICE_ACCOUNT,
KeyManagementRoutes.KEY_REGENERATE_WITH_PATH_PARAM,
KeyManagementRoutes.KEY_LIST,
KeyManagementRoutes.KEY_BLOCK,
KeyManagementRoutes.KEY_UNBLOCK,
KeyManagementRoutes.KEY_GENERATE.value,
KeyManagementRoutes.KEY_UPDATE.value,
KeyManagementRoutes.KEY_DELETE.value,
KeyManagementRoutes.KEY_INFO.value,
KeyManagementRoutes.KEY_REGENERATE.value,
KeyManagementRoutes.KEY_GENERATE_SERVICE_ACCOUNT.value,
KeyManagementRoutes.KEY_REGENERATE_WITH_PATH_PARAM.value,
KeyManagementRoutes.KEY_LIST.value,
KeyManagementRoutes.KEY_BLOCK.value,
KeyManagementRoutes.KEY_UNBLOCK.value,
]
management_routes = [
@ -747,9 +734,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
allowed_cache_controls: Optional[list] = []
config: Optional[dict] = {}
permissions: Optional[dict] = {}
model_max_budget: Optional[
dict
] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_max_budget: Optional[dict] = (
{}
) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {}
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
@ -788,12 +775,11 @@ class GenerateKeyRequest(KeyRequestBase):
description="Type of key that determines default allowed routes.",
)
auto_rotate: Optional[bool] = Field(
default=False,
description="Whether this key should be automatically rotated"
default=False, description="Whether this key should be automatically rotated"
)
rotation_interval: Optional[str] = Field(
default=None,
description="How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True"
description="How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True",
)
@ -1157,12 +1143,12 @@ class NewCustomerRequest(BudgetNewRequest):
blocked: bool = False # allow/disallow requests for this end-user
budget_id: Optional[str] = None # give either a budget_id or max_budget
spend: Optional[float] = None
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
@model_validator(mode="before")
@classmethod
@ -1184,12 +1170,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase):
blocked: bool = False # allow/disallow requests for this end-user
max_budget: Optional[float] = None
budget_id: Optional[str] = None # give either a budget_id or max_budget
allowed_model_region: Optional[
AllowedModelRegion
] = None # require all user requests to use models in this specific region
default_model: Optional[
str
] = None # if no equivalent model in allowed region - default all requests to this model
allowed_model_region: Optional[AllowedModelRegion] = (
None # require all user requests to use models in this specific region
)
default_model: Optional[str] = (
None # if no equivalent model in allowed region - default all requests to this model
)
class DeleteCustomerRequest(LiteLLMPydanticObjectBase):
@ -1263,15 +1249,15 @@ class NewTeamRequest(TeamBase):
guardrails: Optional[List[str]] = None
prompts: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
team_member_budget: Optional[
float
] = None # allow user to set a budget for all team members
team_member_rpm_limit: Optional[
int
] = None # allow user to set RPM limit for all team members
team_member_tpm_limit: Optional[
int
] = None # allow user to set TPM limit for all team members
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
team_member_rpm_limit: Optional[int] = (
None # allow user to set RPM limit for all team members
)
team_member_tpm_limit: Optional[int] = (
None # allow user to set TPM limit for all team members
)
team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m"
model_config = ConfigDict(protected_namespaces=())
@ -1350,9 +1336,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase):
class AddTeamCallback(LiteLLMPydanticObjectBase):
callback_name: str
callback_type: Optional[
Literal["success", "failure", "success_and_failure"]
] = "success_and_failure"
callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = (
"success_and_failure"
)
callback_vars: Dict[str, str]
@model_validator(mode="before")
@ -1621,9 +1607,9 @@ class ConfigList(LiteLLMPydanticObjectBase):
stored_in_db: Optional[bool]
field_default_value: Any
premium_field: bool = False
nested_fields: Optional[
List[FieldDetail]
] = None # For nested dictionary or Pydantic fields
nested_fields: Optional[List[FieldDetail]] = (
None # For nested dictionary or Pydantic fields
)
class UserHeaderMapping(LiteLLMPydanticObjectBase):
@ -1931,7 +1917,7 @@ class UserAPIKeyAuth(
key_alias=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
team_alias=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
)
@classmethod
def get_litellm_cli_user_api_key_auth(cls) -> "UserAPIKeyAuth":
"""
@ -1947,7 +1933,7 @@ class UserAPIKeyAuth(
key_alias=LITTELM_CLI_SERVICE_ACCOUNT_NAME,
team_alias=LITTELM_CLI_SERVICE_ACCOUNT_NAME,
)
@classmethod
def get_litellm_internal_jobs_user_api_key_auth(cls) -> "UserAPIKeyAuth":
"""
@ -1990,9 +1976,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase):
budget_id: Optional[str] = None
created_at: datetime
updated_at: datetime
user: Optional[
Any
] = None # You might want to replace 'Any' with a more specific type if available
user: Optional[Any] = (
None # You might want to replace 'Any' with a more specific type if available
)
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
model_config = ConfigDict(protected_namespaces=())
@ -2887,9 +2873,9 @@ class TeamModelDeleteRequest(BaseModel):
# Organization Member Requests
class OrganizationMemberAddRequest(OrgMemberAddRequest):
organization_id: str
max_budget_in_organization: Optional[
float
] = None # Users max budget within the organization
max_budget_in_organization: Optional[float] = (
None # Users max budget within the organization
)
class OrganizationMemberDeleteRequest(MemberDeleteRequest):
@ -3099,9 +3085,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase):
Maps provider names to their budget configs.
"""
providers: Dict[
str, ProviderBudgetResponseObject
] = {} # Dictionary mapping provider names to their budget configurations
providers: Dict[str, ProviderBudgetResponseObject] = (
{}
) # Dictionary mapping provider names to their budget configurations
class ProxyStateVariables(TypedDict):
@ -3235,9 +3221,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
enforce_rbac: bool = False
roles_jwt_field: Optional[str] = None # v2 on role mappings
role_mappings: Optional[List[RoleMapping]] = None
object_id_jwt_field: Optional[
str
] = None # can be either user / team, inferred from the role mapping
object_id_jwt_field: Optional[str] = (
None # can be either user / team, inferred from the role mapping
)
scope_mappings: Optional[List[ScopeMapping]] = None
enforce_scope_based_access: bool = False
enforce_team_based_model_access: bool = False

View file

@ -14,9 +14,9 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
verbose_proxy_logger.debug("Handling oauth2 proxy request")
# Define the OAuth2 config mappings
oauth2_config_mappings: Dict[str, str] = general_settings.get(
"oauth2_config_mappings", {}
) or {}
oauth2_config_mappings: Dict[str, str] = (
general_settings.get("oauth2_config_mappings") or {}
)
verbose_proxy_logger.debug(f"Oauth2 config mappings: {oauth2_config_mappings}")
if not oauth2_config_mappings:

View file

@ -22,9 +22,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.caching import DualCache
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -597,11 +595,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data[
"messages"
] = self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
)
#########################################################
@ -652,11 +650,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
#########################################################
########## 2. Update the messages with the guardrail response ##########
#########################################################
data[
"messages"
] = self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
data["messages"] = (
self._update_messages_with_updated_bedrock_guardrail_response(
messages=new_messages,
bedrock_guardrail_response=bedrock_guardrail_response,
)
)
#########################################################

View file

@ -303,18 +303,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if "{" in key and "}" in key:
start = key.find("{")
end = key.find("}", start)
hash_tag = key[start:end+1]
hash_tag = key[start : end + 1]
else:
# Fallback for keys without hash tags
hash_tag = "no_hash_tag"
if hash_tag not in groups:
groups[hash_tag] = []
groups[hash_tag].append(key)
return groups
async def _execute_redis_batch_rate_limiter_script(
self,
keys_to_fetch: List[str],
@ -332,10 +331,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
if self.batch_rate_limiter_script is None:
return []
key_groups = self._group_keys_by_hash_tag(keys_to_fetch)
all_cache_values = []
for hash_tag, group_keys in key_groups.items():
try:
group_cache_values = await self.batch_rate_limiter_script(
@ -354,7 +353,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
window_size=self.window_size,
)
all_cache_values.extend(group_cache_values)
return all_cache_values
async def should_rate_limit(
@ -378,7 +377,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
for descriptor in descriptors:
descriptor_key = descriptor["key"]
descriptor_value = descriptor["value"]
rate_limit: Optional[RateLimitDescriptorRateLimitObject] = descriptor.get("rate_limit", {}) or {}
rate_limit: RateLimitDescriptorRateLimitObject = (
descriptor.get("rate_limit") or RateLimitDescriptorRateLimitObject()
)
requests_limit = rate_limit.get("requests_per_unit")
tokens_limit = rate_limit.get("tokens_per_unit")
max_parallel_requests_limit = rate_limit.get("max_parallel_requests")
@ -632,26 +633,28 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
for i, status in enumerate(response["statuses"]):
if status["code"] == "OVER_LIMIT":
descriptor = descriptors[floor(i / 2)]
# Calculate reset time (window_start + window_size)
now = datetime.now().timestamp()
reset_time = now + self.window_size # Conservative estimate
reset_time_formatted = datetime.fromtimestamp(reset_time).strftime("%Y-%m-%d %H:%M:%S UTC")
reset_time_formatted = datetime.fromtimestamp(
reset_time
).strftime("%Y-%m-%d %H:%M:%S UTC")
# Handle negative remaining values more gracefully
remaining_display = max(0, status['limit_remaining'])
remaining_display = max(0, status["limit_remaining"])
# Create detailed error message
rate_limit_type = status['rate_limit_type']
current_limit = status['current_limit']
rate_limit_type = status["rate_limit_type"]
current_limit = status["current_limit"]
detail = (
f"Rate limit exceeded for {descriptor['key']}: {descriptor['value']}. "
f"Limit type: {rate_limit_type}. "
f"Current limit: {current_limit}, Remaining: {remaining_display}. "
f"Limit resets at: {reset_time_formatted}"
)
raise HTTPException(
status_code=429,
detail=detail,
@ -693,7 +696,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
return pipeline_operations
async def _execute_token_increment_script(
self,
pipeline_operations: List["RedisPipelineIncrementOperation"],
@ -703,15 +706,17 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
"""
if self.token_increment_script is None:
return
# Group operations by hash tag for Redis cluster compatibility
operation_keys = [op["key"] for op in pipeline_operations]
key_groups = self._group_keys_by_hash_tag(operation_keys)
for _hash_tag, group_keys in key_groups.items():
# Get operations for this hash tag group
group_operations = [op for op in pipeline_operations if op["key"] in group_keys]
group_operations = [
op for op in pipeline_operations if op["key"] in group_keys
]
keys = []
args = []
@ -731,7 +736,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
args=args,
)
async def async_increment_tokens_with_ttl_preservation(
self,
pipeline_operations: List["RedisPipelineIncrementOperation"],
@ -757,7 +761,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
try:
await self._execute_token_increment_script(pipeline_operations)
verbose_proxy_logger.debug(
f"Successfully executed TTL-preserving increment for {len(pipeline_operations)} keys"
)
@ -811,7 +815,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
# Get metadata from kwargs
litellm_metadata = kwargs["litellm_params"].get(get_metadata_variable_name_from_kwargs(kwargs), {})
litellm_metadata = kwargs["litellm_params"].get(
get_metadata_variable_name_from_kwargs(kwargs), {}
)
if litellm_metadata is None:
return
user_api_key = litellm_metadata.get("user_api_key")
@ -825,7 +831,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Get total tokens from response
total_tokens = 0
# spot fix for /responses api
if (isinstance(response_obj, ModelResponse) or isinstance(response_obj, BaseLiteLLMOpenAIResponseObject)):
if isinstance(response_obj, ModelResponse) or isinstance(
response_obj, BaseLiteLLMOpenAIResponseObject
):
_usage = getattr(response_obj, "usage", None)
if _usage and isinstance(_usage, Usage):
if rate_limit_type == "output":
@ -943,7 +951,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
_get_parent_otel_span_from_kwargs(kwargs)
)
litellm_metadata = kwargs["litellm_params"]["metadata"]
user_api_key = litellm_metadata.get("user_api_key") if litellm_metadata else None
user_api_key = (
litellm_metadata.get("user_api_key") if litellm_metadata else None
)
pipeline_operations: List[RedisPipelineIncrementOperation] = []
if user_api_key:

View file

@ -10,7 +10,6 @@ Has all /sso/* routes
import asyncio
import os
from litellm._uuid import uuid
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
@ -19,6 +18,7 @@ from fastapi.responses import RedirectResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.caching import DualCache
from litellm.constants import MAX_SPENDLOG_ROWS_TO_QUERY
from litellm.llms.custom_httpx.http_handler import (
@ -115,7 +115,10 @@ def process_sso_jwt_access_token(
@router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False)
async def google_login(
request: Request, source: Optional[str] = None, key: Optional[str] = None, existing_key: Optional[str] = None
request: Request,
source: Optional[str] = None,
key: Optional[str] = None,
existing_key: Optional[str] = None,
): # noqa: PLR0915
"""
Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env
@ -664,17 +667,20 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
status_code=401,
detail="Result not returned by SSO provider.",
)
if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the key ID from the state
key_id = state.split(":", 1)[1]
# Get existing_key from query parameters if provided
existing_key = request.query_params.get("existing_key")
verbose_proxy_logger.info(f"CLI SSO callback detected for key: {key_id}, existing_key: {existing_key}")
return await cli_sso_callback(request=request, key=key_id, existing_key=existing_key, result=result)
verbose_proxy_logger.info(
f"CLI SSO callback detected for key: {key_id}, existing_key: {existing_key}"
)
return await cli_sso_callback(
request=request, key=key_id, existing_key=existing_key, result=result
)
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=result,
@ -685,30 +691,30 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
)
async def _regenerate_cli_key(existing_key: str, new_key: str, user_id: Optional[str] = None) -> None:
async def _regenerate_cli_key(
existing_key: str, new_key: str, user_id: Optional[str] = None
) -> None:
"""Regenerate an existing CLI key with a new token"""
from litellm.proxy._types import RegenerateKeyRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
regenerate_key_fn,
)
verbose_proxy_logger.info(f"Regenerating existing CLI key: {existing_key}")
admin_user_dict = UserAPIKeyAuth.get_litellm_cli_user_api_key_auth()
regenerate_request = RegenerateKeyRequest(
key=existing_key,
new_key=new_key,
duration="24hr",
user_id=user_id,
)
await regenerate_key_fn(
key=existing_key,
data=regenerate_request,
user_api_key_dict=admin_user_dict
key=existing_key, data=regenerate_request, user_api_key_dict=admin_user_dict
)
verbose_proxy_logger.info(f"Regenerated CLI key: {new_key}")
@ -720,9 +726,9 @@ async def _create_new_cli_key(
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
)
verbose_proxy_logger.info("Creating new CLI key")
await generate_key_helper_fn(
request_type="key",
duration="24hr",
@ -734,13 +740,20 @@ async def _create_new_cli_key(
table_name="key",
token=key,
)
verbose_proxy_logger.info(f"Created new CLI key: {key}")
async def cli_sso_callback(request: Request, key: Optional[str] = None, existing_key: Optional[str] = None, result: Optional[Union[OpenID, dict]] = None):
async def cli_sso_callback(
request: Request,
key: Optional[str] = None,
existing_key: Optional[str] = None,
result: Optional[Union[OpenID, dict]] = None,
):
"""CLI SSO callback - regenerates existing CLI key or creates new one"""
verbose_proxy_logger.info(f"CLI SSO callback for key: {key}, existing_key: {existing_key}")
verbose_proxy_logger.info(
f"CLI SSO callback for key: {key}, existing_key: {existing_key}"
)
from litellm.proxy.proxy_server import prisma_client
@ -754,8 +767,10 @@ async def cli_sso_callback(request: Request, key: Optional[str] = None, existing
raise HTTPException(
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
)
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(result=result)
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(
result=result
)
verbose_proxy_logger.debug(f"parsed_openid_result: {parsed_openid_result}")
try:
@ -783,7 +798,9 @@ async def cli_sso_callback(request: Request, key: Optional[str] = None, existing
except Exception as e:
verbose_proxy_logger.error(f"Error with CLI key: {e}")
raise HTTPException(status_code=500, detail=f"Failed to process CLI key: {str(e)}")
raise HTTPException(
status_code=500, detail=f"Failed to process CLI key: {str(e)}"
)
@router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False)
@ -874,7 +891,7 @@ async def insert_sso_user(
auto_create_key=False,
)
if result_openid:
if result_openid and isinstance(result_openid, OpenID):
new_user_request.metadata = {"auth_provider": result_openid.provider}
response = await new_user(
@ -1052,11 +1069,13 @@ class SSOAuthenticationHandler:
# or a cryptographicly signed state that we can verify stateless
# For simplification we are using a static state, this is not perfect but some
# SSO providers do not allow stateless verification
redirect_params = SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=state,
generic_authorization_endpoint=generic_authorization_endpoint
redirect_params = (
SSOAuthenticationHandler._get_generic_sso_redirect_params(
state=state,
generic_authorization_endpoint=generic_authorization_endpoint,
)
)
return await generic_sso.get_login_redirect(**redirect_params) # type: ignore
raise ValueError(
"Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso"
@ -1064,26 +1083,26 @@ class SSOAuthenticationHandler:
@staticmethod
def _get_generic_sso_redirect_params(
state: Optional[str] = None,
generic_authorization_endpoint: Optional[str] = None
state: Optional[str] = None,
generic_authorization_endpoint: Optional[str] = None,
) -> dict:
"""
Get redirect parameters for Generic SSO with proper state priority handling.
Priority order:
1. CLI state (if provided)
2. GENERIC_CLIENT_STATE environment variable
3. Generated UUID for Okta (if Okta endpoint detected)
Args:
state: Optional state parameter (e.g., CLI state)
generic_authorization_endpoint: Authorization endpoint URL
Returns:
dict: Redirect parameters for SSO login
"""
redirect_params = {}
if state:
# CLI state takes priority
# the litellm proxy cli sends the "state" parameter to the proxy server for auth. We should maintain the state parameter for the cli if it is provided
@ -1092,8 +1111,13 @@ class SSOAuthenticationHandler:
generic_client_state = os.getenv("GENERIC_CLIENT_STATE", None)
if generic_client_state:
redirect_params["state"] = generic_client_state
elif generic_authorization_endpoint and "okta" in generic_authorization_endpoint:
redirect_params["state"] = uuid.uuid4().hex # set state param for okta - required
elif (
generic_authorization_endpoint
and "okta" in generic_authorization_endpoint
):
redirect_params["state"] = (
uuid.uuid4().hex
) # set state param for okta - required
return redirect_params
@ -1127,11 +1151,11 @@ class SSOAuthenticationHandler:
redirect_url += sso_callback_route
else:
redirect_url += "/" + sso_callback_route
# Append existing_key as query parameter if provided
if existing_key:
redirect_url += f"?existing_key={existing_key}"
return redirect_url
@staticmethod
@ -1314,7 +1338,9 @@ class SSOAuthenticationHandler:
return team_request
@staticmethod
def _get_cli_state(source: Optional[str], key: Optional[str], existing_key: Optional[str] = None) -> Optional[str]:
def _get_cli_state(
source: Optional[str], key: Optional[str], existing_key: Optional[str] = None
) -> Optional[str]:
"""
Checks the request 'source' if a cli state token was passed in
@ -1374,7 +1400,7 @@ class SSOAuthenticationHandler:
if user_email is not None and (user_id is None or len(user_id) == 0):
user_id = user_email
return ParsedOpenIDResult(
user_email=user_email,
user_id=user_id,
@ -1408,13 +1434,16 @@ class SSOAuthenticationHandler:
)
# User is Authe'd in - generate key for the UI to access Proxy
parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result(result=result, generic_client_id=generic_client_id)
parsed_openid_result = (
SSOAuthenticationHandler._get_user_email_and_id_from_result(
result=result, generic_client_id=generic_client_id
)
)
user_email = parsed_openid_result.get("user_email")
user_id = parsed_openid_result.get("user_id")
user_role = parsed_openid_result.get("user_role")
verbose_proxy_logger.info(f"SSO callback result: {result}")
user_info = None
user_id_models: List = []
max_internal_user_budget = litellm.max_internal_user_budget

View file

@ -171,7 +171,7 @@ class TeamMemberPermissionChecks:
"""
all_available_permissions = []
for route in LiteLLMRoutes.key_management_routes.value:
all_available_permissions.append(route.value)
all_available_permissions.append(route)
return all_available_permissions
@staticmethod

View file

@ -1,6 +1,5 @@
# What is this?
## Helper utils for the management endpoints (keys/users/teams)
from litellm._uuid import uuid
from datetime import datetime
from functools import wraps
from typing import Optional, Tuple
@ -9,6 +8,7 @@ from fastapi import HTTPException, Request
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.proxy._types import ( # key request types; user request types; team request types; customer request types
DeleteCustomerRequest,
DeleteTeamRequest,
@ -36,7 +36,7 @@ def get_new_internal_user_defaults(
user_info = litellm.default_internal_user_params or {}
returned_dict: SSOUserDefinedValues = {
"models": user_info.get("models", None),
"models": user_info.get("models") or [],
"max_budget": user_info.get("max_budget", litellm.max_internal_user_budget),
"budget_duration": user_info.get(
"budget_duration", litellm.internal_user_budget_duration

View file

@ -3,7 +3,6 @@ import asyncio
import copy
import json
import traceback
from litellm._uuid import uuid
from base64 import b64encode
from datetime import datetime
from typing import Dict, List, Optional, Tuple, Union
@ -25,6 +24,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -424,10 +424,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
for field_name, field_value in form_data.items():
if isinstance(field_value, (StarletteUploadFile, UploadFile)):
files[
field_name
] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
upload_file=field_value
files[field_name] = (
await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
upload_file=field_value
)
)
else:
form_data_dict[field_name] = field_value
@ -476,7 +476,11 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
user_api_key_request_route=user_api_key_dict.request_route,
user_api_key_spend=user_api_key_dict.spend,
user_api_key_max_budget=user_api_key_dict.max_budget,
user_api_key_budget_reset_at=user_api_key_dict.budget_reset_at.isoformat() if user_api_key_dict.budget_reset_at else None,
user_api_key_budget_reset_at=(
user_api_key_dict.budget_reset_at.isoformat()
if user_api_key_dict.budget_reset_at
else None
),
)
)
@ -496,7 +500,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
kwargs = {
"litellm_params": {
**litellm_params_in_body,
**litellm_params_in_body, # type: ignore
"metadata": _metadata,
"proxy_server_request": {
"url": str(request.url),
@ -509,9 +513,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"passthrough_logging_payload": passthrough_logging_payload,
}
logging_obj.model_call_details[
"passthrough_logging_payload"
] = passthrough_logging_payload
logging_obj.model_call_details["passthrough_logging_payload"] = (
passthrough_logging_payload
)
return kwargs
@ -923,7 +927,6 @@ def create_pass_through_route(
):
# check if target is an adapter.py or a url
from litellm._uuid import uuid
from litellm.proxy.types_utils.utils import get_instance_fn
try:
@ -1367,7 +1370,6 @@ async def create_pass_through_endpoints(
Create new pass-through endpoint
"""
from litellm._uuid import uuid
from litellm.proxy.proxy_server import (
get_config_general_settings,
update_config_general_settings,

View file

@ -10,7 +10,7 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import REDACTED_BY_LITELM_STRING, MAX_STRING_LENGTH_PROMPT_IN_DB
from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB, REDACTED_BY_LITELM_STRING
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
@ -21,6 +21,7 @@ from litellm.types.utils import (
StandardLoggingModelInformation,
StandardLoggingPayload,
StandardLoggingVectorStoreRequest,
VectorStoreSearchResponse,
)
from litellm.utils import get_end_user_id_for_cost_tracking
@ -297,7 +298,9 @@ def get_logging_payload( # noqa: PLR0915
id = f"{id}_cache_hit{time.time()}" # SpendLogs does not allow duplicate request_id
mcp_namespaced_tool_name = None
mcp_tool_call_metadata = clean_metadata.get("mcp_tool_call_metadata", {})
mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall] = clean_metadata.get(
"mcp_tool_call_metadata"
)
if mcp_tool_call_metadata is not None:
mcp_namespaced_tool_name = mcp_tool_call_metadata.get(
"namespaced_tool_name", None
@ -505,23 +508,23 @@ def _sanitize_request_body_for_spend_logs_payload(
# This split ensures we keep more context from the end of conversations
start_ratio = 0.35
end_ratio = 0.65
# Calculate character distribution
start_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * start_ratio)
end_chars = int(MAX_STRING_LENGTH_PROMPT_IN_DB * end_ratio)
# Ensure we don't exceed the total limit
total_keep = start_chars + end_chars
if total_keep > MAX_STRING_LENGTH_PROMPT_IN_DB:
end_chars = MAX_STRING_LENGTH_PROMPT_IN_DB - start_chars
# If the string length is less than what we want to keep, just truncate normally
if len(value) <= MAX_STRING_LENGTH_PROMPT_IN_DB:
return value
# Calculate how many characters are being skipped
skipped_chars = len(value) - total_keep
# Build the truncated string: beginning + truncation marker + end
truncated_value = (
f"{value[:start_chars]}"
@ -567,8 +570,9 @@ def _get_vector_store_request_for_spend_logs_payload(
if vector_store_request_metadata is None:
return None
for vector_store_request in vector_store_request_metadata:
vector_store_search_response = (
vector_store_request.get("vector_store_search_response", {}) or {}
vector_store_search_response: VectorStoreSearchResponse = (
vector_store_request.get("vector_store_search_response")
or VectorStoreSearchResponse()
)
response_data = vector_store_search_response.get("data", []) or []
for response_item in response_data:

View file

@ -17,7 +17,6 @@ import logging
import threading
import time
import traceback
from litellm._uuid import uuid
from collections import defaultdict
from functools import lru_cache
from typing import (
@ -46,6 +45,7 @@ import litellm.litellm_core_utils
import litellm.litellm_core_utils.exception_mapping_utils
from litellm import get_secret_str
from litellm._logging import verbose_router_logger
from litellm._uuid import uuid
from litellm.caching.caching import (
DualCache,
InMemoryCache,
@ -2011,11 +2011,17 @@ class Router:
# Filter out prompt management specific parameters from data before merging
prompt_management_params = {
"bitbucket_config", "dotprompt_config", "prompt_id",
"prompt_variables", "prompt_label", "prompt_version"
"bitbucket_config",
"dotprompt_config",
"prompt_id",
"prompt_variables",
"prompt_label",
"prompt_version",
}
filtered_data = {k: v for k, v in data.items() if k not in prompt_management_params}
filtered_data = {
k: v for k, v in data.items() if k not in prompt_management_params
}
kwargs = {**filtered_data, **kwargs, **optional_params}
kwargs["model"] = model
kwargs["messages"] = messages
@ -3442,7 +3448,7 @@ class Router:
*[try_retrieve_batch(model) for model in filtered_model_list]
)
final_results = {
final_results: Dict = {
"object": "list",
"data": [],
"first_id": None,
@ -4114,7 +4120,9 @@ class Router:
"""
model_group = kwargs.get("model")
response = original_function(*args, **kwargs)
if coroutine_checker.is_async_callable(response) or inspect.isawaitable(response):
if coroutine_checker.is_async_callable(response) or inspect.isawaitable(
response
):
response = await response
## PROCESS RESPONSE HEADERS
response = await self.set_response_headers(
@ -4523,7 +4531,9 @@ class Router:
_time_to_cooldown = self.cooldown_time
if isinstance(_model_info, dict):
deployment_id = _model_info.get("id", None)
deployment_id: Optional[str] = _model_info.get("id")
if deployment_id is None:
return False
increment_deployment_failures_for_current_minute(
litellm_router_instance=self,
deployment_id=deployment_id,
@ -5141,12 +5151,12 @@ class Router:
# Check if this is a prompt management model before validating as LLM provider
litellm_model = deployment.litellm_params.model
is_prompt_management_model = False
if "/" in litellm_model:
split_litellm_model = litellm_model.split("/")[0]
if split_litellm_model in litellm._known_custom_logger_compatible_callbacks:
is_prompt_management_model = True
if is_prompt_management_model:
# For prompt management models, skip LLM provider validation
# The actual model will be resolved at runtime from the prompt file
@ -5236,11 +5246,12 @@ class Router:
# litellm_router_instance=self, model=deployment.to_json(exclude_none=True)
# )
self._initialize_deployment_for_pass_through(
deployment=deployment,
custom_llm_provider=custom_llm_provider,
model=deployment.litellm_params.model,
)
if custom_llm_provider is not None:
self._initialize_deployment_for_pass_through(
deployment=deployment,
custom_llm_provider=custom_llm_provider,
model=deployment.litellm_params.model,
)
#########################################################
# Check if this is an auto-router deployment

View file

@ -1308,7 +1308,7 @@ class MCPListToolsFailedEvent(BaseLiteLLMOpenAIResponseObject):
item_id: str
# MCP Call Events
# MCP Call Events
class MCPCallInProgressEvent(BaseLiteLLMOpenAIResponseObject):
type: Literal[ResponsesAPIStreamEvents.MCP_CALL_IN_PROGRESS]
sequence_number: int

View file

@ -915,7 +915,6 @@ async def test_async_embedding_bedrock():
pytest.fail(f"An exception occurred: {str(e)}")
# Image Generation

View file

@ -1,38 +1,2 @@
============================= test session starts ==============================
platform darwin -- Python 3.13.1, pytest-8.3.5, pluggy-1.5.0 -- /Users/krrishdholakia/Documents/litellm/myenv/bin/python3.13
cachedir: .pytest_cache
rootdir: /Users/krrishdholakia/Documents/litellm
configfile: pyproject.toml
plugins: respx-0.22.0, postgresql-7.0.1, anyio-4.4.0, asyncio-0.26.0, mock-3.14.0, ddtrace-2.19.0rc1, xdist-3.6.1
asyncio: mode=Mode.STRICT, asyncio_default_fixture_loop_scope=None, asyncio_default_test_loop_scope=function
collecting ... collected 8 items
test_main.py::test_url_with_format_param[False-anthropic/claude-3-5-sonnet] PASSED [ 12%]
test_main.py::test_url_with_format_param[False-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 25%]
test_main.py::test_url_with_format_param[False-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 37%]
test_main.py::test_url_with_format_param[False-gemini/gemini-1.5-flash] PASSED [ 50%]
test_main.py::test_url_with_format_param[True-anthropic/claude-3-5-sonnet] PASSED [ 62%]
test_main.py::test_url_with_format_param[True-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 75%]
test_main.py::test_url_with_format_param[True-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 87%]
test_main.py::test_url_with_format_param[True-gemini/gemini-1.5-flash] PASSED [100%]
=============================== warnings summary ===============================
tests/litellm/test_main.py::test_url_with_format_param[False-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0]
tests/litellm/test_main.py::test_url_with_format_param[False-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0]
tests/litellm/test_main.py::test_url_with_format_param[True-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0]
tests/litellm/test_main.py::test_url_with_format_param[True-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0]
/Users/krrishdholakia/Documents/litellm/myenv/lib/python3.13/site-packages/botocore/auth.py:425: DeprecationWarning: datetime.datetime.utcnow() is deprecated and scheduled for removal in a future version. Use timezone-aware objects to represent datetimes in UTC: datetime.datetime.now(datetime.UTC).
datetime_now = datetime.datetime.utcnow()
tests/litellm/test_main.py::test_url_with_format_param[True-anthropic/claude-3-5-sonnet]
/Users/krrishdholakia/Documents/litellm/myenv/lib/python3.13/site-packages/pydantic/main.py:421: UserWarning: Pydantic serializer warnings:
Expected `str` but got `MagicMock` with value `<MagicMock name='mock().j...em__()' id='5209845984'>` - serialized value may not be as expected
return self.__pydantic_serializer__.to_python(
tests/litellm/test_main.py::test_url_with_format_param[True-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0]
/Users/krrishdholakia/Documents/litellm/myenv/lib/python3.13/site-packages/pydantic/main.py:421: UserWarning: Pydantic serializer warnings:
Expected `str` but got `MagicMock` with value `<MagicMock name='mock().j...em__()' id='5210168288'>` - serialized value may not be as expected
return self.__pydantic_serializer__.to_python(
-- Docs: https://docs.pytest.org/en/stable/how-to/capture-warnings.html
======================== 8 passed, 6 warnings in 2.33s =========================
llms/bedrock/chat/invoke_agent/transformation.py:404: error: Incompatible types in assignment (expression has type "object", variable has type "InvokeAgentModelInvocationOutput | None") [assignment]
llms/bedrock/chat/invoke_agent/transformation.py:405: error: Argument 1 to "get" of "Mapping" has incompatible type "str | InvokeAgentModelInvocationOutput"; expected "str" [typeddict-item]