Merge pull request #18873 from BerriAI/litellm_staging_01_09_2026

staging 01/09/2025
This commit is contained in:
Sameer Kankute 2026-01-13 17:31:44 +05:30 • committed by GitHub
commit 54f6f55c98
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 748 additions and 142 deletions

View file

@ -594,9 +594,9 @@ class OpenTelemetry(CustomLogger):
def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> Optional[dict]:
"""Extract dynamic headers from kwargs if available."""
standard_callback_dynamic_params: Optional[
StandardCallbackDynamicParams
] = kwargs.get("standard_callback_dynamic_params")
standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = (
kwargs.get("standard_callback_dynamic_params")
)
if not standard_callback_dynamic_params:
return None
@ -797,7 +797,7 @@ class OpenTelemetry(CustomLogger):
and self._token_usage_histogram
):
in_attrs = {**common_attrs, "gen_ai.token.type": "input"}
out_attrs = {**common_attrs, "gen_ai.token.type": "completion"}
out_attrs = {**common_attrs, "gen_ai.token.type": "output"}
self._token_usage_histogram.record(
usage.get("prompt_tokens", 0), attributes=in_attrs
)
@ -1488,21 +1488,21 @@ class OpenTelemetry(CustomLogger):
if usage:
self.safe_set_attribute(
span=span,
key=SpanAttributes.LLM_USAGE_TOTAL_TOKENS.value,
key=SpanAttributes.GEN_AI_USAGE_TOTAL_TOKENS.value,
value=usage.get("total_tokens"),
)
# The number of tokens used in the LLM response (completion).
self.safe_set_attribute(
span=span,
key=SpanAttributes.LLM_USAGE_COMPLETION_TOKENS.value,
key=SpanAttributes.GEN_AI_USAGE_OUTPUT_TOKENS.value,
value=usage.get("completion_tokens"),
)
# The number of tokens used in the LLM prompt.
self.safe_set_attribute(
span=span,
key=SpanAttributes.LLM_USAGE_PROMPT_TOKENS.value,
key=SpanAttributes.GEN_AI_USAGE_INPUT_TOKENS.value,
value=usage.get("prompt_tokens"),
)
@ -1520,53 +1520,75 @@ class OpenTelemetry(CustomLogger):
self.set_tools_attributes(span, tools)
if kwargs.get("messages"):
for idx, prompt in enumerate(kwargs.get("messages")):
if prompt.get("role"):
self.safe_set_attribute(
span=span,
key=f"{SpanAttributes.LLM_PROMPTS.value}.{idx}.role",
value=prompt.get("role"),
)
transformed_messages = (
self._transform_messages_to_otel_semantic_conventions(
kwargs.get("messages")
)
)
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_INPUT_MESSAGES.value,
value=safe_dumps(transformed_messages),
)
if prompt.get("content"):
if not isinstance(prompt.get("content"), str):
prompt["content"] = str(prompt.get("content"))
self.safe_set_attribute(
span=span,
key=f"{SpanAttributes.LLM_PROMPTS.value}.{idx}.content",
value=prompt.get("content"),
)
if kwargs.get("system_instructions"):
transformed_system_instructions = (
self._transform_messages_to_otel_semantic_conventions(
kwargs.get("system_instructions")
)
)
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_SYSTEM_INSTRUCTIONS.value,
value=safe_dumps(transformed_system_instructions),
)
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_OPERATION_NAME.value,
value=(
"chat"
if standard_logging_payload.get("call_type") == "completion"
else standard_logging_payload.get("call_type") or "chat"
),
)
if standard_logging_payload.get("request_id"):
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_REQUEST_ID.value,
value=standard_logging_payload.get("request_id"),
)
#############################################
########## LLM Response Attributes ##########
#############################################
if response_obj is not None:
if response_obj.get("choices"):
transformed_choices = (
self._transform_choices_to_otel_semantic_conventions(
response_obj.get("choices")
)
)
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_OUTPUT_MESSAGES.value,
value=safe_dumps(transformed_choices),
)
finish_reasons = []
for idx, choice in enumerate(response_obj.get("choices")):
if choice.get("finish_reason"):
finish_reasons.append(choice.get("finish_reason"))
if finish_reasons:
self.safe_set_attribute(
span=span,
key=SpanAttributes.GEN_AI_RESPONSE_FINISH_REASONS.value,
value=safe_dumps(finish_reasons),
)
for idx, choice in enumerate(response_obj.get("choices")):
if choice.get("finish_reason"):
self.safe_set_attribute(
span=span,
key=f"{SpanAttributes.LLM_COMPLETIONS.value}.{idx}.finish_reason",
value=choice.get("finish_reason"),
)
if choice.get("message"):
if choice.get("message").get("role"):
self.safe_set_attribute(
span=span,
key=f"{SpanAttributes.LLM_COMPLETIONS.value}.{idx}.role",
value=choice.get("message").get("role"),
)
if choice.get("message").get("content"):
if not isinstance(
choice.get("message").get("content"), str
):
choice["message"]["content"] = str(
choice.get("message").get("content")
)
self.safe_set_attribute(
span=span,
key=f"{SpanAttributes.LLM_COMPLETIONS.value}.{idx}.content",
value=choice.get("message").get("content"),
)
message = choice.get("message")
tool_calls = message.get("tool_calls")
@ -1608,6 +1630,66 @@ class OpenTelemetry(CustomLogger):
primitive_value = self._cast_as_primitive_value_type(value)
span.set_attribute(key, primitive_value)
def _transform_messages_to_otel_semantic_conventions(
self, messages: Union[List[dict], str]
) -> List[dict]:
"""
Transforms LiteLLM/OpenAI style messages into OTEL GenAI 1.38 compliant format.
OTEL expects a 'parts' array instead of a single 'content' string.
"""
if isinstance(messages, str):
# Handle system_instructions passed as a string
return [
{"role": "system", "parts": [{"type": "text", "content": messages}]}
]
transformed = []
for msg in messages:
role = msg.get("role", "user")
content = msg.get("content", "")
parts = []
if isinstance(content, str):
parts.append({"type": "text", "content": content})
elif isinstance(content, list):
# Handle multi-modal content if necessary
for part in content:
if isinstance(part, dict):
parts.append(part)
else:
parts.append({"type": "text", "content": str(part)})
transformed_msg = {"role": role, "parts": parts}
if "id" in msg:
transformed_msg["id"] = msg["id"]
if "tool_calls" in msg:
transformed_msg["tool_calls"] = msg["tool_calls"]
if "tool_call_id" in msg:
transformed_msg["tool_call_id"] = msg["tool_call_id"]
transformed.append(transformed_msg)
return transformed
def _transform_choices_to_otel_semantic_conventions(
self, choices: List[dict]
) -> List[dict]:
"""
Transforms choices into OTEL GenAI 1.38 compliant format for output.messages.
"""
transformed = []
for choice in choices:
message = choice.get("message") or {}
finish_reason = choice.get("finish_reason")
transformed_msg = self._transform_messages_to_otel_semantic_conventions(
[message]
)[0]
if finish_reason:
transformed_msg["finish_reason"] = finish_reason
transformed.append(transformed_msg)
return transformed
def set_raw_request_attributes(self, span: Span, kwargs, response_obj):
try:
kwargs.get("optional_params", {})

View file

@ -830,9 +830,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
@ -1366,12 +1366,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
@ -1393,12 +1393,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):
@ -1484,15 +1484,15 @@ class NewTeamRequest(TeamBase):
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
model_tpm_limit: Optional[Dict[str, int]] = 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"
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
@ -1580,9 +1580,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")
@ -1895,9 +1895,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):
@ -2291,9 +2291,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=())
@ -2828,6 +2828,18 @@ class SpanAttributes(str, enum.Enum):
LLM_RESPONSE_MODEL = "gen_ai.response.model"
LLM_USAGE_COMPLETION_TOKENS = "gen_ai.usage.completion_tokens"
LLM_USAGE_PROMPT_TOKENS = "gen_ai.usage.prompt_tokens"
# OTEL 1.38 attributes
GEN_AI_INPUT_MESSAGES = "gen_ai.input.messages"
GEN_AI_OUTPUT_MESSAGES = "gen_ai.output.messages"
GEN_AI_USAGE_INPUT_TOKENS = "gen_ai.usage.input_tokens"
GEN_AI_USAGE_OUTPUT_TOKENS = "gen_ai.usage.output_tokens"
GEN_AI_USAGE_TOTAL_TOKENS = "gen_ai.usage.total_tokens"
GEN_AI_OPERATION_NAME = "gen_ai.operation.name"
GEN_AI_REQUEST_ID = "gen_ai.request.id"
GEN_AI_SYSTEM_INSTRUCTIONS = "gen_ai.system_instructions"
GEN_AI_RESPONSE_FINISH_REASONS = "gen_ai.response.finish_reasons"
LLM_TOKEN_TYPE = "gen_ai.token.type"
# To be added
# LLM_RESPONSE_FINISH_REASON = "gen_ai.response.finish_reasons"
@ -3253,9 +3265,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):
@ -3470,9 +3482,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):
@ -3615,9 +3627,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
@ -3868,20 +3880,44 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
num_requests_per_day: Optional[int] = None
num_requests_per_month: Optional[int] = None
# Per-request costs
cost_per_request: float = Field(description="Total cost per request (includes margin)")
input_cost_per_request: float = Field(description="Input token cost per request (before margin)")
output_cost_per_request: float = Field(description="Output token cost per request (before margin)")
margin_cost_per_request: float = Field(default=0.0, description="Margin/fee added per request")
cost_per_request: float = Field(
description="Total cost per request (includes margin)"
)
input_cost_per_request: float = Field(
description="Input token cost per request (before margin)"
)
output_cost_per_request: float = Field(
description="Output token cost per request (before margin)"
)
margin_cost_per_request: float = Field(
default=0.0, description="Margin/fee added per request"
)
# Daily costs (if num_requests_per_day provided)
daily_cost: Optional[float] = Field(default=None, description="Total daily cost (includes margin)")
daily_input_cost: Optional[float] = Field(default=None, description="Daily input token cost")
daily_output_cost: Optional[float] = Field(default=None, description="Daily output token cost")
daily_margin_cost: Optional[float] = Field(default=None, description="Daily margin/fee")
daily_cost: Optional[float] = Field(
default=None, description="Total daily cost (includes margin)"
)
daily_input_cost: Optional[float] = Field(
default=None, description="Daily input token cost"
)
daily_output_cost: Optional[float] = Field(
default=None, description="Daily output token cost"
)
daily_margin_cost: Optional[float] = Field(
default=None, description="Daily margin/fee"
)
# Monthly costs (if num_requests_per_month provided)
monthly_cost: Optional[float] = Field(default=None, description="Total monthly cost (includes margin)")
monthly_input_cost: Optional[float] = Field(default=None, description="Monthly input token cost")
monthly_output_cost: Optional[float] = Field(default=None, description="Monthly output token cost")
monthly_margin_cost: Optional[float] = Field(default=None, description="Monthly margin/fee")
monthly_cost: Optional[float] = Field(
default=None, description="Total monthly cost (includes margin)"
)
monthly_input_cost: Optional[float] = Field(
default=None, description="Monthly input token cost"
)
monthly_output_cost: Optional[float] = Field(
default=None, description="Monthly output token cost"
)
monthly_margin_cost: Optional[float] = Field(
default=None, description="Monthly margin/fee"
)
# Pricing info
input_cost_per_token: Optional[float] = None
output_cost_per_token: Optional[float] = None

View file

@ -472,7 +472,9 @@ class ResponseAPILoggingUtils:
output_tokens_details = getattr(response_api_usage, "output_tokens_details", None)
if output_tokens_details:
completion_tokens_details = CompletionTokensDetailsWrapper(
reasoning_tokens=getattr(output_tokens_details, "reasoning_tokens", None)
reasoning_tokens=getattr(output_tokens_details, "reasoning_tokens", None),
image_tokens=getattr(output_tokens_details, "image_tokens", None),
text_tokens=getattr(output_tokens_details, "text_tokens", None),
)
chat_usage = Usage(

View file

@ -4993,6 +4993,25 @@ def _strip_model_name(model: str, custom_llm_provider: Optional[str]) -> str:
return model
def _get_model_cost_key(potential_key: str) -> Optional[str]:
"""
Get the actual key from model_cost, with case-insensitive fallback.
Returns the key if found (exact match preferred, then case-insensitive), or None if not found.
"""
# Try exact match first (most common case, O(1))
if potential_key in litellm.model_cost:
return potential_key
# Fallback to case-insensitive match
potential_key_lower = potential_key.lower()
for key in litellm.model_cost:
if key.lower() == potential_key_lower:
return key
return None
def _get_model_info_from_model_cost(key: str) -> dict:
return litellm.model_cost[key]
@ -5142,10 +5161,10 @@ def _is_potential_model_name_in_model_cost(
potential_model_names: PotentialModelNamesAndCustomLLMProvider,
) -> bool:
"""
Check if the potential model name is in the model cost.
Check if the potential model name is in the model cost (case-insensitive).
"""
return any(
potential_model_name in litellm.model_cost
_get_model_cost_key(str(potential_model_name)) is not None
for potential_model_name in potential_model_names.values()
)
@ -5223,44 +5242,51 @@ def _get_model_info_helper( # noqa: PLR0915
_model_info: Optional[Dict[str, Any]] = None
key: Optional[str] = None
if combined_model_name in litellm.model_cost:
key = combined_model_name
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None and model in litellm.model_cost:
key = model
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if (
_model_info is None
and combined_stripped_model_name in litellm.model_cost
):
key = combined_stripped_model_name
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None and stripped_model_name in litellm.model_cost:
key = stripped_model_name
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None and split_model in litellm.model_cost:
key = split_model
# Use case-insensitive lookup for all model name checks
_matched_key = _get_model_cost_key(combined_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(combined_stripped_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(stripped_model_name)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None:
_matched_key = _get_model_cost_key(split_model)
if _matched_key is not None:
key = _matched_key
_model_info = _get_model_info_from_model_cost(key=cast(str, key))
if not _check_provider_match(
model_info=_model_info, custom_llm_provider=custom_llm_provider
):
_model_info = None
if _model_info is None or key is None:
raise ValueError(

View file

@ -372,3 +372,86 @@ def test_get_model_info_cost_calculator_bedrock_region_cris_stripped(model, prov
print("info", info)
assert info["key"] == "us.anthropic.claude-3-haiku-20240307-v1:0"
assert info["litellm_provider"] == "bedrock"
def test_get_model_info_case_insensitive_lookup(monkeypatch):
"""
Test that model info lookup is case-insensitive.
This ensures that users can use lowercase model names even when the model cost
map has mixed-case keys (e.g., "Qwen/Qwen3-Next-80B-A3B-Thinking").
Related Slack discussion: Users were getting "does not support parameters: ['tools']"
errors when using lowercase model names like "qwen/qwen3-next-80b-a3b-thinking"
because the lookup was case-sensitive.
"""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
# Register a test model with mixed-case name
litellm.register_model(
{
"together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": {
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.0002,
"litellm_provider": "together_ai",
"supports_function_calling": True,
}
}
)
# Test 1: Exact case should work
info = litellm.get_model_info(
model="Qwen/Qwen3-Next-80B-A3B-Thinking", custom_llm_provider="together_ai"
)
assert info is not None
assert info["supports_function_calling"] is True
# Test 2: Lowercase should also work (case-insensitive lookup)
info_lower = litellm.get_model_info(
model="qwen/qwen3-next-80b-a3b-thinking", custom_llm_provider="together_ai"
)
assert info_lower is not None
assert info_lower["supports_function_calling"] is True
# Test 3: Mixed case should also work
info_mixed = litellm.get_model_info(
model="QWEN/qwen3-NEXT-80b-a3b-thinking", custom_llm_provider="together_ai"
)
assert info_mixed is not None
assert info_mixed["supports_function_calling"] is True
def test_get_model_info_case_insensitive_supports_function_calling(monkeypatch):
"""
Test that supports_function_calling check works with case-insensitive model lookup.
"""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
# Register a model with mixed-case name that supports function calling
litellm.register_model(
{
"test_provider/TestModel-ABC": {
"input_cost_per_token": 0.0001,
"output_cost_per_token": 0.0002,
"litellm_provider": "test_provider",
"supports_function_calling": True,
}
}
)
# Test that supports_function_calling works with lowercase model name
from litellm.utils import supports_function_calling
# Exact case
assert (
supports_function_calling("TestModel-ABC", custom_llm_provider="test_provider")
is True
)
# Lowercase (should now work with case-insensitive lookup)
assert (
supports_function_calling("testmodel-abc", custom_llm_provider="test_provider")
is True
)

View file

@ -27,6 +27,9 @@ exporter = InMemorySpanExporter()
@pytest.mark.parametrize("streaming", [True, False])
async def test_async_otel_callback(streaming):
litellm.set_verbose = True
# Clear exporter at the start to ensure clean state
exporter.clear()
litellm.callbacks = [OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter))]
@ -83,9 +86,9 @@ def validate_litellm_request(span):
"llm.user",
"gen_ai.response.id",
"gen_ai.response.model",
"llm.usage.total_tokens",
"gen_ai.usage.completion_tokens",
"gen_ai.usage.prompt_tokens",
"gen_ai.usage.total_tokens",
"gen_ai.usage.output_tokens",
"gen_ai.usage.input_tokens",
]
# get the str of all the span attributes
@ -149,6 +152,10 @@ async def test_awesome_otel_with_message_logging_off(streaming, global_redact):
tests when OpenTelemetry(message_logging=False) is set
"""
litellm.set_verbose = True
# Clear exporter at the start to ensure clean state
exporter.clear()
litellm.callbacks = [OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter))]
if global_redact is False:
otel_logger = OpenTelemetry(
@ -201,9 +208,9 @@ def validate_redacted_message_span_attributes(span):
"llm.request.type",
"gen_ai.response.id",
"gen_ai.response.model",
"llm.usage.total_tokens",
"gen_ai.usage.completion_tokens",
"gen_ai.usage.prompt_tokens",
"gen_ai.usage.total_tokens",
"gen_ai.usage.output_tokens",
"gen_ai.usage.input_tokens",
]
_all_attributes = set(
@ -230,6 +237,8 @@ def validate_redacted_message_span_attributes(span):
attr.startswith("metadata.")
or attr.startswith("hidden_params")
or attr.startswith("gen_ai.cost.")
or attr.startswith("gen_ai.operation.")
or attr.startswith("gen_ai.request.")
), f"Non-metadata attribute found: {attr}"
pass

View file

@ -80,6 +80,53 @@ def test_get_usage(response_obj, expected_values):
assert usage.total_tokens == expected_values[2]
def test_get_usage_from_image_generation_response():
"""
Test that image generation usage (with input_tokens/output_tokens format)
is correctly transformed to standard usage format with image_tokens preserved.
Note: get_usage_from_response_obj() is used by multiple endpoints including
/images/generations and Response API (/responses), both of which use the
input_tokens/output_tokens format instead of prompt_tokens/completion_tokens.
This tests the fix for the bug where image_tokens were being lost during
spend log creation for /images/generations endpoint.
"""
# Simulating image generation response usage from OpenAI
response_obj = {
"usage": {
"input_tokens": 13,
"output_tokens": 372,
"total_tokens": 385,
"input_tokens_details": {
"image_tokens": 0,
"text_tokens": 13,
},
"output_tokens_details": {
"image_tokens": 272,
"text_tokens": 100,
},
}
}
usage = StandardLoggingPayloadSetup.get_usage_from_response_obj(response_obj)
# Check basic token counts are mapped correctly
assert usage.prompt_tokens == 13
assert usage.completion_tokens == 372
assert usage.total_tokens == 385
# Check that prompt_tokens_details contains image_tokens and text_tokens
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.image_tokens == 0
assert usage.prompt_tokens_details.text_tokens == 13
# Check that completion_tokens_details contains image_tokens and text_tokens
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.image_tokens == 272
assert usage.completion_tokens_details.text_tokens == 100
def test_get_additional_headers():
additional_headers = {
"x-ratelimit-limit-requests": "2000",

View file

@ -1869,3 +1869,237 @@ class TestOpenTelemetryExternalSpan(unittest.TestCase):
parent_span_finished.attributes,
"Parent span should have model attribute from LiteLLM even on failure"
)
class TestOpenTelemetrySemanticConventions138(unittest.TestCase):
"""
Test suite for OpenTelemetry 1.38 Semantic Conventions compliance.
These tests verify that LiteLLM emits span attributes following the
OpenTelemetry GenAI semantic conventions v1.38, including:
- gen_ai.input.messages (JSON string with parts array)
- gen_ai.output.messages (JSON string with parts array)
- gen_ai.usage.input_tokens / output_tokens (new naming)
- gen_ai.response.finish_reasons (JSON array)
See: https://github.com/BerriAI/litellm/issues/17794
"""
def test_input_messages_uses_parts_structure(self):
"""
Test that gen_ai.input.messages uses the OTEL 1.38 parts array structure.
Expected format:
[{"role": "user", "parts": [{"type": "text", "content": "Hello"}]}]
"""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello world"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "completion",
"metadata": {},
},
}
response_obj = {
"id": "test-response-id",
"model": "gpt-4",
"choices": [
{
"finish_reason": "stop",
"message": {"role": "assistant", "content": "Hi there!"},
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
# Find the call that set gen_ai.input.messages
input_messages_calls = [
call for call in mock_span.set_attribute.call_args_list
if call[0][0] == "gen_ai.input.messages"
]
self.assertEqual(len(input_messages_calls), 1, "Should have exactly one gen_ai.input.messages attribute")
input_messages_value = input_messages_calls[0][0][1]
parsed = json.loads(input_messages_value)
# Verify structure
self.assertIsInstance(parsed, list)
self.assertEqual(len(parsed), 1)
self.assertEqual(parsed[0]["role"], "user")
self.assertIn("parts", parsed[0])
self.assertEqual(parsed[0]["parts"][0]["type"], "text")
self.assertEqual(parsed[0]["parts"][0]["content"], "Hello world")
def test_output_messages_uses_parts_structure(self):
"""
Test that gen_ai.output.messages uses the OTEL 1.38 parts array structure.
Expected format:
[{"role": "assistant", "parts": [{"type": "text", "content": "Hi!"}], "finish_reason": "stop"}]
"""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "completion",
"metadata": {},
},
}
response_obj = {
"id": "test-response-id",
"model": "gpt-4",
"choices": [
{
"finish_reason": "stop",
"message": {"role": "assistant", "content": "Hello back!"},
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
# Find the call that set gen_ai.output.messages
output_messages_calls = [
call for call in mock_span.set_attribute.call_args_list
if call[0][0] == "gen_ai.output.messages"
]
self.assertEqual(len(output_messages_calls), 1, "Should have exactly one gen_ai.output.messages attribute")
output_messages_value = output_messages_calls[0][0][1]
parsed = json.loads(output_messages_value)
# Verify structure
self.assertIsInstance(parsed, list)
self.assertEqual(len(parsed), 1)
self.assertEqual(parsed[0]["role"], "assistant")
self.assertIn("parts", parsed[0])
self.assertEqual(parsed[0]["parts"][0]["type"], "text")
self.assertEqual(parsed[0]["parts"][0]["content"], "Hello back!")
self.assertEqual(parsed[0]["finish_reason"], "stop")
def test_usage_tokens_use_new_naming_convention(self):
"""
Test that token usage uses the OTEL 1.38 naming convention:
- gen_ai.usage.input_tokens (not prompt_tokens)
- gen_ai.usage.output_tokens (not completion_tokens)
"""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "completion",
"metadata": {},
},
}
response_obj = {
"id": "test-response-id",
"model": "gpt-4",
"choices": [],
"usage": {"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150},
}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
# Verify new naming convention is used
mock_span.set_attribute.assert_any_call("gen_ai.usage.input_tokens", 100)
mock_span.set_attribute.assert_any_call("gen_ai.usage.output_tokens", 50)
mock_span.set_attribute.assert_any_call("gen_ai.usage.total_tokens", 150)
def test_finish_reasons_is_json_array(self):
"""
Test that gen_ai.response.finish_reasons is a proper JSON array.
Expected: '["stop"]' (not "['stop']")
"""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "completion",
"metadata": {},
},
}
response_obj = {
"id": "test-response-id",
"model": "gpt-4",
"choices": [
{"finish_reason": "stop", "message": {"role": "assistant", "content": "Hi"}},
],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
# Find the call that set gen_ai.response.finish_reasons
finish_reasons_calls = [
call for call in mock_span.set_attribute.call_args_list
if call[0][0] == "gen_ai.response.finish_reasons"
]
self.assertEqual(len(finish_reasons_calls), 1, "Should have exactly one gen_ai.response.finish_reasons attribute")
finish_reasons_value = finish_reasons_calls[0][0][1]
# Verify it's valid JSON (not Python repr)
parsed = json.loads(finish_reasons_value)
self.assertEqual(parsed, ["stop"])
def test_operation_name_is_chat_for_completion(self):
"""
Test that gen_ai.operation.name is 'chat' for completion calls.
"""
otel = OpenTelemetry()
mock_span = MagicMock()
kwargs = {
"model": "gpt-4",
"messages": [{"role": "user", "content": "Hello"}],
"optional_params": {},
"litellm_params": {"custom_llm_provider": "openai"},
"standard_logging_object": {
"id": "test-id",
"call_type": "completion",
"metadata": {},
},
}
response_obj = {
"id": "test-response-id",
"model": "gpt-4",
"choices": [],
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
}
otel.set_attributes(span=mock_span, kwargs=kwargs, response_obj=response_obj)
mock_span.set_attribute.assert_any_call("gen_ai.operation.name", "chat")

View file

@ -222,3 +222,90 @@ class TestResponseAPILoggingUtils:
assert result.prompt_tokens == 15
assert result.completion_tokens == 25
assert result.total_tokens == 40 # 15 + 25
def test_transform_response_api_usage_with_image_tokens(self):
"""Test transformation handles image_tokens from image generation responses.
Note: _transform_response_api_usage_to_chat_usage() is used by multiple
endpoints including /images/generations and Response API (/responses),
both of which use the input_tokens/output_tokens format.
This tests the fix for image generation responses that include image_tokens
in both input_tokens_details and output_tokens_details.
Example from gpt-image-1.5:
- input: text prompt with 13 tokens
- output: generated image with 272 image tokens + 100 text tokens
"""
# Setup - simulating image generation usage from OpenAI
usage = {
"input_tokens": 13,
"output_tokens": 372,
"total_tokens": 385,
"input_tokens_details": {
"image_tokens": 0,
"text_tokens": 13,
},
"output_tokens_details": {
"image_tokens": 272,
"text_tokens": 100,
},
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
# Assert - verify basic token counts
assert isinstance(result, Usage)
assert result.prompt_tokens == 13
assert result.completion_tokens == 372
assert result.total_tokens == 385
# Assert - verify prompt_tokens_details includes image_tokens and text_tokens
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.image_tokens == 0
assert result.prompt_tokens_details.text_tokens == 13
# Assert - verify completion_tokens_details includes image_tokens and text_tokens
assert result.completion_tokens_details is not None
assert result.completion_tokens_details.image_tokens == 272
assert result.completion_tokens_details.text_tokens == 100
def test_transform_response_api_usage_mixed_details(self):
"""Test transformation handles mixed token details (cached + image + audio)."""
# Setup - hypothetical usage with mixed token types
usage = {
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
"input_tokens_details": {
"cached_tokens": 50,
"audio_tokens": 10,
"image_tokens": 20,
"text_tokens": 20,
},
"output_tokens_details": {
"reasoning_tokens": 30,
"image_tokens": 100,
"text_tokens": 70,
},
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
# Assert - all token detail types should be preserved
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cached_tokens == 50
assert result.prompt_tokens_details.audio_tokens == 10
assert result.prompt_tokens_details.image_tokens == 20
assert result.prompt_tokens_details.text_tokens == 20
assert result.completion_tokens_details is not None
assert result.completion_tokens_details.reasoning_tokens == 30
assert result.completion_tokens_details.image_tokens == 100
assert result.completion_tokens_details.text_tokens == 70