mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
a7470b3291
38 changed files with 646 additions and 563 deletions
2
.github/workflows/test-linting.yml
vendored
2
.github/workflows/test-linting.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -915,7 +915,6 @@ async def test_async_embedding_bedrock():
|
|||
pytest.fail(f"An exception occurred: {str(e)}")
|
||||
|
||||
|
||||
|
||||
# Image Generation
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue