mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
6d36219a4d
10 changed files with 504 additions and 262 deletions
|
|
@ -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"},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
from typing import Any, Union
|
||||
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue