mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(vertex_ai): consider tools when validating context caching min tokens (#43319)
* fix(vertex_ai): consider tools when validating context caching min tokens Pass tools to is_prompt_caching_valid_prompt in both sync and async check_and_create_cache before popping them into the cachedContents request body. This allows agent-shaped requests with heavy tool schemas and small message histories to reach the minimum token threshold and benefit from prompt caching. Fixes #42804 * test(vertex_ai): avoid doubles on internal code and assert tools in cache payload
This commit is contained in:
parent
be35b22dfc
commit
9ba552d527
2 changed files with 125 additions and 0 deletions
|
|
@ -322,6 +322,7 @@ class ContextCachingEndpoints(VertexBase):
|
|||
if not is_prompt_caching_valid_prompt(
|
||||
model=model,
|
||||
messages=cached_messages,
|
||||
tools=optional_params.get("tools"),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
):
|
||||
verbose_logger.debug(
|
||||
|
|
@ -481,6 +482,7 @@ class ContextCachingEndpoints(VertexBase):
|
|||
if not is_prompt_caching_valid_prompt(
|
||||
model=model,
|
||||
messages=cached_messages,
|
||||
tools=optional_params.get("tools"),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
):
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -1503,6 +1503,129 @@ class TestContextCachingEndpoints:
|
|||
# Restart the patcher so teardown_method can stop it cleanly
|
||||
self._token_check_patcher.start()
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider", ["gemini", "vertex_ai"]
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_create_cache_considers_tools_for_min_tokens(
|
||||
self, custom_llm_provider, is_async
|
||||
):
|
||||
"""Test that context caching accounts for tools when validating minimum token count.
|
||||
|
||||
Fixes #42804: When messages alone are below the threshold, but tools push the total
|
||||
over the minimum token count, context caching must proceed and include tools.
|
||||
"""
|
||||
self._token_check_patcher.stop()
|
||||
|
||||
short_cached_messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "Short system instruction.",
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
]
|
||||
non_cached_messages = [
|
||||
{"role": "user", "content": "Hello world"},
|
||||
]
|
||||
all_messages = short_cached_messages + non_cached_messages
|
||||
|
||||
large_tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": f"synthetic_tool_{i}",
|
||||
"description": "A very descriptive explanation of a synthetic tool designed to add tokens to the prompt cache prefix " * 8,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
f"arg_{j}": {"type": "string", "description": "Argument description for caching verification " * 4}
|
||||
for j in range(10)
|
||||
},
|
||||
"required": [f"arg_{j}" for j in range(5)],
|
||||
},
|
||||
},
|
||||
}
|
||||
for i in range(12)
|
||||
]
|
||||
|
||||
optional_params = {
|
||||
**self.sample_optional_params,
|
||||
"tools": large_tools,
|
||||
}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"name": "cachedContents/test_cache_id",
|
||||
"model": "gemini-1.5-pro",
|
||||
}
|
||||
mock_response.status_code = 200
|
||||
self.mock_client.post.return_value = mock_response
|
||||
self.mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch.object(
|
||||
self.context_caching,
|
||||
"_get_token_and_url_context_caching",
|
||||
return_value=("fake_token", "https://fake.url/cachedContents"),
|
||||
), patch.object(
|
||||
self.context_caching,
|
||||
"check_cache",
|
||||
return_value=None,
|
||||
), patch.object(
|
||||
self.context_caching,
|
||||
"async_check_cache",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
):
|
||||
if is_async:
|
||||
result = await self.context_caching.async_check_and_create_cache(
|
||||
messages=all_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,
|
||||
cached_content=None,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="us-central1",
|
||||
vertex_auth_header="test_token",
|
||||
)
|
||||
else:
|
||||
result = self.context_caching.check_and_create_cache(
|
||||
messages=all_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,
|
||||
cached_content=None,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_project="test_project",
|
||||
vertex_location="us-central1",
|
||||
vertex_auth_header="test_token",
|
||||
)
|
||||
|
||||
messages, returned_params, returned_cache = result
|
||||
assert messages == non_cached_messages
|
||||
assert returned_cache == "cachedContents/test_cache_id"
|
||||
assert "tools" not in returned_params
|
||||
|
||||
post_mock = self.mock_async_client.post if is_async else self.mock_client.post
|
||||
post_mock.assert_called_once()
|
||||
call_kwargs = post_mock.call_args.kwargs
|
||||
assert call_kwargs["json"]["tools"] == large_tools
|
||||
assert call_kwargs["json"]["contents"] == [
|
||||
{"role": "user", "parts": [{"text": "Short system instruction."}]}
|
||||
]
|
||||
|
||||
self._token_check_patcher.start()
|
||||
|
||||
|
||||
def _model_turn_final_messages(self, final_cached_role):
|
||||
tool_call = {
|
||||
"id": "call_abc123",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue