mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(vertex_ai): narrow transformation scope and preserve schema defaults
- Restrict recursive camelCase transformation to 'contents', 'systemInstruction', and 'tools' to avoid breaking existing users. - Add logic to preserve user-supplied 'default' values in JSON Schema properties. - Fix 'labels' preservation to only trigger at top-level. - Ensure type-safe return for context caching with explicit cast.
This commit is contained in:
parent
94c2db996a
commit
acee5b1d5a
5 changed files with 83 additions and 5 deletions
|
|
@ -208,5 +208,11 @@ def transform_openai_messages_to_gemini_context_caching(
|
|||
{"system_instruction": data["system_instruction"]}, parent_key=None
|
||||
)["systemInstruction"]
|
||||
del data["system_instruction"]
|
||||
if "tools" in data:
|
||||
data["tools"] = _transform_part_to_httpx_format(
|
||||
{"tools": data["tools"]}, parent_key=None
|
||||
)["tools"]
|
||||
|
||||
return data
|
||||
from typing import cast
|
||||
|
||||
return cast(CachedContentRequestBody, data)
|
||||
|
|
|
|||
|
|
@ -267,9 +267,12 @@ def _transform_part_to_httpx_format(part: dict, parent_key: Optional[str] = None
|
|||
elif k == "properties":
|
||||
# 'properties' keys in a JSON Schema are user-defined
|
||||
should_preserve_keys = True
|
||||
elif k == "labels":
|
||||
elif k == "labels" and parent_key is None:
|
||||
# 'labels' is a top-level map for Vertex AI billing/tracking, keys are user-defined
|
||||
should_preserve_keys = True
|
||||
elif k == "default" and parent_key in ["properties", None]:
|
||||
# 'default' values in a JSON Schema property are user-defined and should NOT be transformed
|
||||
should_preserve_keys = True
|
||||
|
||||
if should_preserve_keys:
|
||||
if k == "properties":
|
||||
|
|
@ -836,6 +839,10 @@ def _transform_request_body( # noqa: PLR0915
|
|||
{"system_instruction": data["system_instruction"]}, parent_key=None
|
||||
)["systemInstruction"]
|
||||
del data["system_instruction"]
|
||||
if "tools" in data:
|
||||
data["tools"] = _transform_part_to_httpx_format(
|
||||
{"tools": data["tools"]}, parent_key=None
|
||||
)["tools"]
|
||||
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -364,3 +364,64 @@ def test_vertex_transformation_field_casing():
|
|||
transformed = _transform_part_to_httpx_format(part)
|
||||
assert "labels" in transformed
|
||||
assert "response" in transformed["labels"]
|
||||
|
||||
# 6. Test default preservation in schema
|
||||
part = {"properties": {"my_field": {"type": "object", "default": {"snake_case_key": 1}}}}
|
||||
transformed = _transform_part_to_httpx_format(part, parent_key="properties")
|
||||
assert "properties" in transformed
|
||||
assert "my_field" in transformed["properties"]
|
||||
assert "default" in transformed["properties"]["my_field"]
|
||||
assert "snake_case_key" in transformed["properties"]["my_field"]["default"]
|
||||
|
||||
# 7. Test labels at non-top-level are camelCased
|
||||
part = {"some_inner_object": {"labels": {"user_key": "value"}}}
|
||||
transformed = _transform_part_to_httpx_format(part, parent_key="something")
|
||||
assert "someInnerObject" in transformed
|
||||
assert "labels" in transformed["someInnerObject"]
|
||||
assert "userKey" in transformed["someInnerObject"]["labels"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_request_body_tools_transformation():
|
||||
from litellm.llms.vertex_ai.gemini.transformation import _transform_request_body
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
tools = [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"name": "my_func",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"my_param": {"type": "string", "format": "email"}},
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
# Test for Vertex AI - should transform tools
|
||||
transformed = _transform_request_body(
|
||||
messages=messages,
|
||||
model="gemini-1.5-flash",
|
||||
optional_params={"tools": tools},
|
||||
custom_llm_provider="vertex_ai",
|
||||
litellm_params={},
|
||||
cached_content=None,
|
||||
)
|
||||
|
||||
assert "tools" in transformed
|
||||
assert "functionDeclarations" in transformed["tools"][0]
|
||||
|
||||
# Test for Gemini (AI Studio) - should NOT transform tools
|
||||
original = _transform_request_body(
|
||||
messages=messages,
|
||||
model="gemini-1.5-flash",
|
||||
optional_params={"tools": tools},
|
||||
custom_llm_provider=LlmProviders.GEMINI.value,
|
||||
litellm_params={},
|
||||
cached_content=None,
|
||||
)
|
||||
assert "tools" in original
|
||||
assert "function_declarations" in original["tools"][0]
|
||||
|
|
|
|||
|
|
@ -342,7 +342,9 @@ class TestTransformationWithTTL:
|
|||
|
||||
assert "ttl" in result
|
||||
assert result["ttl"] == "7200s"
|
||||
assert "system_instruction" in result
|
||||
|
||||
system_instruction_key = "systemInstruction" if custom_llm_provider in ["vertex_ai", "vertex_ai_beta"] else "system_instruction"
|
||||
assert system_instruction_key in result
|
||||
|
||||
if custom_llm_provider == "gemini":
|
||||
assert result["model"] == "models/gemini-2.5-pro"
|
||||
|
|
|
|||
|
|
@ -1555,5 +1555,7 @@ def test_system_prompt_only_adds_blank_user_message():
|
|||
#########################################################
|
||||
# system message was passed in
|
||||
#########################################################
|
||||
assert len(data["system_instruction"]) == 1
|
||||
assert data["system_instruction"]["parts"][0]["text"] == SYSTEM_INSTRUCTION
|
||||
assert (
|
||||
len(data["systemInstruction"]["parts"]) == 1
|
||||
) # was renamed to camelCase for Vertex AI
|
||||
assert data["systemInstruction"]["parts"][0]["text"] == SYSTEM_INSTRUCTION
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue