Merge pull request #14122 from BerriAI/litellm_dev_08_30_2025_p1

Braintrust - fix logging when OTEL is enabled + Gemini - add 'thoughtSignature' support via 'thinking_blocks'
This commit is contained in:
Krish Dholakia 2025-09-01 22:42:12 -07:00 • committed by GitHub
commit 6d36219a4d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 504 additions and 262 deletions

View file

@ -1,13 +1,11 @@
# What is this?
## Log success + failure events to Braintrust
import copy
import os
from datetime import datetime
from typing import Dict, Optional
import httpx
from pydantic import BaseModel
import litellm
from litellm import verbose_logger
@ -24,7 +22,6 @@ API_BASE = "https://api.braintrustdata.com/v1"
def get_utc_datetime():
import datetime as dt
from datetime import datetime
if hasattr(dt, "UTC"):
return datetime.now(dt.UTC) # type: ignore
@ -45,9 +42,9 @@ class BraintrustLogger(CustomLogger):
"Authorization": "Bearer " + self.api_key,
"Content-Type": "application/json",
}
self._project_id_cache: Dict[
str, str
] = {} # Cache mapping project names to IDs
self._project_id_cache: Dict[str, str] = (
{}
) # Cache mapping project names to IDs
self.global_braintrust_http_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
@ -108,43 +105,6 @@ class BraintrustLogger(CustomLogger):
except httpx.HTTPStatusError as e:
raise Exception(f"Failed to register project: {e.response.text}")
@staticmethod
def add_metadata_from_header(litellm_params: dict, metadata: dict) -> dict:
"""
Adds metadata from proxy request headers to Braintrust logging if keys start with "braintrust_"
and overwrites litellm_params.metadata if already included.
For example if you want to append your trace to an existing `trace_id` via header, send
`headers: { ..., langfuse_existing_trace_id: your-existing-trace-id }` via proxy request.
"""
if litellm_params is None:
return metadata
if litellm_params.get("proxy_server_request") is None:
return metadata
if metadata is None:
metadata = {}
proxy_headers = (
litellm_params.get("proxy_server_request", {}).get("headers", {}) or {}
)
for metadata_param_key in proxy_headers:
if metadata_param_key.startswith("braintrust"):
trace_param_key = metadata_param_key.replace("braintrust", "", 1)
if trace_param_key in metadata:
verbose_logger.warning(
f"Overwriting Braintrust `{trace_param_key}` from request header"
)
else:
verbose_logger.debug(
f"Found Braintrust `{trace_param_key}` in request header"
)
metadata[trace_param_key] = proxy_headers.get(metadata_param_key)
return metadata
async def create_default_project_and_experiment(self):
project = await self.global_braintrust_http_handler.post(
f"{self.api_base}/project", headers=self.headers, json={"name": "litellm"}
@ -169,7 +129,9 @@ class BraintrustLogger(CustomLogger):
verbose_logger.debug("REACHES BRAINTRUST SUCCESS")
try:
litellm_call_id = kwargs.get("litellm_call_id")
standard_logging_object = kwargs.get("standard_logging_object", {})
prompt = {"messages": kwargs.get("messages")}
output = None
choices = []
if response_obj is not None and (
@ -192,33 +154,13 @@ class BraintrustLogger(CustomLogger):
):
output = response_obj["data"]
litellm_params = kwargs.get("litellm_params", {})
metadata = (
litellm_params.get("metadata", {}) or {}
) # if litellm_params['metadata'] == None
metadata = self.add_metadata_from_header(litellm_params, metadata)
clean_metadata = {}
try:
metadata = copy.deepcopy(
metadata
) # Avoid modifying the original metadata
except Exception:
new_metadata = {}
for key, value in metadata.items():
if (
isinstance(value, list)
or isinstance(value, dict)
or isinstance(value, str)
or isinstance(value, int)
or isinstance(value, float)
):
new_metadata[key] = copy.deepcopy(value)
metadata = new_metadata
litellm_params = kwargs.get("litellm_params", {}) or {}
dynamic_metadata = litellm_params.get("metadata", {}) or {}
# Get project_id from metadata or create default if needed
project_id = metadata.get("project_id")
project_id = dynamic_metadata.get("project_id")
if project_id is None:
project_name = metadata.get("project_name")
project_name = dynamic_metadata.get("project_name")
project_id = (
self.get_project_id_sync(project_name) if project_name else None
)
@ -229,8 +171,9 @@ class BraintrustLogger(CustomLogger):
project_id = self.default_project_id
tags = []
if isinstance(metadata, dict):
for key, value in metadata.items():
if isinstance(dynamic_metadata, dict):
for key, value in dynamic_metadata.items():
# generate langfuse tags - Default Tags sent to Langfuse from LiteLLM Proxy
if (
litellm.langfuse_default_tags is not None
@ -239,25 +182,12 @@ class BraintrustLogger(CustomLogger):
):
tags.append(f"{key}:{value}")
# clean litellm metadata before logging
if key in [
"headers",
"endpoint",
"caching_groups",
"previous_models",
]:
continue
else:
clean_metadata[key] = value
if (
isinstance(value, str) and key not in standard_logging_object
): # support logging dynamic metadata to braintrust
standard_logging_object[key] = value
cost = kwargs.get("response_cost", None)
if cost is not None:
clean_metadata["litellm_response_cost"] = cost
# metadata.model is required for braintrust to calculate the "Estimated cost" metric
litellm_model = kwargs.get("model", None)
if litellm_model is not None:
clean_metadata["model"] = litellm_model
metrics: Optional[dict] = None
usage_obj = getattr(response_obj, "usage", None)
@ -275,12 +205,12 @@ class BraintrustLogger(CustomLogger):
}
# Allow metadata override for span name
span_name = metadata.get("span_name", "Chat Completion")
span_name = dynamic_metadata.get("span_name", "Chat Completion")
request_data = {
"id": litellm_call_id,
"input": prompt["messages"],
"metadata": clean_metadata,
"metadata": standard_logging_object,
"tags": tags,
"span_attributes": {"name": span_name, "type": "llm"},
}
@ -312,6 +242,7 @@ class BraintrustLogger(CustomLogger):
verbose_logger.debug("REACHES BRAINTRUST SUCCESS")
try:
litellm_call_id = kwargs.get("litellm_call_id")
standard_logging_object = kwargs.get("standard_logging_object", {})
prompt = {"messages": kwargs.get("messages")}
output = None
choices = []
@ -336,32 +267,12 @@ class BraintrustLogger(CustomLogger):
output = response_obj["data"]
litellm_params = kwargs.get("litellm_params", {})
metadata = (
litellm_params.get("metadata", {}) or {}
) # if litellm_params['metadata'] == None
metadata = self.add_metadata_from_header(litellm_params, metadata)
clean_metadata = {}
new_metadata = {}
for key, value in metadata.items():
if (
isinstance(value, list)
or isinstance(value, str)
or isinstance(value, int)
or isinstance(value, float)
):
new_metadata[key] = value
elif isinstance(value, BaseModel):
new_metadata[key] = value.model_dump_json()
elif isinstance(value, dict):
for k, v in value.items():
if isinstance(v, datetime):
value[k] = v.isoformat()
new_metadata[key] = value
dynamic_metadata = litellm_params.get("metadata", {}) or {}
# Get project_id from metadata or create default if needed
project_id = metadata.get("project_id")
project_id = dynamic_metadata.get("project_id")
if project_id is None:
project_name = metadata.get("project_name")
project_name = dynamic_metadata.get("project_name")
project_id = (
await self.get_project_id_async(project_name)
if project_name
@ -374,8 +285,9 @@ class BraintrustLogger(CustomLogger):
project_id = self.default_project_id
tags = []
if isinstance(metadata, dict):
for key, value in metadata.items():
if isinstance(dynamic_metadata, dict):
for key, value in dynamic_metadata.items():
# generate langfuse tags - Default Tags sent to Langfuse from LiteLLM Proxy
if (
litellm.langfuse_default_tags is not None
@ -384,25 +296,12 @@ class BraintrustLogger(CustomLogger):
):
tags.append(f"{key}:{value}")
# clean litellm metadata before logging
if key in [
"headers",
"endpoint",
"caching_groups",
"previous_models",
]:
continue
else:
clean_metadata[key] = value
if (
isinstance(value, str) and key not in standard_logging_object
): # support logging dynamic metadata to braintrust
standard_logging_object[key] = value
cost = kwargs.get("response_cost", None)
if cost is not None:
clean_metadata["litellm_response_cost"] = cost
# metadata.model is required for braintrust to calculate the "Estimated cost" metric
litellm_model = kwargs.get("model", None)
if litellm_model is not None:
clean_metadata["model"] = litellm_model
metrics: Optional[dict] = None
usage_obj = getattr(response_obj, "usage", None)
@ -430,13 +329,13 @@ class BraintrustLogger(CustomLogger):
)
# Allow metadata override for span name
span_name = metadata.get("span_name", "Chat Completion")
span_name = dynamic_metadata.get("span_name", "Chat Completion")
request_data = {
"id": litellm_call_id,
"input": prompt["messages"],
"output": output,
"metadata": clean_metadata,
"metadata": standard_logging_object,
"tags": tags,
"span_attributes": {"name": span_name, "type": "llm"},
}

View file

@ -1,5 +1,6 @@
import json
from typing import Any, Union
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH

View file

@ -105,6 +105,64 @@ def _process_gemini_image(image_url: str, format: Optional[str] = None) -> PartT
raise e
def _snake_to_camel(snake_str: str) -> str:
"""Convert snake_case to camelCase"""
components = snake_str.split("_")
return components[0] + "".join(x.capitalize() for x in components[1:])
def _camel_to_snake(camel_str: str) -> str:
"""Convert camelCase to snake_case"""
import re
return re.sub(r"(?<!^)(?=[A-Z])", "_", camel_str).lower()
def _get_equivalent_key(key: str, available_keys: set) -> Optional[str]:
"""
Get the equivalent key from available keys, checking both camelCase and snake_case variants
"""
if key in available_keys:
return key
# Try camelCase version
camel_key = _snake_to_camel(key)
if camel_key in available_keys:
return camel_key
# Try snake_case version
snake_key = _camel_to_snake(key)
if snake_key in available_keys:
return snake_key
return None
def check_if_part_exists_in_parts(
parts: List[PartType], part: PartType, excluded_keys: List[str] = []
) -> bool:
"""
Check if a part exists in a list of parts
Handles both camelCase and snake_case key variations (e.g., function_call vs functionCall)
"""
keys_to_compare = set(part.keys()) - set(excluded_keys)
for p in parts:
p_keys = set(p.keys())
# Check if all keys in part have equivalent values in p
match_found = True
for key in keys_to_compare:
equivalent_key = _get_equivalent_key(key, p_keys)
if equivalent_key is None or p.get(equivalent_key, None) != part.get(
key, None
):
match_found = False
break
if match_found:
return True
return False
def _gemini_convert_messages_with_history( # noqa: PLR0915
messages: List[AllMessageValues],
) -> List[ContentType]:
@ -236,10 +294,33 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
assistant_msg = ChatCompletionAssistantMessage(**msg_dict) # type: ignore
_message_content = assistant_msg.get("content", None)
reasoning_content = assistant_msg.get("reasoning_content", None)
thinking_blocks = assistant_msg.get("thinking_blocks")
if reasoning_content is not None:
assistant_content.append(
PartType(thought=True, text=reasoning_content)
)
if thinking_blocks is not None:
for block in thinking_blocks:
block_thinking_str = block.get("thinking")
block_signature = block.get("signature")
if (
block_thinking_str is not None
and block_signature is not None
):
try:
assistant_content.append(
PartType(
thoughtSignature=block_signature,
**json.loads(block_thinking_str),
)
)
except Exception:
assistant_content.append(
PartType(
thoughtSignature=block_signature,
text=block_thinking_str,
)
)
if _message_content is not None and isinstance(_message_content, list):
_parts = []
for element in _message_content:
@ -262,9 +343,17 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
assistant_msg.get("tool_calls", []) is not None
or assistant_msg.get("function_call") is not None
): # support assistant tool invoke conversion
assistant_content.extend(
convert_to_gemini_tool_call_invoke(assistant_msg)
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
assistant_msg
)
## check if gemini_tool_call already exists in assistant_content
for gemini_tool_call_part in gemini_tool_call_parts:
if not check_if_part_exists_in_parts(
assistant_content,
gemini_tool_call_part,
excluded_keys=["thoughtSignature"],
):
assistant_content.append(gemini_tool_call_part)
last_message_with_tool_calls = assistant_msg
msg_i += 1
@ -476,6 +565,7 @@ async def async_transform_request_body(
optional_params=optional_params,
)
def _default_user_message_when_system_message_passed() -> ChatCompletionUserMessage:
"""
Returns a default user message when a "system" message is passed in gemini fails.
@ -484,6 +574,7 @@ def _default_user_message_when_system_message_passed() -> ChatCompletionUserMess
"""
return ChatCompletionUserMessage(content=".", role="user")
def _transform_system_message(
supports_system_message: bool, messages: List[AllMessageValues]
) -> Tuple[Optional[SystemInstructions], List[AllMessageValues]]:

View file

@ -43,6 +43,7 @@ from litellm.types.llms.gemini import BidiGenerateContentServerMessage
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionResponseMessage,
ChatCompletionThinkingBlock,
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolParamFunctionChunk,
@ -792,7 +793,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
content_str += _content_str
return content_str, reasoning_content_str
def _extract_thinking_blocks_from_parts(
self, parts: List[HttpxPartType]
) -> List[ChatCompletionThinkingBlock]:
"""Extract thinking blocks from parts if present"""
thinking_blocks: List[ChatCompletionThinkingBlock] = []
for part in parts:
if "thoughtSignature" in part:
part_copy = part.copy()
part_copy.pop("thoughtSignature")
thinking_blocks.append(
ChatCompletionThinkingBlock(
type="thinking",
thinking=json.dumps(part_copy),
signature=part["thoughtSignature"],
)
)
return thinking_blocks
def _extract_image_response_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[ImageURLObject]:
@ -804,10 +823,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if mime_type.startswith("image/"):
# Convert base64 data to data URI format
data_uri = f"data:{mime_type};base64,{data}"
return ImageURLObject(
url=data_uri,
detail="auto"
)
return ImageURLObject(url=data_uri, detail="auto")
return None
def _extract_audio_response_from_parts(
@ -1127,7 +1143,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif web_search_queries:
web_search_requests = len(grounding_metadata)
return web_search_requests
@staticmethod
def _create_streaming_choice(
chat_completion_message: ChatCompletionResponseMessage,
@ -1151,9 +1167,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
index=candidate.get("index", idx),
delta=Delta(
content=chat_completion_message.get("content"),
reasoning_content=chat_completion_message.get(
"reasoning_content"
),
reasoning_content=chat_completion_message.get("reasoning_content"),
tool_calls=tools,
image=image_response,
function_call=functions,
@ -1164,13 +1178,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return choice
@staticmethod
def _extract_candidate_metadata(candidate: Candidates) -> Tuple[List[dict], List[dict], List, List]:
def _extract_candidate_metadata(
candidate: Candidates,
) -> Tuple[List[dict], List[dict], List, List]:
"""
Extract metadata from a single candidate response.
Returns:
grounding_metadata: List[dict]
url_context_metadata: List[dict]
url_context_metadata: List[dict]
safety_ratings: List
citation_metadata: List
"""
@ -1178,7 +1194,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
url_context_metadata: List[dict] = []
safety_ratings: List = []
citation_metadata: List = []
if "groundingMetadata" in candidate:
if isinstance(candidate["groundingMetadata"], list):
grounding_metadata.extend(candidate["groundingMetadata"]) # type: ignore
@ -1194,8 +1210,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if "urlContextMetadata" in candidate:
# Add URL context metadata to grounding metadata
url_context_metadata.append(cast(dict, candidate["urlContextMetadata"]))
return grounding_metadata, url_context_metadata, safety_ratings, citation_metadata
return (
grounding_metadata,
url_context_metadata,
safety_ratings,
citation_metadata,
)
@staticmethod
def _process_candidates(
@ -1227,6 +1248,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
tools: Optional[List[ChatCompletionToolCallChunk]] = []
functions: Optional[ChatCompletionToolCallFunctionChunk] = None
cumulative_tool_call_index: int = 0
thinking_blocks: Optional[List[ChatCompletionThinkingBlock]] = None
for idx, candidate in enumerate(_candidates):
if "content" not in candidate:
@ -1239,7 +1261,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
candidate_safety_ratings,
candidate_citation_metadata,
) = VertexGeminiConfig._extract_candidate_metadata(candidate)
grounding_metadata.extend(candidate_grounding_metadata)
url_context_metadata.extend(candidate_url_context_metadata)
safety_ratings.extend(candidate_safety_ratings)
@ -1264,6 +1286,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
)
thinking_blocks = (
VertexGeminiConfig()._extract_thinking_blocks_from_parts(
parts=candidate["content"]["parts"]
)
)
if audio_response is not None:
cast(Dict[str, Any], chat_completion_message)[
"audio"
@ -1271,7 +1299,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
chat_completion_message["content"] = None # OpenAI spec
if image_response is not None:
# Handle image response - combine with text content into structured format
cast(Dict[str, Any], chat_completion_message)["image"] = image_response
cast(Dict[str, Any], chat_completion_message)[
"image"
] = image_response
if content is not None:
chat_completion_message["content"] = content
@ -1298,15 +1328,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if functions is not None:
chat_completion_message["function_call"] = functions
if thinking_blocks is not None:
chat_completion_message["thinking_blocks"] = thinking_blocks # type: ignore
if isinstance(model_response, ModelResponseStream):
choice = VertexGeminiConfig._create_streaming_choice(
chat_completion_message=chat_completion_message,
candidate=candidate,
idx=idx,
tools=tools,
functions=functions,
candidate=candidate,
idx=idx,
tools=tools,
functions=functions,
chat_completion_logprobs=chat_completion_logprobs,
image_response=image_response
image_response=image_response,
)
model_response.choices.append(choice)
elif isinstance(model_response, ModelResponse):

View file

@ -19,12 +19,4 @@ router_settings:
litellm_settings:
callbacks: ["otel"]
cache: true
cache_params:
type: redis
ttl: 600
supported_call_types: ["acompletion", "completion"]
model_group_settings:
forward_client_headers_to_llm_api:
- fake-openai-endpoint
success_callback: ["braintrust"]

View file

@ -43,10 +43,14 @@ from openai.types.responses.response import (
# Handle OpenAI SDK version compatibility for Text type
try:
from openai.types.responses.response_create_params import Text as ResponseText
from openai.types.responses.response_create_params import (
Text as ResponseText, # type: ignore
)
except (ImportError, AttributeError):
# Fall back to the concrete config type available in all SDK versions
from openai.types.responses.response_text_config_param import ResponseTextConfigParam as ResponseText
from openai.types.responses.response_text_config_param import (
ResponseTextConfigParam as ResponseText,
)
from openai.types.responses.response_create_params import (
Reasoning,

View file

@ -41,6 +41,7 @@ class PartType(TypedDict, total=False):
function_call: FunctionCall
function_response: FunctionResponse
thought: bool
thoughtSignature: str
class HttpxFunctionCall(TypedDict):
@ -72,6 +73,7 @@ class HttpxPartType(TypedDict, total=False):
executableCode: HttpxExecutableCode
codeExecutionResult: HttpxCodeExecutionResult
thought: bool
thoughtSignature: str
class HttpxContentType(TypedDict, total=False):
@ -245,10 +247,11 @@ class UsageMetadata(TypedDict, total=False):
class TokenCountDetailsResponse(TypedDict):
"""
Response structure for token count details with modality breakdown.
Example:
{'totalTokens': 12, 'promptTokensDetails': [{'modality': 'TEXT', 'tokenCount': 12}]}
"""
totalTokens: int
promptTokensDetails: List[PromptTokensDetails]

View file

@ -436,7 +436,10 @@ def test_gemini_with_empty_function_call_arguments():
async def test_claude_tool_use_with_gemini():
response = await litellm.anthropic.messages.acreate(
messages=[
{"role": "user", "content": "Hello, can you tell me the weather in Boston. Please respond with a tool call?"}
{
"role": "user",
"content": "Hello, can you tell me the weather in Boston. Please respond with a tool call?",
}
],
model="gemini/gemini-2.5-flash",
stream=True,
@ -578,11 +581,17 @@ def test_gemini_tool_use():
assert stop_reason is not None
assert stop_reason == "tool_calls"
@pytest.mark.asyncio
async def test_gemini_image_generation_async():
litellm._turn_on_debug()
response = await litellm.acompletion(
messages=[{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}],
messages=[
{
"role": "user",
"content": "Generate an image of a banana wearing a costume that says LiteLLM",
}
],
model="gemini/gemini-2.5-flash-image-preview",
)
@ -597,12 +606,16 @@ async def test_gemini_image_generation_async():
assert IMAGE_URL["url"].startswith("data:image/png;base64,")
@pytest.mark.asyncio
async def test_gemini_image_generation_async_stream():
#litellm._turn_on_debug()
# litellm._turn_on_debug()
response = await litellm.acompletion(
messages=[{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}],
messages=[
{
"role": "user",
"content": "Generate an image of a banana wearing a costume that says LiteLLM",
}
],
model="gemini/gemini-2.5-flash-image-preview",
stream=True,
)
@ -611,35 +624,144 @@ async def test_gemini_image_generation_async_stream():
model_response_image = None
async for chunk in response:
print("CHUNK: ", chunk)
if hasattr(chunk.choices[0].delta, "image") and chunk.choices[0].delta.image is not None:
if (
hasattr(chunk.choices[0].delta, "image")
and chunk.choices[0].delta.image is not None
):
model_response_image = chunk.choices[0].delta.image
print("MODEL_RESPONSE_IMAGE: ", model_response_image)
assert model_response_image is not None
assert model_response_image["url"].startswith("data:image/png;base64,")
break
#########################################################
# Important: Validate we did get an image in the response
#########################################################
assert model_response_image is not None
assert model_response_image["url"].startswith("data:image/png;base64,")
def test_system_message_with_no_user_message():
"""
Test that the system message is translated correctly for non-OpenAI providers.
"""
messages = [
{
"role": "system",
"content": "Be a good bot!",
},
]
"""
Test that the system message is translated correctly for non-OpenAI providers.
"""
messages = [
{
"role": "system",
"content": "Be a good bot!",
},
]
response = litellm.completion(
model="gemini/gemini-2.5-flash",
messages=messages,
response = litellm.completion(
model="gemini/gemini-2.5-flash",
messages=messages,
)
assert response is not None
assert response.choices[0].message.content is not None
def get_current_weather(location, unit="fahrenheit"):
"""Get the current weather in a given location"""
if "tokyo" in location.lower():
return json.dumps({"location": "Tokyo", "temperature": "10", "unit": "celsius"})
elif "san francisco" in location.lower():
return json.dumps(
{"location": "San Francisco", "temperature": "72", "unit": "fahrenheit"}
)
assert response is not None
elif "paris" in location.lower():
return json.dumps({"location": "Paris", "temperature": "22", "unit": "celsius"})
else:
return json.dumps({"location": location, "temperature": "unknown"})
assert response.choices[0].message.content is not None
def test_gemini_with_thinking():
from litellm import completion
litellm._turn_on_debug()
litellm.modify_params = True
model = "gemini/gemini-2.5-flash"
messages = [
{
"role": "user",
"content": "What's the weather like in San Francisco, Tokyo, and Paris? - give me 3 responses",
}
]
tools = [
{
"type": "function",
"function": {
"name": "get_current_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state",
},
"unit": {
"type": "string",
"enum": ["celsius", "fahrenheit"],
},
},
"required": ["location"],
},
},
}
]
response = litellm.completion(
model=model,
messages=messages,
tools=tools,
tool_choice="auto", # auto is default, but we'll be explicit
reasoning_effort="low",
)
print("Response\n", response)
response_message = response.choices[0].message
tool_calls = response_message.tool_calls
print("Expecting there to be 3 tool calls")
assert len(tool_calls) > 0 # this has to call the function for SF, Tokyo and paris
# Step 2: check if the model wanted to call a function
print(f"tool_calls: {tool_calls}")
if tool_calls:
# Step 3: call the function
# Note: the JSON response may not always be valid; be sure to handle errors
available_functions = {
"get_current_weather": get_current_weather,
} # only one function in this example, but you can have multiple
messages.append(response_message) # extend conversation with assistant's reply
print("Response message\n", response_message)
# Step 4: send the info for each function call and function response to the model
for tool_call in tool_calls:
function_name = tool_call.function.name
if function_name not in available_functions:
# the model called a function that does not exist in available_functions - don't try calling anything
return
function_to_call = available_functions[function_name]
function_args = json.loads(tool_call.function.arguments)
function_response = function_to_call(
location=function_args.get("location"),
unit=function_args.get("unit"),
)
messages.append(
{
"tool_call_id": tool_call.id,
"role": "tool",
"name": function_name,
"content": function_response,
}
) # extend conversation with function response
print(f"messages: {messages}")
second_response = litellm.completion(
model=model,
messages=messages,
seed=22,
reasoning_effort="low",
tools=tools,
drop_params=True,
) # get a new response from the model where it can see the function response
print("second response\n", second_response)

View file

@ -11,7 +11,7 @@ from litellm.integrations.braintrust_logging import BraintrustLogger
class TestBraintrustSpanName(unittest.TestCase):
"""Test custom span_name functionality in Braintrust logging."""
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
@patch("litellm.integrations.braintrust_logging.HTTPHandler")
def test_default_span_name(self, MockHTTPHandler):
"""Test that default span name is 'Chat Completion' when not provided."""
# Mock HTTP response
@ -22,39 +22,43 @@ class TestBraintrustSpanName(unittest.TestCase):
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop",
}
],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
"response_cost": 0.001,
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Chat Completion')
json_data = call_args.kwargs["json"]
self.assertEqual(
json_data["events"][0]["span_attributes"]["name"], "Chat Completion"
)
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
@patch("litellm.integrations.braintrust_logging.HTTPHandler")
def test_custom_span_name(self, MockHTTPHandler):
"""Test that custom span name is used when provided in metadata."""
# Mock HTTP response
@ -65,39 +69,43 @@ class TestBraintrustSpanName(unittest.TestCase):
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop",
}
],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {"span_name": "Custom Operation"}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
"response_cost": 0.001,
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Custom Operation')
json_data = call_args.kwargs["json"]
self.assertEqual(
json_data["events"][0]["span_attributes"]["name"], "Custom Operation"
)
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
@patch("litellm.integrations.braintrust_logging.HTTPHandler")
def test_span_name_with_other_metadata(self, MockHTTPHandler):
"""Test that span_name works alongside other metadata fields."""
# Mock HTTP response
@ -108,21 +116,23 @@ class TestBraintrustSpanName(unittest.TestCase):
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop",
}
],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
@ -132,34 +142,40 @@ class TestBraintrustSpanName(unittest.TestCase):
"project_id": "custom-project",
"user_id": "user123",
"session_id": "session456",
"environment": "production"
"environment": "production",
}
},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
"response_cost": 0.001,
"standard_logging_object": {
"user_id": "user123",
},
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
# Check span name
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Multi Metadata Test')
# Check that other metadata is preserved (except for filtered keys)
event_metadata = json_data['events'][0]['metadata']
self.assertEqual(event_metadata['user_id'], 'user123')
self.assertEqual(event_metadata['session_id'], 'session456')
self.assertEqual(event_metadata['environment'], 'production')
# Span name should be in span_attributes, not in metadata
self.assertIn('span_name', event_metadata) # span_name is also kept in metadata
json_data = call_args.kwargs["json"]
@patch('litellm.integrations.braintrust_logging.get_async_httpx_client')
# Check span name
self.assertEqual(
json_data["events"][0]["span_attributes"]["name"], "Multi Metadata Test"
)
# Check that other metadata is preserved (except for filtered keys)
event_metadata = json_data["events"][0]["metadata"]
print(event_metadata)
self.assertEqual(event_metadata["user_id"], "user123")
self.assertEqual(event_metadata["session_id"], "session456")
self.assertEqual(event_metadata["environment"], "production")
# Span name should be in span_attributes, not in metadata
self.assertIn("span_name", event_metadata) # span_name is also kept in metadata
@patch("litellm.integrations.braintrust_logging.get_async_httpx_client")
async def test_async_custom_span_name(self, mock_get_http_handler):
"""Test async logging with custom span name."""
# Mock async HTTP response
@ -170,38 +186,44 @@ class TestBraintrustSpanName(unittest.TestCase):
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
choices=[
{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop",
}
],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {"span_name": "Async Custom Operation"}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
"response_cost": 0.001,
}
# Execute
await logger.async_log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
await logger.async_log_success_event(
kwargs, response_obj, datetime.now(), datetime.now()
)
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Async Custom Operation')
json_data = call_args.kwargs["json"]
self.assertEqual(
json_data["events"][0]["span_attributes"]["name"], "Async Custom Operation"
)
if __name__ == "__main__":
unittest.main()
unittest.main()

View file

@ -0,0 +1,75 @@
from litellm.llms.vertex_ai.gemini.transformation import check_if_part_exists_in_parts
def test_check_if_part_exists_in_parts():
parts = [
{"text": "Hello", "thought": True},
{"text": "World", "thought": False},
]
part = {"text": "Hello", "thought": True}
new_part = {"text": "Hello World", "thought": True}
assert check_if_part_exists_in_parts(parts, part)
assert not check_if_part_exists_in_parts(parts, new_part, ["thought"])
assert check_if_part_exists_in_parts(parts, new_part, ["text"])
def test_check_if_part_exists_in_parts_camel_case_snake_case():
"""Test that function handles both camelCase and snake_case key variations"""
# Test snake_case to camelCase matching
parts_with_snake_case = [
{
"function_call": {
"name": "get_current_weather",
"args": {"location": "San Francisco, CA"},
}
},
{"text": "Some other content"},
]
part_with_camel_case = {
"functionCall": {
"name": "get_current_weather",
"args": {"location": "San Francisco, CA"},
}
}
# Should find match between function_call and functionCall
assert check_if_part_exists_in_parts(parts_with_snake_case, part_with_camel_case)
# Test camelCase to snake_case matching
parts_with_camel_case = [
{"functionCall": {"name": "calculate_sum", "args": {"a": 1, "b": 2}}}
]
part_with_snake_case = {
"function_call": {"name": "calculate_sum", "args": {"a": 1, "b": 2}}
}
# Should find match between functionCall and function_call
assert check_if_part_exists_in_parts(parts_with_camel_case, part_with_snake_case)
# Test no match when values differ
part_with_different_values = {
"function_call": {"name": "different_function", "args": {"x": 5}}
}
assert not check_if_part_exists_in_parts(
parts_with_snake_case, part_with_different_values
)
# Test multiple keys with mixed casing
parts_mixed = [
{
"function_call": {"name": "test"},
"thoughtSignature": "reasoning",
"text": "content",
}
]
part_mixed_casing = {
"functionCall": {"name": "test"},
"thought_signature": "reasoning",
"text": "content",
}
assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing)