Merge branch 'main' into akshoop/fastuuid-dep-make-optional

This commit is contained in:
Alex Shoop
2025-09-24 02:15:02 +09:00
committed by GitHub
6 changed files with 104 additions and 80 deletions
+5 -2
View File
@@ -330,7 +330,9 @@ def cost_per_token( # noqa: PLR0915
elif custom_llm_provider == "bedrock":
return bedrock_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "openai":
return openai_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
return openai_cost_per_token(
model=model, usage=usage_block, service_tier=service_tier
)
elif custom_llm_provider == "databricks":
return databricks_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "fireworks_ai":
@@ -351,6 +353,7 @@ def cost_per_token( # noqa: PLR0915
from litellm.llms.dashscope.cost_calculator import (
cost_per_token as dashscope_cost_per_token,
)
return dashscope_cost_per_token(model=model, usage=usage_block)
else:
model_info = _cached_get_model_info_helper(
@@ -663,7 +666,7 @@ def completion_cost( # noqa: PLR0915
completion_response=completion_response
)
rerank_billed_units: Optional[RerankBilledUnits] = None
# Extract service_tier from optional_params if not provided directly
if service_tier is None and optional_params is not None:
service_tier = optional_params.get("service_tier")
+13 -5
View File
@@ -1228,7 +1228,9 @@ 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(
@@ -4191,16 +4193,22 @@ class StandardLoggingPayloadSetup:
# Get the actual s3_path from the configured cold storage logger instance
s3_path = "" # default value
# Try to get the actual logger instance from the logger name
try:
custom_logger = litellm.logging_callback_manager.get_active_custom_logger_for_callback_name(configured_cold_storage_logger)
if custom_logger and hasattr(custom_logger, 's3_path') and custom_logger.s3_path:
custom_logger = litellm.logging_callback_manager.get_active_custom_logger_for_callback_name(
configured_cold_storage_logger
)
if (
custom_logger
and hasattr(custom_logger, "s3_path")
and custom_logger.s3_path
):
s3_path = custom_logger.s3_path
except Exception:
# If any error occurs in getting the logger instance, use default empty s3_path
pass
s3_object_key = get_s3_object_key(
s3_path=s3_path, # Use actual s3_path from logger configuration
team_alias_prefix="", # Don't split by team alias for cold storage
@@ -11,8 +11,8 @@ from litellm.types.utils import (
ImageResponse,
ModelInfo,
PassthroughCallTypes,
Usage,
ServiceTier,
Usage,
)
from litellm.utils import get_model_info
@@ -118,21 +118,21 @@ def _generic_cost_per_character(
def _get_service_tier_cost_key(base_key: str, service_tier: Optional[str]) -> str:
"""
Get the appropriate cost key based on service tier.
Args:
base_key: The base cost key (e.g., "input_cost_per_token")
service_tier: The service tier ("flex", "priority", or None for standard)
Returns:
str: The cost key to use (e.g., "input_cost_per_token_flex" or "input_cost_per_token")
"""
if service_tier is None:
return base_key
# Only use service tier specific keys for "flex" and "priority"
if service_tier.lower() in [ServiceTier.FLEX.value, ServiceTier.PRIORITY.value]:
return f"{base_key}_{service_tier.lower()}"
# For any other service tier, use standard pricing
return base_key
@@ -152,15 +152,15 @@ def _get_token_base_cost(
# Get service tier aware cost keys
input_cost_key = _get_service_tier_cost_key("input_cost_per_token", service_tier)
output_cost_key = _get_service_tier_cost_key("output_cost_per_token", service_tier)
cache_creation_cost_key = _get_service_tier_cost_key("cache_creation_input_token_cost", service_tier)
cache_read_cost_key = _get_service_tier_cost_key("cache_read_input_token_cost", service_tier)
prompt_base_cost = cast(
float, _get_cost_per_unit(model_info, input_cost_key)
cache_creation_cost_key = _get_service_tier_cost_key(
"cache_creation_input_token_cost", service_tier
)
completion_base_cost = cast(
float, _get_cost_per_unit(model_info, output_cost_key)
cache_read_cost_key = _get_service_tier_cost_key(
"cache_read_input_token_cost", service_tier
)
prompt_base_cost = cast(float, _get_cost_per_unit(model_info, input_cost_key))
completion_base_cost = cast(float, _get_cost_per_unit(model_info, output_cost_key))
cache_creation_cost = cast(
float, _get_cost_per_unit(model_info, cache_creation_cost_key)
)
@@ -168,9 +168,7 @@ def _get_token_base_cost(
float,
_get_cost_per_unit(model_info, "cache_creation_input_token_cost_above_1hr"),
)
cache_read_cost = cast(
float, _get_cost_per_unit(model_info, cache_read_cost_key)
)
cache_read_cost = cast(float, _get_cost_per_unit(model_info, cache_read_cost_key))
## CHECK IF ABOVE THRESHOLD
threshold: Optional[float] = None
@@ -183,7 +181,6 @@ def _get_token_base_cost(
1000 if "k" in threshold_str else 1
)
if usage.prompt_tokens > threshold:
prompt_base_cost = cast(
float, _get_cost_per_unit(model_info, key, prompt_base_cost)
)
@@ -278,7 +275,7 @@ def _get_cost_per_unit(
verbose_logger.exception(
f"litellm.litellm_core_utils.llm_cost_calc.utils.py::calculate_cost_per_component(): Exception occured - {cost_per_unit}\nDefaulting to 0.0"
)
# If the service tier key doesn't exist or is None, try to fall back to the standard key
if cost_per_unit is None:
# Check if any service tier suffix exists in the cost key using ServiceTier enum
@@ -286,7 +283,7 @@ def _get_cost_per_unit(
suffix = f"_{service_tier.value}"
if suffix in cost_key:
# Extract the base key by removing the matched suffix
base_key = cost_key.replace(suffix, '')
base_key = cost_key.replace(suffix, "")
fallback_cost = model_info.get(base_key)
if isinstance(fallback_cost, float):
return fallback_cost
@@ -300,7 +297,7 @@ def _get_cost_per_unit(
f"litellm.litellm_core_utils.llm_cost_calc.utils.py::_get_cost_per_unit(): Exception occured - {fallback_cost}\nDefaulting to 0.0"
)
break # Only try the first matching suffix
return default_value
@@ -495,7 +492,10 @@ def _calculate_input_cost(
def generic_cost_per_token(
model: str, usage: Usage, custom_llm_provider: str, service_tier: Optional[str] = None
model: str,
usage: Usage,
custom_llm_provider: str,
service_tier: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@@ -547,7 +547,9 @@ def generic_cost_per_token(
cache_creation_cost,
cache_creation_cost_above_1hr,
cache_read_cost,
) = _get_token_base_cost(model_info=model_info, usage=usage, service_tier=service_tier)
) = _get_token_base_cost(
model_info=model_info, usage=usage, service_tier=service_tier
)
prompt_cost = _calculate_input_cost(
prompt_tokens_details=prompt_tokens_details,
+7 -2
View File
@@ -18,7 +18,9 @@ def cost_router(call_type: CallTypes) -> Literal["cost_per_token", "cost_per_sec
return "cost_per_token"
def cost_per_token(model: str, usage: Usage, service_tier: Optional[str] = None) -> Tuple[float, float]:
def cost_per_token(
model: str, usage: Usage, service_tier: Optional[str] = None
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@@ -31,7 +33,10 @@ def cost_per_token(model: str, usage: Usage, service_tier: Optional[str] = None)
"""
## CALCULATE INPUT COST
return generic_cost_per_token(
model=model, usage=usage, custom_llm_provider="openai", service_tier=service_tier
model=model,
usage=usage,
custom_llm_provider="openai",
service_tier=service_tier,
)
# ### Non-cached text tokens
# non_cached_text_tokens = usage.prompt_tokens
+14 -20
View File
@@ -9,29 +9,18 @@ from typing import (
Literal,
Mapping,
Optional,
Tuple,
Union,
)
from aiohttp import FormData
from openai._models import BaseModel as OpenAIObject
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
from openai.types.chat.chat_completion import ChatCompletion
from openai.types.completion_usage import (
CompletionTokensDetails,
CompletionUsage,
PromptTokensDetails,
)
from openai.types.moderation import (
Categories,
CategoryAppliedInputTypes,
CategoryScores,
)
from openai.types.moderation_create_response import Moderation, ModerationCreateResponse
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator
from typing_extensions import Callable, Dict, Required, TypedDict, override
from typing_extensions import Required, TypedDict
import litellm
from litellm._uuid import uuid
from litellm.types.llms.base import (
BaseLiteLLMOpenAIResponseObject,
@@ -57,7 +46,6 @@ from .llms.openai import (
OpenAIRealtimeStreamList,
WebSearchOptions,
)
from .rerank import RerankResponse
if TYPE_CHECKING:
from .vector_stores import VectorStoreSearchResponse
@@ -123,12 +111,18 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
max_output_tokens: Required[Optional[int]]
input_cost_per_token: Required[float]
input_cost_per_token_flex: Optional[float] # OpenAI flex service tier pricing
input_cost_per_token_priority: Optional[float] # OpenAI priority service tier pricing
input_cost_per_token_priority: Optional[
float
] # OpenAI priority service tier pricing
cache_creation_input_token_cost: Optional[float]
cache_creation_input_token_cost_above_1hr: Optional[float]
cache_read_input_token_cost: Optional[float]
cache_read_input_token_cost_flex: Optional[float] # OpenAI flex service tier pricing
cache_read_input_token_cost_priority: Optional[float] # OpenAI priority service tier pricing
cache_read_input_token_cost_flex: Optional[
float
] # OpenAI flex service tier pricing
cache_read_input_token_cost_priority: Optional[
float
] # OpenAI priority service tier pricing
input_cost_per_character: Optional[float] # only for vertex ai models
input_cost_per_audio_token: Optional[float]
input_cost_per_token_above_128k_tokens: Optional[float] # only for vertex ai models
@@ -147,7 +141,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_token_batches: Optional[float]
output_cost_per_token: Required[float]
output_cost_per_token_flex: Optional[float] # OpenAI flex service tier pricing
output_cost_per_token_priority: Optional[float] # OpenAI priority service tier pricing
output_cost_per_token_priority: Optional[
float
] # OpenAI priority service tier pricing
output_cost_per_character: Optional[float] # only for vertex ai models
output_cost_per_audio_token: Optional[float]
output_cost_per_token_above_128k_tokens: Optional[
@@ -1141,9 +1137,6 @@ class StreamingChatCompletionChunk(OpenAIChatCompletionChunk):
super().__init__(**kwargs)
from openai.types.chat import ChatCompletionChunk
class ModelResponseBase(OpenAIObject):
id: str
"""A unique identifier for the completion."""
@@ -2592,6 +2585,7 @@ class SpecialEnums(Enum):
class ServiceTier(Enum):
"""Enum for service tier types used in cost calculations."""
FLEX = "flex"
PRIORITY = "priority"
+42 -30
View File
@@ -527,9 +527,6 @@ def get_dynamic_callbacks(
return returned_callbacks
def function_setup( # noqa: PLR0915
original_function: str, rules_obj, start_time, *args, **kwargs
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
@@ -551,9 +548,9 @@ def function_setup( # noqa: PLR0915
function_id: Optional[str] = kwargs["id"] if "id" in kwargs else None
## DYNAMIC CALLBACKS ##
dynamic_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = (
kwargs.pop("callbacks", None)
)
dynamic_callbacks: Optional[
List[Union[str, Callable, CustomLogger]]
] = kwargs.pop("callbacks", None)
all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks)
if len(all_callbacks) > 0:
@@ -793,7 +790,7 @@ def function_setup( # noqa: PLR0915
call_type=call_type,
):
stream = True
logging_obj = get_litellm_logging_class()( # Victim for object pool
logging_obj = get_litellm_logging_class()( # Victim for object pool
model=model, # type: ignore
messages=messages,
stream=stream,
@@ -1302,9 +1299,9 @@ def client(original_function): # noqa: PLR0915
exception=e,
retry_policy=kwargs.get("retry_policy"),
)
kwargs["retry_policy"] = (
reset_retry_policy()
) # prevent infinite loops
kwargs[
"retry_policy"
] = reset_retry_policy() # prevent infinite loops
litellm.num_retries = (
None # set retries to None to prevent infinite loops
)
@@ -3153,10 +3150,10 @@ def pre_process_non_default_params(
if "response_format" in non_default_params:
if provider_config is not None:
non_default_params["response_format"] = (
provider_config.get_json_schema_from_pydantic_object(
response_format=non_default_params["response_format"]
)
non_default_params[
"response_format"
] = provider_config.get_json_schema_from_pydantic_object(
response_format=non_default_params["response_format"]
)
else:
non_default_params["response_format"] = type_to_response_format_param(
@@ -3285,16 +3282,16 @@ def pre_process_optional_params(
True # so that main.py adds the function call to the prompt
)
if "tools" in non_default_params:
optional_params["functions_unsupported_model"] = (
non_default_params.pop("tools")
)
optional_params[
"functions_unsupported_model"
] = non_default_params.pop("tools")
non_default_params.pop(
"tool_choice", None
) # causes ollama requests to hang
elif "functions" in non_default_params:
optional_params["functions_unsupported_model"] = (
non_default_params.pop("functions")
)
optional_params[
"functions_unsupported_model"
] = non_default_params.pop("functions")
elif (
litellm.add_function_to_prompt
): # if user opts to add it to prompt instead
@@ -4387,9 +4384,9 @@ def _count_characters(text: str) -> int:
def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) -> str:
_choices: Union[List[Union[Choices, StreamingChoices]], List[StreamingChoices]] = (
response_obj.choices
)
_choices: Union[
List[Union[Choices, StreamingChoices]], List[StreamingChoices]
] = response_obj.choices
response_str = ""
for choice in _choices:
@@ -4878,16 +4875,24 @@ def _get_model_info_helper( # noqa: PLR0915
max_input_tokens=_model_info.get("max_input_tokens", None),
max_output_tokens=_model_info.get("max_output_tokens", None),
input_cost_per_token=_input_cost_per_token,
input_cost_per_token_flex=_model_info.get("input_cost_per_token_flex", None),
input_cost_per_token_priority=_model_info.get("input_cost_per_token_priority", None),
input_cost_per_token_flex=_model_info.get(
"input_cost_per_token_flex", None
),
input_cost_per_token_priority=_model_info.get(
"input_cost_per_token_priority", None
),
cache_creation_input_token_cost=_model_info.get(
"cache_creation_input_token_cost", None
),
cache_read_input_token_cost=_model_info.get(
"cache_read_input_token_cost", None
),
cache_read_input_token_cost_flex=_model_info.get("cache_read_input_token_cost_flex", None),
cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None),
cache_read_input_token_cost_flex=_model_info.get(
"cache_read_input_token_cost_flex", None
),
cache_read_input_token_cost_priority=_model_info.get(
"cache_read_input_token_cost_priority", None
),
cache_creation_input_token_cost_above_1hr=_model_info.get(
"cache_creation_input_token_cost_above_1hr", None
),
@@ -4912,8 +4917,12 @@ def _get_model_info_helper( # noqa: PLR0915
"output_cost_per_token_batches"
),
output_cost_per_token=_output_cost_per_token,
output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None),
output_cost_per_token_priority=_model_info.get("output_cost_per_token_priority", None),
output_cost_per_token_flex=_model_info.get(
"output_cost_per_token_flex", None
),
output_cost_per_token_priority=_model_info.get(
"output_cost_per_token_priority", None
),
output_cost_per_audio_token=_model_info.get(
"output_cost_per_audio_token", None
),
@@ -6929,7 +6938,10 @@ class ProviderConfigManager:
return litellm.LlamaAPIConfig()
elif litellm.LlmProviders.TEXT_COMPLETION_OPENAI == provider:
return litellm.OpenAITextCompletionConfig()
elif litellm.LlmProviders.COHERE_CHAT == provider or litellm.LlmProviders.COHERE == provider:
elif (
litellm.LlmProviders.COHERE_CHAT == provider
or litellm.LlmProviders.COHERE == provider
):
return litellm.CohereChatConfig()
elif litellm.LlmProviders.SNOWFLAKE == provider:
return litellm.SnowflakeConfig()