mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
test: add tool_choice tests for vertex AI context caching
- Update existing assertions to expect tool_choice=None in get_cache_key calls - Verify tool_choice is included in cached content request body - Add sync/async tests with explicit tool_choice value: pop, cache key, request body Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
ceff1831b0
commit
cee00d61e5
1 changed files with 167 additions and 4 deletions
|
|
@ -201,9 +201,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_params == optional_params
|
||||
assert returned_cache == "existing_cache_name"
|
||||
|
||||
# Verify cache key was generated with tools and model
|
||||
# Verify cache key was generated with tools, tool_choice, and model
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages, tools=self.sample_tools, model="gemini-1.5-pro"
|
||||
messages=cached_messages, tools=self.sample_tools, tool_choice=None, model="gemini-1.5-pro"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -280,6 +280,8 @@ class TestContextCachingEndpoints:
|
|||
call_args = self.mock_client.post.call_args
|
||||
assert "tools" in call_args.kwargs["json"]
|
||||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
assert "tool_choice" in call_args.kwargs["json"]
|
||||
assert call_args.kwargs["json"]["tool_choice"] is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
|
|
@ -474,9 +476,9 @@ class TestContextCachingEndpoints:
|
|||
assert returned_params == optional_params
|
||||
assert returned_cache == "existing_cache_name"
|
||||
|
||||
# Verify cache key was generated with tools and model
|
||||
# Verify cache key was generated with tools, tool_choice, and model
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages, tools=self.sample_tools, model="gemini-1.5-pro"
|
||||
messages=cached_messages, tools=self.sample_tools, tool_choice=None, model="gemini-1.5-pro"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -558,6 +560,8 @@ class TestContextCachingEndpoints:
|
|||
call_args = self.mock_async_client.post.call_args
|
||||
assert "tools" in call_args.kwargs["json"]
|
||||
assert call_args.kwargs["json"]["tools"] == self.sample_tools
|
||||
assert "tool_choice" in call_args.kwargs["json"]
|
||||
assert call_args.kwargs["json"]["tool_choice"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -906,6 +910,165 @@ class TestContextCachingEndpoints:
|
|||
# Restart the patcher so teardown_method can stop it cleanly
|
||||
self._token_check_patcher.start()
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.transform_openai_messages_to_gemini_context_caching"
|
||||
)
|
||||
@patch.object(ContextCachingEndpoints, "check_cache")
|
||||
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
|
||||
def test_check_and_create_cache_with_tool_choice(
|
||||
self,
|
||||
mock_get_token_url,
|
||||
mock_check_cache,
|
||||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test that tool_choice is popped from optional_params, included in cache key, and set in request body"""
|
||||
cached_messages = [self.sample_messages[0]]
|
||||
non_cached_messages = [self.sample_messages[1]]
|
||||
mock_separate.return_value = (cached_messages, non_cached_messages)
|
||||
|
||||
mock_cache_obj.get_cache_key.return_value = "test_cache_key"
|
||||
mock_check_cache.return_value = None
|
||||
mock_get_token_url.return_value = ("token", "https://test-url.com")
|
||||
mock_transform.return_value = {"model": "gemini-1.5-pro", "contents": []}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"name": "new_cache_name",
|
||||
"model": "gemini-1.5-pro",
|
||||
}
|
||||
self.mock_client.post.return_value = mock_response
|
||||
|
||||
tool_choice_value = {"type": "function", "function": {"name": "get_weather"}}
|
||||
optional_params = {
|
||||
"tools": self.sample_tools.copy(),
|
||||
"tool_choice": tool_choice_value,
|
||||
}
|
||||
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
messages=self.sample_messages,
|
||||
optional_params=optional_params,
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model="gemini-1.5-pro",
|
||||
client=self.mock_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="test_location",
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# tool_choice should be popped from optional_params
|
||||
assert "tool_choice" not in optional_params
|
||||
assert "tools" not in optional_params
|
||||
|
||||
# tool_choice should be passed to get_cache_key
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages,
|
||||
tools=self.sample_tools,
|
||||
tool_choice=tool_choice_value,
|
||||
model="gemini-1.5-pro",
|
||||
)
|
||||
|
||||
# tool_choice should be in the request body
|
||||
call_args = self.mock_client.post.call_args
|
||||
assert call_args.kwargs["json"]["tool_choice"] == tool_choice_value
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai", "vertex_ai_beta"]
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.separate_cached_messages"
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.local_cache_obj"
|
||||
)
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.transform_openai_messages_to_gemini_context_caching"
|
||||
)
|
||||
@patch.object(ContextCachingEndpoints, "async_check_cache")
|
||||
@patch.object(ContextCachingEndpoints, "_get_token_and_url_context_caching")
|
||||
@patch(
|
||||
"litellm.llms.vertex_ai.context_caching.vertex_ai_context_caching.get_async_httpx_client"
|
||||
)
|
||||
async def test_async_check_and_create_cache_with_tool_choice(
|
||||
self,
|
||||
mock_get_client,
|
||||
mock_get_token_url,
|
||||
mock_async_check_cache,
|
||||
mock_transform,
|
||||
mock_cache_obj,
|
||||
mock_separate,
|
||||
custom_llm_provider,
|
||||
):
|
||||
"""Test that tool_choice is popped, included in cache key, and set in request body (async)"""
|
||||
cached_messages = [self.sample_messages[0]]
|
||||
non_cached_messages = [self.sample_messages[1]]
|
||||
mock_separate.return_value = (cached_messages, non_cached_messages)
|
||||
|
||||
mock_cache_obj.get_cache_key.return_value = "test_cache_key"
|
||||
mock_async_check_cache.return_value = None
|
||||
mock_get_token_url.return_value = ("token", "https://test-url.com")
|
||||
mock_transform.return_value = {"model": "gemini-1.5-pro", "contents": []}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"name": "new_cache_name",
|
||||
"model": "gemini-1.5-pro",
|
||||
}
|
||||
self.mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
tool_choice_value = {"type": "function", "function": {"name": "get_weather"}}
|
||||
optional_params = {
|
||||
"tools": self.sample_tools.copy(),
|
||||
"tool_choice": tool_choice_value,
|
||||
}
|
||||
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
messages=self.sample_messages,
|
||||
optional_params=optional_params,
|
||||
api_key="test_key",
|
||||
api_base=None,
|
||||
model="gemini-1.5-pro",
|
||||
client=self.mock_async_client,
|
||||
timeout=30.0,
|
||||
logging_obj=self.mock_logging,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="test_location",
|
||||
vertex_auth_header="vertext_test_token",
|
||||
)
|
||||
|
||||
# tool_choice should be popped from optional_params
|
||||
assert "tool_choice" not in optional_params
|
||||
assert "tools" not in optional_params
|
||||
|
||||
# tool_choice should be passed to get_cache_key
|
||||
mock_cache_obj.get_cache_key.assert_called_once_with(
|
||||
messages=cached_messages,
|
||||
tools=self.sample_tools,
|
||||
tool_choice=tool_choice_value,
|
||||
model="gemini-1.5-pro",
|
||||
)
|
||||
|
||||
# tool_choice should be in the request body
|
||||
call_args = self.mock_async_client.post.call_args
|
||||
assert call_args.kwargs["json"]["tool_choice"] == tool_choice_value
|
||||
|
||||
|
||||
class TestCheckCachePagination:
|
||||
"""Test pagination logic in check_cache and async_check_cache methods."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue