_transform_request_body

This commit is contained in:
Ishaan Jaff
2025-03-26 00:05:45 -07:00
parent 830ecbdb8c
commit 2bef0481af
3 changed files with 88 additions and 55 deletions
+73 -50
View File
@@ -3,66 +3,89 @@ from typing import Dict, Optional
import litellm
def _ensure_extra_body_is_safe(extra_body: Optional[Dict]) -> Optional[Dict]:
"""
Ensure that the extra_body sent in the request is safe, otherwise users will see this error
class LitellmCoreRequestUtils:
"Object of type TextPromptClient is not JSON serializable
@staticmethod
def _ensure_extra_body_is_safe(extra_body: Optional[Dict]) -> Optional[Dict]:
"""
Ensure that the extra_body sent in the request is safe, otherwise users will see this error
"Object of type TextPromptClient is not JSON serializable
Relevant Issue: https://github.com/BerriAI/litellm/issues/4140
"""
if extra_body is None:
return None
Relevant Issue: https://github.com/BerriAI/litellm/issues/4140
"""
if extra_body is None:
return None
if not isinstance(extra_body, dict):
return extra_body
if "metadata" in extra_body and isinstance(extra_body["metadata"], dict):
if "prompt" in extra_body["metadata"]:
_prompt = extra_body["metadata"].get("prompt")
# users can send Langfuse TextPromptClient objects, so we need to convert them to dicts
# Langfuse TextPromptClients have .__dict__ attribute
if _prompt is not None and hasattr(_prompt, "__dict__"):
extra_body["metadata"]["prompt"] = _prompt.__dict__
if not isinstance(extra_body, dict):
return extra_body
if "metadata" in extra_body and isinstance(extra_body["metadata"], dict):
if "prompt" in extra_body["metadata"]:
_prompt = extra_body["metadata"].get("prompt")
@staticmethod
def pick_cheapest_chat_models_from_llm_provider(custom_llm_provider: str, n=1):
"""
Pick the n cheapest chat models from the LLM provider.
# users can send Langfuse TextPromptClient objects, so we need to convert them to dicts
# Langfuse TextPromptClients have .__dict__ attribute
if _prompt is not None and hasattr(_prompt, "__dict__"):
extra_body["metadata"]["prompt"] = _prompt.__dict__
Args:
custom_llm_provider (str): The name of the LLM provider.
n (int): The number of cheapest models to return.
return extra_body
Returns:
list[str]: A list of the n cheapest chat models.
"""
if custom_llm_provider not in litellm.models_by_provider:
return []
known_models = litellm.models_by_provider.get(custom_llm_provider, [])
model_costs = []
def pick_cheapest_chat_models_from_llm_provider(custom_llm_provider: str, n=1):
"""
Pick the n cheapest chat models from the LLM provider.
Args:
custom_llm_provider (str): The name of the LLM provider.
n (int): The number of cheapest models to return.
Returns:
list[str]: A list of the n cheapest chat models.
"""
if custom_llm_provider not in litellm.models_by_provider:
return []
known_models = litellm.models_by_provider.get(custom_llm_provider, [])
model_costs = []
for model in known_models:
try:
model_info = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
for model in known_models:
try:
model_info = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
except Exception:
continue
if model_info.get("mode") != "chat":
continue
_cost = model_info.get("input_cost_per_token", 0) + model_info.get(
"output_cost_per_token", 0
)
except Exception:
continue
if model_info.get("mode") != "chat":
continue
_cost = model_info.get("input_cost_per_token", 0) + model_info.get(
"output_cost_per_token", 0
)
model_costs.append((model, _cost))
model_costs.append((model, _cost))
# Sort by cost (ascending)
model_costs.sort(key=lambda x: x[1])
# Sort by cost (ascending)
model_costs.sort(key=lambda x: x[1])
# Return the top n cheapest models
return [model for model, _ in model_costs[:n]]
# Return the top n cheapest models
return [model for model, _ in model_costs[:n]]
@staticmethod
def select_model_for_request_transformation(
model: str,
base_model: Optional[str] = None,
litellm_params: Optional[Dict] = None,
) -> str:
"""
If `base_model` is passed in by user, use it for the request transformation
Else, use the model passed in the request
"""
if base_model is not None:
return base_model
elif (
litellm_params is not None and litellm_params.get("base_model") is not None
):
return litellm_params["base_model"]
else:
return model
@@ -12,6 +12,7 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.llm_request_utils import LitellmCoreRequestUtils
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_image_obj,
convert_to_gemini_tool_call_invoke,
@@ -284,6 +285,11 @@ def _transform_request_body(
Common transformation logic across sync + async Gemini /generateContent calls.
"""
# Separate system prompt from rest of message
model = LitellmCoreRequestUtils.select_model_for_request_transformation(
model=model,
litellm_params=litellm_params,
)
supports_system_message = get_supports_system_message(
model=model, custom_llm_provider=custom_llm_provider
)
+9 -5
View File
@@ -84,7 +84,7 @@ from litellm.litellm_core_utils.get_llm_provider_logic import (
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
)
from litellm.litellm_core_utils.llm_request_utils import _ensure_extra_body_is_safe
from litellm.litellm_core_utils.llm_request_utils import LitellmCoreRequestUtils
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
LiteLLMResponseObjectHandler,
_handle_invalid_parallel_tool_calls,
@@ -2855,8 +2855,10 @@ def get_optional_params( # noqa: PLR0915
special_params = passed_params.pop("kwargs")
# Use `base_model` for paramter mapping if passed in by user
if base_model is not None:
model = base_model
model = LitellmCoreRequestUtils.select_model_for_request_transformation(
model=model,
base_model=base_model,
)
for k, v in special_params.items():
if k.startswith("aws_") and (
custom_llm_provider != "bedrock" and custom_llm_provider != "sagemaker"
@@ -3738,8 +3740,10 @@ def get_optional_params( # noqa: PLR0915
**extra_body,
}
optional_params["extra_body"] = _ensure_extra_body_is_safe(
extra_body=optional_params["extra_body"]
optional_params["extra_body"] = (
LitellmCoreRequestUtils._ensure_extra_body_is_safe(
extra_body=optional_params["extra_body"]
)
)
else:
# if user passed in non-default kwargs for specific providers/models, pass them along