mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
add context caching support for vertex ai
This commit is contained in:
parent
d2f8787bdc
commit
acf8ec5b3c
4 changed files with 142 additions and 61 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue