mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-23 18:25:22 +00:00
[Feat] Support for Vertex AI Gemma Models on Custom Endpoints (#15397)
* TestVertexGemmaiCompletion * test vertex Gemma * fix file name * fix file naming * add VertexAIGemmaModels * add cost_router for vertexai * fix main.py * fix VertexGemmaConfig * fix Vertex AI Gemma-AI Models Handler * docs gemma * fix ids * test fix * ruff check fixes * docs fix * docs fix
This commit is contained in:
@@ -16,7 +16,6 @@ import TabItem from '@theme/TabItem';
|
||||
| AI21 (Jamba) | `vertex_ai/jamba-*` | [Vertex AI - AI21 Models](https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/ai21) |
|
||||
| Qwen | `vertex_ai/qwen/*` | [Vertex AI - Qwen Models](https://cloud.google.com/vertex-ai/generative-ai/docs/maas/qwen) |
|
||||
| OpenAI (GPT-OSS) | `vertex_ai/openai/gpt-oss-*` | [Vertex AI - GPT-OSS Models](https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/) |
|
||||
| Model Garden | `vertex_ai/openai/{MODEL_ID}` or `vertex_ai/{MODEL_ID}` | [Vertex Model Garden](https://cloud.google.com/model-garden?hl=en) |
|
||||
|
||||
## Vertex AI - Anthropic (Claude)
|
||||
|
||||
@@ -793,112 +792,3 @@ curl http://0.0.0.0:4000/v1/chat/completions \
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Model Garden
|
||||
|
||||
:::tip
|
||||
|
||||
All OpenAI compatible models from Vertex Model Garden are supported.
|
||||
|
||||
:::
|
||||
|
||||
#### Using Model Garden
|
||||
|
||||
**Almost all Vertex Model Garden models are OpenAI compatible.**
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="openai" label="OpenAI Compatible Models">
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/openai/{MODEL_ID}` |
|
||||
| Vertex Documentation | [Model Garden LiteLLM Inference](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/open-models/use-cases/model_garden_litellm_inference.ipynb), [Vertex Model Garden](https://cloud.google.com/model-garden?hl=en) |
|
||||
| Supported Operations | `/chat/completions`, `/embeddings` |
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/openai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="Proxy">
|
||||
|
||||
|
||||
**1. Add to config**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: llama3-1-8b-instruct
|
||||
litellm_params:
|
||||
model: vertex_ai/openai/5464397967697903616
|
||||
vertex_ai_project: "my-test-project"
|
||||
vertex_ai_location: "us-east-1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING at http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**3. Test it!**
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "llama3-1-8b-instruct", # 👈 the 'model_name' in config
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="non-openai" label="Non-OpenAI Compatible Models">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Vertex AI - Self Deployed Models
|
||||
|
||||
Deploy and use your own models on Vertex AI through Model Garden or custom endpoints.
|
||||
|
||||
## Model Garden
|
||||
|
||||
:::tip
|
||||
|
||||
All OpenAI compatible models from Vertex Model Garden are supported.
|
||||
|
||||
:::
|
||||
|
||||
### Using Model Garden
|
||||
|
||||
**Almost all Vertex Model Garden models are OpenAI compatible.**
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="openai" label="OpenAI Compatible Models">
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/openai/{MODEL_ID}` |
|
||||
| Vertex Documentation | [Model Garden LiteLLM Inference](https://github.com/GoogleCloudPlatform/generative-ai/blob/main/open-models/use-cases/model_garden_litellm_inference.ipynb), [Vertex Model Garden](https://cloud.google.com/model-garden?hl=en) |
|
||||
| Supported Operations | `/chat/completions`, `/embeddings` |
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/openai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="Proxy">
|
||||
|
||||
|
||||
**1. Add to config**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: llama3-1-8b-instruct
|
||||
litellm_params:
|
||||
model: vertex_ai/openai/5464397967697903616
|
||||
vertex_ai_project: "my-test-project"
|
||||
vertex_ai_location: "us-east-1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING at http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**3. Test it!**
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "llama3-1-8b-instruct", # 👈 the 'model_name' in config
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="non-openai" label="Non-OpenAI Compatible Models">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
## set ENV variables
|
||||
os.environ["VERTEXAI_PROJECT"] = "hardy-device-38811"
|
||||
os.environ["VERTEXAI_LOCATION"] = "us-central1"
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/<your-endpoint-id>",
|
||||
messages=[{ "content": "Hello, how are you?","role": "user"}]
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
|
||||
## Gemma Models (Custom Endpoints)
|
||||
|
||||
Deploy Gemma models on custom Vertex AI prediction endpoints with OpenAI-compatible format.
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `vertex_ai/gemma/{MODEL_NAME}` |
|
||||
| Vertex Documentation | [Vertex AI Prediction](https://cloud.google.com/vertex-ai/docs/predictions/get-predictions) |
|
||||
| Required Parameter | `api_base` - Full prediction endpoint URL |
|
||||
|
||||
### Usage
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="proxy" label="Proxy">
|
||||
|
||||
**1. Add to config.yaml**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemma-model
|
||||
litellm_params:
|
||||
model: vertex_ai/gemma/gemma-3-12b-it-1222199011122
|
||||
api_base: https://ENDPOINT.us-central1-PROJECT.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict
|
||||
vertex_project: "my-project-id"
|
||||
vertex_location: "us-central1"
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
**3. Test it**
|
||||
|
||||
```bash
|
||||
curl http://0.0.0.0:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gemma-model",
|
||||
"messages": [{"role": "user", "content": "What is machine learning?"}],
|
||||
"max_tokens": 100
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
response = completion(
|
||||
model="vertex_ai/gemma/gemma-3-12b-it-1222199011122",
|
||||
messages=[{"role": "user", "content": "What is machine learning?"}],
|
||||
api_base="https://ENDPOINT.us-central1-PROJECT.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
|
||||
vertex_project="my-project-id",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
@@ -1,4 +1,5 @@
|
||||
import re
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, get_type_hints
|
||||
|
||||
import httpx
|
||||
@@ -24,6 +25,68 @@ class VertexAIError(BaseLLMException):
|
||||
super().__init__(message=message, status_code=status_code, headers=headers)
|
||||
|
||||
|
||||
class VertexAIModelRoute(str, Enum):
|
||||
"""Enum for Vertex AI model routing"""
|
||||
PARTNER_MODELS = "partner_models"
|
||||
GEMINI = "gemini"
|
||||
GEMMA = "gemma"
|
||||
MODEL_GARDEN = "model_garden"
|
||||
NON_GEMINI = "non_gemini"
|
||||
|
||||
|
||||
def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None) -> VertexAIModelRoute:
|
||||
"""
|
||||
Determine which handler to use for a Vertex AI model based on the model name.
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "openai/gpt-oss-120b")
|
||||
litellm_params: Optional litellm parameters dict that may contain base_model for routing
|
||||
|
||||
Returns:
|
||||
VertexAIModelRoute: The route enum indicating which handler should be used
|
||||
|
||||
Examples:
|
||||
>>> get_vertex_ai_model_route("llama3-405b")
|
||||
VertexAIModelRoute.PARTNER_MODELS
|
||||
|
||||
>>> get_vertex_ai_model_route("gemini-pro")
|
||||
VertexAIModelRoute.GEMINI
|
||||
|
||||
>>> get_vertex_ai_model_route("gemma/gemma-3-12b-it")
|
||||
VertexAIModelRoute.GEMMA
|
||||
|
||||
>>> get_vertex_ai_model_route("openai/gpt-oss-120b")
|
||||
VertexAIModelRoute.MODEL_GARDEN
|
||||
"""
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
|
||||
VertexAIPartnerModels,
|
||||
)
|
||||
|
||||
# Check base_model in litellm_params for gemini override
|
||||
if litellm_params and litellm_params.get("base_model") is not None:
|
||||
if "gemini" in litellm_params["base_model"]:
|
||||
return VertexAIModelRoute.GEMINI
|
||||
|
||||
# Check for partner models (llama, mistral, claude, etc.)
|
||||
if VertexAIPartnerModels.is_vertex_partner_model(model=model):
|
||||
return VertexAIModelRoute.PARTNER_MODELS
|
||||
|
||||
# Check for gemma models
|
||||
if "gemma/" in model:
|
||||
return VertexAIModelRoute.GEMMA
|
||||
|
||||
# Check for model garden openai models
|
||||
if "openai" in model:
|
||||
return VertexAIModelRoute.MODEL_GARDEN
|
||||
|
||||
# Check for gemini models
|
||||
if "gemini" in model:
|
||||
return VertexAIModelRoute.GEMINI
|
||||
|
||||
# Default to non-gemini (legacy vertex models like chat-bison, text-bison, etc.)
|
||||
return VertexAIModelRoute.NON_GEMINI
|
||||
|
||||
|
||||
def get_supports_system_message(
|
||||
model: str, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"]
|
||||
) -> bool:
|
||||
|
||||
@@ -44,6 +44,7 @@ def cost_router(
|
||||
or "mistral" in model
|
||||
or "jamba" in model
|
||||
or "codestral" in model
|
||||
or "gemma" in model
|
||||
):
|
||||
return "cost_per_token"
|
||||
elif custom_llm_provider == "vertex_ai" and (
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
"""Vertex AI Gemma-AI Models Handler"""
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
"""
|
||||
API Handler for calling Vertex AI Gemma Models
|
||||
|
||||
These models use a custom prediction endpoint format that wraps messages in 'instances'
|
||||
with @requestFormat: "chatCompletions" and returns responses wrapped in 'predictions'.
|
||||
|
||||
Usage:
|
||||
|
||||
response = litellm.completion(
|
||||
model="vertex_ai/gemma/gemma-3-12b-it-1222199011122",
|
||||
messages=[{"role": "user", "content": "What is machine learning?"}],
|
||||
vertex_project="your-project-id",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
|
||||
Sent to this route when `model` is in the format `vertex_ai/gemma/{MODEL_NAME}`
|
||||
|
||||
The API expects a custom endpoint URL format:
|
||||
https://{ENDPOINT_NUMBER}.{location}-{REGION_NUMBER}.prediction.vertexai.goog/v1/projects/{PROJECT_ID}/locations/{location}/endpoints/{ENDPOINT_ID}:predict
|
||||
"""
|
||||
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import httpx # type: ignore
|
||||
|
||||
from litellm.utils import ModelResponse
|
||||
|
||||
from ..common_utils import VertexAIError
|
||||
from ..vertex_llm_base import VertexBase
|
||||
|
||||
|
||||
class VertexAIGemmaModels(VertexBase):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
encoding,
|
||||
logging_obj,
|
||||
api_base: Optional[str],
|
||||
optional_params: dict,
|
||||
custom_prompt_dict: dict,
|
||||
headers: Optional[dict],
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
litellm_params: dict,
|
||||
vertex_project=None,
|
||||
vertex_location=None,
|
||||
vertex_credentials=None,
|
||||
logger_fn=None,
|
||||
acompletion: bool = False,
|
||||
client=None,
|
||||
):
|
||||
"""
|
||||
Handles calling Vertex AI Gemma Models
|
||||
|
||||
Sent to this route when `model` is in the format `vertex_ai/gemma/{MODEL_NAME}`
|
||||
"""
|
||||
try:
|
||||
import vertexai
|
||||
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexLLM,
|
||||
)
|
||||
from litellm.llms.vertex_ai.vertex_gemma_models.transformation import (
|
||||
VertexGemmaConfig,
|
||||
)
|
||||
except Exception as e:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""",
|
||||
)
|
||||
|
||||
if not (
|
||||
hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models")
|
||||
):
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""",
|
||||
)
|
||||
try:
|
||||
model = model.replace("gemma/", "")
|
||||
vertex_httpx_logic = VertexLLM()
|
||||
|
||||
access_token, project_id = vertex_httpx_logic._ensure_access_token(
|
||||
credentials=vertex_credentials,
|
||||
project_id=vertex_project,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
gemma_transformation = VertexGemmaConfig()
|
||||
|
||||
## CONSTRUCT API BASE
|
||||
stream: bool = optional_params.get("stream", False) or False
|
||||
optional_params["stream"] = stream
|
||||
|
||||
# If api_base is not provided, it should be set as an environment variable
|
||||
# or passed explicitly because the endpoint URL is unique per deployment
|
||||
if api_base is None:
|
||||
raise VertexAIError(
|
||||
status_code=400,
|
||||
message="api_base is required for Vertex AI Gemma models. Please provide the full endpoint URL.",
|
||||
)
|
||||
|
||||
# Check if we need to append :predict
|
||||
if not api_base.endswith(":predict"):
|
||||
_, api_base = self._check_custom_proxy(
|
||||
api_base=api_base,
|
||||
custom_llm_provider="vertex_ai",
|
||||
gemini_api_key=None,
|
||||
endpoint="predict",
|
||||
stream=stream,
|
||||
auth_header=None,
|
||||
url=api_base,
|
||||
)
|
||||
# If api_base already ends with :predict, use it as-is
|
||||
|
||||
# Use the custom transformation handler for gemma models
|
||||
return gemma_transformation.completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
api_key=access_token,
|
||||
custom_prompt_dict=custom_prompt_dict,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
acompletion=acompletion,
|
||||
litellm_params=litellm_params,
|
||||
logger_fn=logger_fn,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
encoding=encoding,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
if hasattr(e, "status_code"):
|
||||
raise e
|
||||
raise VertexAIError(status_code=500, message=str(e))
|
||||
|
||||
@@ -0,0 +1,350 @@
|
||||
"""
|
||||
Transformation logic for Vertex AI Gemma Models
|
||||
|
||||
Handles the custom request/response format:
|
||||
- Request: Wraps messages in 'instances' with @requestFormat: "chatCompletions"
|
||||
- Response: Extracts data from 'predictions' wrapper
|
||||
|
||||
The actual message transformation reuses OpenAIGPTConfig since Gemma uses OpenAI-compatible format.
|
||||
"""
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, Union, cast
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
class VertexGemmaConfig(OpenAIGPTConfig):
|
||||
"""
|
||||
Configuration and transformation class for Vertex AI Gemma models
|
||||
|
||||
Extends OpenAIGPTConfig to wrap/unwrap the instances/predictions format
|
||||
used by Vertex AI's Gemma deployment endpoint.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform request to Vertex Gemma format.
|
||||
|
||||
Uses parent class to create OpenAI-compatible request, then wraps it
|
||||
in the Vertex Gemma instances format.
|
||||
"""
|
||||
# Get the base OpenAI request from parent class
|
||||
openai_request = super().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# Remove 'model' from the request as it's not needed in the instance
|
||||
openai_request.pop("model", None)
|
||||
|
||||
# Wrap in Vertex Gemma format
|
||||
return {
|
||||
"instances": [
|
||||
{
|
||||
"@requestFormat": "chatCompletions",
|
||||
**openai_request,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
async def async_transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Async version of transform_request.
|
||||
"""
|
||||
# Get the base OpenAI request from parent class
|
||||
openai_request = await super().async_transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# Remove 'model' from the request as it's not needed in the instance
|
||||
openai_request.pop("model", None)
|
||||
|
||||
# Wrap in Vertex Gemma format
|
||||
return {
|
||||
"instances": [
|
||||
{
|
||||
"@requestFormat": "chatCompletions",
|
||||
**openai_request,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
def _unwrap_predictions_response(
|
||||
self,
|
||||
response_json: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Unwrap the Vertex Gemma predictions format to OpenAI format.
|
||||
|
||||
Vertex Gemma wraps the OpenAI-compatible response in a 'predictions' field.
|
||||
This method extracts it so the parent class can process it normally.
|
||||
"""
|
||||
if "predictions" not in response_json:
|
||||
raise BaseLLMException(
|
||||
status_code=422,
|
||||
message="Invalid response format: missing 'predictions' field",
|
||||
)
|
||||
|
||||
return response_json["predictions"]
|
||||
|
||||
def completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
custom_prompt_dict: dict,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
optional_params: dict,
|
||||
acompletion: bool,
|
||||
litellm_params: dict,
|
||||
logger_fn: Optional[Callable] = None,
|
||||
client: Optional[httpx.Client] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
encoding=None,
|
||||
custom_llm_provider: str = "vertex_ai",
|
||||
):
|
||||
"""
|
||||
Make completion request to Vertex Gemma endpoint.
|
||||
Supports both sync and async requests.
|
||||
"""
|
||||
# Handle streaming
|
||||
stream = optional_params.get("stream", False)
|
||||
if stream:
|
||||
raise BaseLLMException(
|
||||
status_code=400,
|
||||
message="Streaming is not yet supported for Vertex AI Gemma models",
|
||||
)
|
||||
|
||||
if acompletion:
|
||||
return self._async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
timeout=timeout,
|
||||
encoding=encoding,
|
||||
)
|
||||
else:
|
||||
return self._sync_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model_response=model_response,
|
||||
print_verbose=print_verbose,
|
||||
logging_obj=logging_obj,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
timeout=timeout,
|
||||
encoding=encoding,
|
||||
)
|
||||
|
||||
def _sync_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding: Any,
|
||||
):
|
||||
"""Synchronous completion request"""
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
# Transform the request using parent class methods
|
||||
request_data = self.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params.copy(),
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Set up headers
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Log the request
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": request_data,
|
||||
"api_base": api_base,
|
||||
},
|
||||
)
|
||||
|
||||
# Make the HTTP request
|
||||
http_handler = HTTPHandler(concurrent_limit=1)
|
||||
response = http_handler.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BaseLLMException(
|
||||
status_code=response.status_code,
|
||||
message=f"Request failed: {response.text}",
|
||||
)
|
||||
|
||||
response_json = response.json()
|
||||
|
||||
# Unwrap predictions to get OpenAI-compatible response
|
||||
openai_response = self._unwrap_predictions_response(response_json)
|
||||
|
||||
# Use litellm's standard response converter
|
||||
model_response = cast(
|
||||
ModelResponse,
|
||||
convert_to_model_response_object(
|
||||
response_object=openai_response,
|
||||
model_response_object=model_response,
|
||||
_response_headers={},
|
||||
),
|
||||
)
|
||||
|
||||
# Ensure model is set correctly
|
||||
model_response.model = model
|
||||
|
||||
# Log the response
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=response_json,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
async def _async_completion(
|
||||
self,
|
||||
model: str,
|
||||
messages: list,
|
||||
api_base: str,
|
||||
api_key: str,
|
||||
model_response: ModelResponse,
|
||||
print_verbose: Callable,
|
||||
logging_obj: Any,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
timeout: Optional[Union[float, httpx.Timeout]],
|
||||
encoding: Any,
|
||||
):
|
||||
"""Asynchronous completion request"""
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
# Transform the request using parent class async methods
|
||||
request_data = await self.async_transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params.copy(),
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Set up headers
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
# Log the request
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": request_data,
|
||||
"api_base": api_base,
|
||||
},
|
||||
)
|
||||
|
||||
# Make the HTTP request
|
||||
http_handler = AsyncHTTPHandler(concurrent_limit=1)
|
||||
response = await http_handler.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise BaseLLMException(
|
||||
status_code=response.status_code,
|
||||
message=f"Request failed: {response.text}",
|
||||
)
|
||||
|
||||
response_json = response.json()
|
||||
|
||||
# Unwrap predictions to get OpenAI-compatible response
|
||||
openai_response = self._unwrap_predictions_response(response_json)
|
||||
|
||||
# Use litellm's standard response converter
|
||||
model_response = cast(
|
||||
ModelResponse,
|
||||
convert_to_model_response_object(
|
||||
response_object=openai_response,
|
||||
model_response_object=model_response,
|
||||
_response_headers={},
|
||||
),
|
||||
)
|
||||
|
||||
# Ensure model is set correctly
|
||||
model_response.model = model
|
||||
|
||||
# Log the response
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=response_json,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
+36
-9
@@ -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,
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
"""Tests for Vertex AI Gemma-AI models"""
|
||||
|
||||
+217
@@ -0,0 +1,217 @@
|
||||
"""
|
||||
Mocked tests for Vertex AI Gemma Models
|
||||
|
||||
Maps to: litellm/llms/vertex_ai/vertex_gemma_models/transformation.py
|
||||
"""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestVertexGemmaCompletion:
|
||||
"""Test completion flow for Vertex AI Gemma models using litellm.acompletion()"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_basic_request(self):
|
||||
"""
|
||||
Test litellm.acompletion() with Vertex AI Gemma model
|
||||
|
||||
Expected URL:
|
||||
https://32277599999999999.us-central1-10582012152.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict
|
||||
|
||||
Expected Request Body (sent to Vertex):
|
||||
{
|
||||
"instances": [
|
||||
{
|
||||
"@requestFormat": "chatCompletions",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What is machine learning?"
|
||||
}
|
||||
],
|
||||
"max_tokens": 100
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
Expected Vertex Response:
|
||||
{
|
||||
"deployedModelId": "1207280419999999999",
|
||||
"model": "projects/993702345710/locations/us-central1/models/gemma-3-12b-it-1222199011122",
|
||||
"modelDisplayName": "gemma-3-12b-it-1222199011122",
|
||||
"modelVersionId": "1",
|
||||
"predictions": {
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "length",
|
||||
"index": 0,
|
||||
"logprobs": null,
|
||||
"message": {
|
||||
"content": "Okay, let's break down machine learning...",
|
||||
"reasoning_content": null,
|
||||
"role": "assistant",
|
||||
"tool_calls": []
|
||||
},
|
||||
"stop_reason": null
|
||||
}
|
||||
],
|
||||
"created": 1759863903,
|
||||
"id": "chatcmpl-aaa4288f-2b8e-4bc0-8b14-4e444decd2c4",
|
||||
"model": "google/gemma-3-12b-it",
|
||||
"object": "chat.completion",
|
||||
"prompt_logprobs": null,
|
||||
"usage": {
|
||||
"completion_tokens": 100,
|
||||
"prompt_tokens": 14,
|
||||
"prompt_tokens_details": null,
|
||||
"total_tokens": 114
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Expected LiteLLM Response: Standard OpenAI format
|
||||
"""
|
||||
# Real Vertex response from user's spec
|
||||
mock_vertex_response = {
|
||||
"deployedModelId": "1207280419999999999",
|
||||
"model": "projects/993702345710/locations/us-central1/models/gemma-3-12b-it-1222199011122",
|
||||
"modelDisplayName": "gemma-3-12b-it-1222199011122",
|
||||
"modelVersionId": "1",
|
||||
"predictions": {
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "length",
|
||||
"index": 0,
|
||||
"logprobs": None,
|
||||
"message": {
|
||||
"content": "Okay, let's break down machine learning. Here's a comprehensive explanation, covering the core concepts, types, and some examples, tailored to different levels of understanding. I'll structure it into sections: **The Core Idea**, **Types of Machine Learning**, **How It Works (Simplified)**, **Examples**, and **Why It's Useful**.\n\n**1. The Core Idea: Learning from Data**\n\nAt its heart, machine learning (ML) is about enabling computers",
|
||||
"reasoning_content": None,
|
||||
"role": "assistant",
|
||||
"tool_calls": [],
|
||||
},
|
||||
"stop_reason": None,
|
||||
}
|
||||
],
|
||||
"created": 1759863903,
|
||||
"id": "chatcmpl-aaa4288f-2b8e-4bc0-8b14-4e444decd2c4",
|
||||
"model": "google/gemma-3-12b-it",
|
||||
"object": "chat.completion",
|
||||
"prompt_logprobs": None,
|
||||
"usage": {
|
||||
"completion_tokens": 100,
|
||||
"prompt_tokens": 14,
|
||||
"prompt_tokens_details": None,
|
||||
"total_tokens": 114,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
# Mock the async HTTP handler
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler"
|
||||
) as mock_http_handler:
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = mock_vertex_response
|
||||
mock_http_handler.return_value.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
# Call litellm.acompletion()
|
||||
response = await litellm.acompletion(
|
||||
model="vertex_ai/gemma/gemma-3-12b-it-1222199011122",
|
||||
messages=[{"role": "user", "content": "What is machine learning?"}],
|
||||
max_tokens=100,
|
||||
api_base="https://32277599999999999.us-central1-10582012152.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
|
||||
vertex_project="PROJECT_ID",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
|
||||
# Verify the request sent to Vertex
|
||||
call_args = mock_http_handler.return_value.post.call_args
|
||||
assert call_args is not None, "HTTP handler was not called"
|
||||
|
||||
request_data = call_args.kwargs["json"]
|
||||
request_url = call_args.kwargs["url"]
|
||||
|
||||
# Validate exact URL matches what we sent
|
||||
expected_url = "https://32277599999999999.us-central1-10582012152.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict"
|
||||
assert request_url == expected_url, f"Expected URL: {expected_url}\nActual URL: {request_url}"
|
||||
|
||||
# Validate Request Body matches expected format
|
||||
assert "instances" in request_data
|
||||
assert len(request_data["instances"]) == 1
|
||||
|
||||
outer_instance = request_data["instances"][0]
|
||||
assert outer_instance["@requestFormat"] == "chatCompletions"
|
||||
|
||||
# The actual instance with messages is nested inside
|
||||
assert "instances" in outer_instance
|
||||
inner_instance = outer_instance["instances"][0]
|
||||
assert inner_instance["@requestFormat"] == "chatCompletions"
|
||||
assert "messages" in inner_instance
|
||||
assert inner_instance["messages"][0]["role"] == "user"
|
||||
assert inner_instance["messages"][0]["content"] == "What is machine learning?"
|
||||
assert inner_instance["max_tokens"] == 100
|
||||
|
||||
# Validate LiteLLM Response (OpenAI format)
|
||||
assert response.id == "chatcmpl-aaa4288f-2b8e-4bc0-8b14-4e444decd2c4"
|
||||
assert response.object == "chat.completion"
|
||||
assert response.created == 1759863903
|
||||
# Model name has the gemma/ prefix stripped during processing
|
||||
assert response.model == "gemma-3-12b-it-1222199011122"
|
||||
|
||||
# Validate choices
|
||||
assert len(response.choices) == 1
|
||||
assert response.choices[0].index == 0
|
||||
assert response.choices[0].finish_reason == "length"
|
||||
assert response.choices[0].message.role == "assistant"
|
||||
assert "machine learning" in response.choices[0].message.content.lower()
|
||||
|
||||
# Validate usage
|
||||
assert response.usage.prompt_tokens == 14
|
||||
assert response.usage.completion_tokens == 100
|
||||
assert response.usage.total_tokens == 114
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_error_handling(self):
|
||||
"""
|
||||
Test litellm.acompletion() error handling when Vertex returns invalid response
|
||||
|
||||
Expected: Proper error handling when 'predictions' field is missing
|
||||
"""
|
||||
from litellm.exceptions import APIConnectionError
|
||||
|
||||
# Invalid response without predictions field
|
||||
invalid_response = {
|
||||
"deployedModelId": "123",
|
||||
"error": {
|
||||
"code": 400,
|
||||
"message": "Invalid request"
|
||||
}
|
||||
}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler"
|
||||
) as mock_http_handler:
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = invalid_response
|
||||
mock_http_handler.return_value.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
# Should raise exception (wrapped as APIConnectionError by LiteLLM)
|
||||
with pytest.raises(APIConnectionError) as exc_info:
|
||||
await litellm.acompletion(
|
||||
model="vertex_ai/gemma/gemma-3-12b-it",
|
||||
messages=[{"role": "user", "content": "Test"}],
|
||||
api_base="https://test.prediction.vertexai.goog/v1/projects/test/locations/us-central1/endpoints/123:predict",
|
||||
vertex_project="test-project",
|
||||
vertex_location="us-central1",
|
||||
)
|
||||
|
||||
# Verify the error message contains the original error
|
||||
assert "missing 'predictions' field" in str(exc_info.value)
|
||||
|
||||
Reference in New Issue
Block a user