fix main.py

This commit is contained in:
Ishaan Jaffer 2025-10-09 18:40:34 -07:00
parent 5ce27bfad6
commit 8b17b5bfc2

View file

@ -85,6 +85,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig
from litellm.llms.bedrock.common_utils import BedrockModelInfo
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.common_utils import (
VertexAIModelRoute,
get_vertex_ai_model_route,
)
from litellm.realtime_api.main import _realtime_health_check
from litellm.secret_managers.main import get_secret_bool, get_secret_str
from litellm.types.router import GenericLiteLLMParams
@ -150,7 +154,6 @@ from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
from .llms.bedrock.embed.embedding import BedrockEmbedding
from .llms.bedrock.image.image_handler import BedrockImageGeneration
from .llms.bytez.chat.transformation import BytezChatConfig
from .llms.lemonade.chat.transformation import LemonadeChatConfig
from .llms.codestral.completion.handler import CodestralTextCompletion
from .llms.cohere.embed import handler as cohere_embed
from .llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler
@ -162,6 +165,7 @@ from .llms.gemini.common_utils import get_api_key_from_env
from .llms.groq.chat.handler import GroqChatCompletion
from .llms.heroku.chat.transformation import HerokuChatConfig
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
from .llms.lemonade.chat.transformation import LemonadeChatConfig
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
from .llms.oci.chat.transformation import OCIChatConfig
from .llms.ollama.completion import handler as ollama
@ -192,6 +196,7 @@ from .llms.vertex_ai.multimodal_embeddings.embedding_handler import (
from .llms.vertex_ai.text_to_speech.text_to_speech_handler import VertexTextToSpeechAPI
from .llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels
from .llms.vertex_ai.vertex_embeddings.embedding_handler import VertexEmbedding
from .llms.vertex_ai.vertex_gemma_models.main import VertexAIGemmaModels
from .llms.vertex_ai.vertex_model_garden.main import VertexAIModelGardenModels
from .llms.vllm.completion import handler as vllm_handler
from .llms.watsonx.chat.handler import WatsonXChatHandler
@ -255,6 +260,7 @@ vertex_multimodal_embedding = VertexMultimodalEmbedding()
vertex_image_generation = VertexImageGeneration()
google_batch_embeddings = GoogleBatchEmbeddings()
vertex_partner_models_chat_completion = VertexAIPartnerModels()
vertex_gemma_chat_completion = VertexAIGemmaModels()
vertex_model_garden_chat_completion = VertexAIModelGardenModels()
vertex_text_to_speech = VertexTextToSpeechAPI()
sagemaker_llm = SagemakerLLM()
@ -2875,7 +2881,7 @@ def completion( # type: ignore # noqa: PLR0915
extra_headers=headers,
)
elif custom_llm_provider == "vertex_ai":
elif custom_llm_provider == "vertex_ai":
vertex_ai_project = (
optional_params.pop("vertex_project", None)
or optional_params.pop("vertex_ai_project", None)
@ -2897,7 +2903,9 @@ def completion( # type: ignore # noqa: PLR0915
api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE")
new_params = safe_deep_copy(optional_params or {})
if vertex_partner_models_chat_completion.is_vertex_partner_model(model):
model_route = get_vertex_ai_model_route(model=model, litellm_params=litellm_params)
if model_route == VertexAIModelRoute.PARTNER_MODELS:
model_response = vertex_partner_models_chat_completion.completion(
model=model,
messages=messages,
@ -2918,10 +2926,7 @@ def completion( # type: ignore # noqa: PLR0915
timeout=timeout,
client=client,
)
elif "gemini" in model or (
litellm_params.get("base_model") is not None
and "gemini" in litellm_params["base_model"]
):
elif model_route == VertexAIModelRoute.GEMINI:
model_response = vertex_chat_completion.completion( # type: ignore
model=model,
messages=messages,
@ -2943,7 +2948,29 @@ def completion( # type: ignore # noqa: PLR0915
api_base=api_base,
extra_headers=headers,
)
elif "openai" in model:
elif model_route == VertexAIModelRoute.GEMMA:
# Vertex Gemma Models with custom prediction endpoint
model_response = vertex_gemma_chat_completion.completion(
model=model,
messages=messages,
model_response=model_response,
print_verbose=print_verbose,
optional_params=new_params,
litellm_params=litellm_params, # type: ignore
logger_fn=logger_fn,
encoding=encoding,
api_base=api_base,
vertex_location=vertex_ai_location,
vertex_project=vertex_ai_project,
vertex_credentials=vertex_credentials,
logging_obj=logging,
acompletion=acompletion,
headers=headers,
custom_prompt_dict=custom_prompt_dict,
timeout=timeout,
client=client,
)
elif model_route == VertexAIModelRoute.MODEL_GARDEN:
# Vertex Model Garden - OpenAI compatible models
model_response = vertex_model_garden_chat_completion.completion(
model=model,
@ -2965,7 +2992,7 @@ def completion( # type: ignore # noqa: PLR0915
timeout=timeout,
client=client,
)
else:
else: # VertexAIModelRoute.NON_GEMINI
model_response = vertex_ai_non_gemini.completion(
model=model,
messages=messages,