diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index f3ca699546f..bb40b7665c1 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -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, ) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 33a480aa6bb..d8e36471ae5 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -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, ) ) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index ccaf28e5906..3d313456d19 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -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, diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 9376b28cbec..8e071d7759b 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -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(