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:
flex-myeonghyeon 2026-04-13 14:32:42 +09:00
parent ceff1831b0
commit cee00d61e5

View file

@ -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."""