add context caching support for vertex ai

This commit is contained in:
Otavio Brito 2025-10-05 20:55:18 -03:00
parent d2f8787bdc
commit acf8ec5b3c
4 changed files with 142 additions and 61 deletions

View file

@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
"""
import re
from typing import List, Optional, Tuple
from typing import List, Optional, Tuple, Literal
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody
@ -155,13 +155,18 @@ def separate_cached_messages(
def transform_openai_messages_to_gemini_context_caching(
model: str, messages: List[AllMessageValues], cache_key: str
model: str,
messages: List[AllMessageValues],
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
cache_key: str,
vertex_project: Optional[str],
vertex_location: Optional[str],
) -> CachedContentRequestBody:
# Extract TTL from cached messages BEFORE system message transformation
ttl = extract_ttl_from_cached_messages(messages)
supports_system_message = get_supports_system_message(
model=model, custom_llm_provider="gemini"
model=model, custom_llm_provider=custom_llm_provider
)
transformed_system_messages, new_messages = _transform_system_message(
@ -170,9 +175,14 @@ def transform_openai_messages_to_gemini_context_caching(
transformed_messages = _gemini_convert_messages_with_history(messages=new_messages)
model_name = "models/{}".format(model)
if custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta":
model_name = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/{model_name}"
data = CachedContentRequestBody(
contents=transformed_messages,
model="models/{}".format(model),
model=model_name,
displayName=cache_key,
)

View file

@ -23,6 +23,9 @@ from .transformation import (
transform_openai_messages_to_gemini_context_caching,
)
from litellm.types.llms.vertex_ai import *
local_cache_obj = Cache(
type=LiteLLMCacheType.LOCAL
) # only used for calling 'get_cache_key' function
@ -41,8 +44,11 @@ class ContextCachingEndpoints(VertexBase):
def _get_token_and_url_context_caching(
self,
gemini_api_key: Optional[str],
custom_llm_provider: Literal["gemini"],
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
api_base: Optional[str],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> Tuple[Optional[str], str]:
"""
Internal function. Returns the token and url for the call.
@ -58,7 +64,10 @@ class ContextCachingEndpoints(VertexBase):
url = "https://generativelanguage.googleapis.com/v1beta/{}?key={}".format(
endpoint, gemini_api_key
)
elif custom_llm_provider == "vertex_ai":
auth_header = vertex_auth_header
endpoint = "cachedContents"
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/{endpoint}"
else:
raise NotImplementedError
@ -80,6 +89,10 @@ class ContextCachingEndpoints(VertexBase):
api_key: str,
api_base: Optional[str],
logging_obj: Logging,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> Optional[str]:
"""
Checks if content already cached.
@ -94,8 +107,11 @@ class ContextCachingEndpoints(VertexBase):
_, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key,
custom_llm_provider="gemini",
custom_llm_provider=custom_llm_provider,
api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
try:
## LOGGING
@ -145,6 +161,10 @@ class ContextCachingEndpoints(VertexBase):
api_key: str,
api_base: Optional[str],
logging_obj: Logging,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str]
) -> Optional[str]:
"""
Checks if content already cached.
@ -159,8 +179,11 @@ class ContextCachingEndpoints(VertexBase):
_, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key,
custom_llm_provider="gemini",
custom_llm_provider=custom_llm_provider,
api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
try:
## LOGGING
@ -212,6 +235,10 @@ class ContextCachingEndpoints(VertexBase):
client: Optional[HTTPHandler],
timeout: Optional[Union[float, httpx.Timeout]],
logging_obj: Logging,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
extra_headers: Optional[dict] = None,
cached_content: Optional[str] = None,
) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
@ -240,8 +267,11 @@ class ContextCachingEndpoints(VertexBase):
## AUTHORIZATION ##
token, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key,
custom_llm_provider="gemini",
custom_llm_provider=custom_llm_provider,
api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
headers = {
@ -273,6 +303,10 @@ class ContextCachingEndpoints(VertexBase):
api_key=api_key,
api_base=api_base,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
if google_cache_name:
return non_cached_messages, optional_params, google_cache_name
@ -280,7 +314,12 @@ class ContextCachingEndpoints(VertexBase):
## TRANSFORM REQUEST
cached_content_request_body = (
transform_openai_messages_to_gemini_context_caching(
model=model, messages=cached_messages, cache_key=generated_cache_key
model=model,
messages=cached_messages,
cache_key=generated_cache_key,
custom_llm_provider=custom_llm_provider,
vertex_project=vertex_project,
vertex_location=vertex_location,
)
)
@ -328,6 +367,10 @@ class ContextCachingEndpoints(VertexBase):
client: Optional[AsyncHTTPHandler],
timeout: Optional[Union[float, httpx.Timeout]],
logging_obj: Logging,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
extra_headers: Optional[dict] = None,
cached_content: Optional[str] = None,
) -> Tuple[List[AllMessageValues], dict, Optional[str]]:
@ -356,8 +399,11 @@ class ContextCachingEndpoints(VertexBase):
## AUTHORIZATION ##
token, url = self._get_token_and_url_context_caching(
gemini_api_key=api_key,
custom_llm_provider="gemini",
custom_llm_provider=custom_llm_provider,
api_base=api_base,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
headers = {
@ -386,6 +432,10 @@ class ContextCachingEndpoints(VertexBase):
api_key=api_key,
api_base=api_base,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header
)
if google_cache_name:
@ -394,7 +444,12 @@ class ContextCachingEndpoints(VertexBase):
## TRANSFORM REQUEST
cached_content_request_body = (
transform_openai_messages_to_gemini_context_caching(
model=model, messages=cached_messages, cache_key=generated_cache_key
model=model,
messages=cached_messages,
cache_key=generated_cache_key,
custom_llm_provider=custom_llm_provider,
vertex_project=vertex_project,
vertex_location=vertex_location,
)
)

View file

@ -514,34 +514,35 @@ def sync_transform_request_body(
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
litellm_params: dict,
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> RequestBody:
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
context_caching_endpoints = ContextCachingEndpoints()
if gemini_api_key is not None:
(
messages,
optional_params,
cached_content,
) = context_caching_endpoints.check_and_create_cache(
messages=messages,
optional_params=optional_params,
api_key=gemini_api_key,
api_base=api_base,
model=model,
client=client,
timeout=timeout,
extra_headers=extra_headers,
cached_content=optional_params.pop("cached_content", None),
logging_obj=logging_obj,
)
else: # [TODO] implement context caching for gemini as well
cached_content = None
if "cached_content" in optional_params:
cached_content = optional_params.pop("cached_content")
elif "cachedContent" in optional_params:
cached_content = optional_params.pop("cachedContent")
(
messages,
optional_params,
cached_content,
) = context_caching_endpoints.check_and_create_cache(
messages=messages,
optional_params=optional_params,
api_key=gemini_api_key or "dummy",
api_base=api_base,
model=model,
client=client,
timeout=timeout,
extra_headers=extra_headers,
cached_content=optional_params.pop("cached_content", None),
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
)
return _transform_request_body(
messages=messages,
@ -565,34 +566,34 @@ async def async_transform_request_body(
logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore
custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"],
litellm_params: dict,
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
) -> RequestBody:
from ..context_caching.vertex_ai_context_caching import ContextCachingEndpoints
context_caching_endpoints = ContextCachingEndpoints()
if gemini_api_key is not None:
(
messages,
optional_params,
cached_content,
) = await context_caching_endpoints.async_check_and_create_cache(
messages=messages,
optional_params=optional_params,
api_key=gemini_api_key,
api_base=api_base,
model=model,
client=client,
timeout=timeout,
extra_headers=extra_headers,
cached_content=optional_params.pop("cached_content", None),
logging_obj=logging_obj,
)
else: # [TODO] implement context caching for gemini as well
cached_content = None
if "cached_content" in optional_params:
cached_content = optional_params.pop("cached_content")
elif "cachedContent" in optional_params:
cached_content = optional_params.pop("cachedContent")
(
messages,
optional_params,
cached_content,
) = await context_caching_endpoints.async_check_and_create_cache(
messages=messages,
optional_params=optional_params,
api_key=gemini_api_key or "dummy",
api_base=api_base,
model=model,
client=client,
timeout=timeout,
extra_headers=extra_headers,
cached_content=optional_params.pop("cached_content", None),
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
)
return _transform_request_body(
messages=messages,

View file

@ -1698,7 +1698,6 @@ class VertexLLM(VertexBase):
gemini_api_key: Optional[str] = None,
extra_headers: Optional[dict] = None,
) -> CustomStreamWrapper:
request_body = await async_transform_request_body(**data) # type: ignore
should_use_v1beta1_features = self.is_using_v1beta1_features(
optional_params=optional_params
@ -1732,6 +1731,13 @@ class VertexLLM(VertexBase):
litellm_params=litellm_params,
)
request_body = await async_transform_request_body(
**data,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=auth_header) # type: ignore
## LOGGING
logging_obj.pre_call(
input=messages,
@ -1819,7 +1825,12 @@ class VertexLLM(VertexBase):
litellm_params=litellm_params,
)
request_body = await async_transform_request_body(**data) # type: ignore
request_body = await async_transform_request_body(
**data,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=auth_header) # type: ignore
_async_client_params = {}
if timeout:
_async_client_params["timeout"] = timeout
@ -1994,7 +2005,11 @@ class VertexLLM(VertexBase):
)
## TRANSFORMATION ##
data = sync_transform_request_body(**transform_request_params)
data = sync_transform_request_body(
**transform_request_params,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=auth_header)
## LOGGING
logging_obj.pre_call(