feat(vertex_ai): support managed batches for custom_endpoint deployments

The Vertex batch API refuses both the v1beta1 BYOE endpoint field and
Model-Garden-sourced model resources, so a custom_endpoint batch job instead
runs batch-owned replicas of the deployment's own serving container:
unmanagedContainerModel with the containerSpec read verbatim from the
endpoint's deployed model, dedicatedResources copied from the online
deployment, and instanceConfig.keyField round-tripping the custom_id.

Uploads for these deployments stage rows as @requestFormat chatCompletions
instances (the vLLM adapter's native OpenAI mode) under a
custom-endpoints/<id> GCS path derived from the deployment api_base, and the
output read unwraps prediction.predictions (already a full OpenAI
chat.completion) into OpenAI batch output rows.
This commit is contained in:
mubashir1osmani 2026-09-09 16:44:04 -04:00
parent 697da95bf4
commit 7c35312657
8 changed files with 617 additions and 63 deletions

View file

@ -1,5 +1,5 @@
import json
from collections.abc import Coroutine, Sequence
from collections.abc import Coroutine, Mapping, Sequence
from typing import TYPE_CHECKING, Final, Protocol
from urllib.parse import urlparse
@ -17,12 +17,18 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.llms.vertex_ai.common_utils import VertexAIError, get_vertex_base_url
from litellm.llms.vertex_ai.common_utils import (
VERTEX_CUSTOM_ENDPOINT_KEY_FIELD,
VertexAIError,
get_vertex_base_url,
)
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.llms.vertex_ai.vertex_llm_base import _graft_default_vertex_path
from litellm.types.llms.openai import CreateBatchRequest
from litellm.types.llms.vertex_ai import (
VERTEX_CREDENTIALS_TYPES,
BatchDedicatedResources,
UnmanagedContainerModel,
VertexAIBatchPredictionJob,
VertexBatchPredictionResponse,
)
@ -58,8 +64,18 @@ class _FetchedResponseView(TypedDict):
response: ReadOnly[httpx.Response]
class _VertexOnlineDedicatedResources(TypedDict, total=False):
"""The dedicatedResources block on an online endpoint deployment; replica bounds are named
min/max there, unlike the batch job's starting/max."""
machineSpec: ReadOnly[Mapping[str, object]]
minReplicaCount: ReadOnly[int]
maxReplicaCount: ReadOnly[int]
class _VertexEndpointDeployedModel(TypedDict, total=False):
model: ReadOnly[str]
dedicatedResources: ReadOnly[_VertexOnlineDedicatedResources]
class _VertexEndpointResponse(TypedDict, total=False):
@ -72,6 +88,16 @@ class _VertexEndpointPayloadView(TypedDict):
payload: ReadOnly[_VertexEndpointResponse]
class _VertexModelResourceResponse(TypedDict, total=False):
containerSpec: ReadOnly[Mapping[str, object]]
class _VertexModelResourcePayloadView(TypedDict):
"""Holds one decoded GET models/<id> response so the payload reads back typed."""
payload: ReadOnly[_VertexModelResourceResponse]
def _gateway_api_base_or_none(api_base: str | None) -> str | None:
"""
A deployment `api_base` whose path names a concrete Vertex resource (contains `/projects/`,
@ -110,15 +136,6 @@ class VertexAIBatchPrediction(VertexLLM):
max_retries: int | None,
custom_endpoint: bool | None = None,
) -> LiteLLMBatch | Coroutine[object, object, LiteLLMBatch]:
if custom_endpoint:
raise VertexAIError(
status_code=400,
message=(
"Vertex AI batch prediction is not supported for `custom_endpoint` deployments. "
"The OpenAI-compatible custom endpoint path has no batch surface in LiteLLM; "
"use a publisher model or fine-tuned Gemini endpoint deployment instead."
),
)
sync_handler: Final = _get_httpx_client()
access_token, project_id = self._ensure_access_token(
@ -139,18 +156,37 @@ class VertexAIBatchPrediction(VertexLLM):
vertex_location=vertex_location or "us-central1",
)
)
if custom_endpoint and "/custom-endpoints/" not in transformed_batch_request.get("model", ""):
raise VertexAIError(
status_code=400,
message=(
"Vertex AI batch prediction on a `custom_endpoint` deployment requires an input "
"file uploaded through LiteLLM against that deployment (its file id carries a "
"custom-endpoints/<endpoint id> path); this input file targets a publisher or "
"fine-tuned Gemini model instead."
),
)
gateway_api_base: Final = _gateway_api_base_or_none(api_base)
vertex_batch_request: Final = self._resolve_fine_tuned_endpoint_model(
resolved_batch_request: Final = self._resolve_fine_tuned_endpoint_model(
vertex_batch_request=transformed_batch_request,
headers=headers,
sync_handler=sync_handler,
api_base=gateway_api_base,
vertex_location=vertex_location or "us-central1",
)
vertex_batch_request: Final = self._resolve_custom_endpoint_container(
vertex_batch_request=resolved_batch_request,
headers=headers,
sync_handler=sync_handler,
api_base=gateway_api_base,
vertex_location=vertex_location or "us-central1",
)
is_unmanaged_container_job: Final = "unmanagedContainerModel" in vertex_batch_request
default_api_base: Final = self.create_vertex_batch_url(
vertex_location=vertex_location or "us-central1",
vertex_project=vertex_project or project_id,
vertex_api_version="v1beta1" if is_unmanaged_container_job else "v1",
)
if len(default_api_base.split(":")) > 1:
@ -169,7 +205,7 @@ class VertexAIBatchPrediction(VertexLLM):
model=None,
vertex_project=vertex_project or project_id,
vertex_location=vertex_location or "us-central1",
vertex_api_version="v1",
vertex_api_version="v1beta1" if is_unmanaged_container_job else "v1",
)
if _is_async is True:
@ -263,6 +299,113 @@ class VertexAIBatchPrediction(VertexLLM):
resolved_request: Final[VertexAIBatchPredictionJob] = {**vertex_batch_request, "model": deployed_model}
return resolved_request
def _resolve_custom_endpoint_container(
self,
vertex_batch_request: VertexAIBatchPredictionJob,
headers: dict[str, str], # mutable-ok: HTTPHandler.get only accepts dict headers
sync_handler: HTTPHandler,
api_base: str | None,
vertex_location: str,
) -> VertexAIBatchPredictionJob:
"""
A `custom_endpoint` deployment serves an OpenAI-compatible container on a Vertex endpoint.
The batch API accepts neither that endpoint (the v1beta1 BYOE `endpoint` field is refused
with "specify model or unmanaged_container_model") nor its Model-Garden-sourced model
resource ("Unknown ModelSource source_type: MODEL_GARDEN"), so the job instead runs
batch-owned replicas of the same container: `unmanagedContainerModel` with the
containerSpec read verbatim from the endpoint's deployed model (a hand-built spec loses
model-source args and crash-loops) plus `dedicatedResources` copied from the endpoint's
own deployment.
"""
model: Final = vertex_batch_request.get("model", "")
if "/custom-endpoints/" not in model:
return vertex_batch_request
endpoint_resource: Final = model.replace("/custom-endpoints/", "/endpoints/")
endpoint_url: Final = self._build_endpoint_resolution_url(
api_base=api_base,
model=endpoint_resource,
vertex_location=vertex_location,
)
endpoint_fetched: Final[_FetchedResponseView] = {
"response": safe_get(sync_handler, endpoint_url, headers=headers)
}
endpoint_response: Final = endpoint_fetched["response"]
if endpoint_response.status_code != 200:
raise VertexAIError(
status_code=endpoint_response.status_code,
message=f"Failed to resolve custom Vertex endpoint '{endpoint_resource}': {endpoint_response.text}",
)
endpoint_view: Final[_VertexEndpointPayloadView] = {"payload": endpoint_response.json()}
deployed_models: Final = endpoint_view["payload"].get("deployedModels") or ()
deployed: Final = deployed_models[0] if deployed_models else _VertexEndpointDeployedModel()
deployed_model_resource: Final = deployed.get("model", "")
if not deployed_model_resource:
raise VertexAIError(
status_code=400,
message=(
f"Vertex endpoint '{endpoint_resource}' has no deployed model, so there is no "
"serving container to run batch predictions with"
),
)
model_url: Final = self._build_endpoint_resolution_url(
api_base=api_base,
model=deployed_model_resource,
vertex_location=vertex_location,
)
model_fetched: Final[_FetchedResponseView] = {
"response": safe_get(sync_handler, model_url, headers=headers)
}
model_response: Final = model_fetched["response"]
if model_response.status_code != 200:
raise VertexAIError(
status_code=model_response.status_code,
message=f"Failed to read model resource '{deployed_model_resource}': {model_response.text}",
)
model_view: Final[_VertexModelResourcePayloadView] = {"payload": model_response.json()}
container_spec: Final = model_view["payload"].get("containerSpec")
if not container_spec:
raise VertexAIError(
status_code=400,
message=(
f"Model resource '{deployed_model_resource}' carries no containerSpec, so its "
"serving container cannot be replicated for batch prediction"
),
)
online_resources: Final = deployed.get("dedicatedResources") or _VertexOnlineDedicatedResources()
machine_spec: Final = online_resources.get("machineSpec")
if not machine_spec:
raise VertexAIError(
status_code=400,
message=(
f"Vertex endpoint '{endpoint_resource}' exposes no dedicatedResources machine "
"spec to size the batch replicas from"
),
)
batch_resources: Final[BatchDedicatedResources] = {
"machineSpec": machine_spec,
"startingReplicaCount": online_resources.get("minReplicaCount", 1),
"maxReplicaCount": online_resources.get("maxReplicaCount", 1),
}
unmanaged: Final[UnmanagedContainerModel] = {"containerSpec": container_spec}
resolved: Final[VertexAIBatchPredictionJob] = {
"displayName": vertex_batch_request["displayName"],
"inputConfig": vertex_batch_request["inputConfig"],
"outputConfig": vertex_batch_request["outputConfig"],
"unmanagedContainerModel": unmanaged,
"dedicatedResources": batch_resources,
# keyField strips the custom_id tag from each instance before it reaches the
# container (vLLM rejects unknown fields) and echoes it back as `key` in the output
# row; it only takes effect alongside an explicit instanceType.
"instanceConfig": {
"instanceType": "object",
"keyField": VERTEX_CUSTOM_ENDPOINT_KEY_FIELD,
},
}
return resolved
async def _async_create_batch(
self,
vertex_batch_request: VertexAIBatchPredictionJob,
@ -298,11 +441,12 @@ class VertexAIBatchPrediction(VertexLLM):
self,
vertex_location: str,
vertex_project: str,
vertex_api_version: str = "v1",
) -> str:
"""Return the base url for the vertex garden models"""
# POST https://LOCATION-aiplatform.googleapis.com/v1/projects/PROJECT_ID/locations/LOCATION/batchPredictionJobs
base_url: Final = get_vertex_base_url(vertex_location)
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/batchPredictionJobs"
return f"{base_url}/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/batchPredictionJobs"
def retrieve_batch(
self,

View file

@ -129,13 +129,23 @@ class VertexAIBatchTransformation:
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str:
"""
Gets the output file id from the Vertex AI Batch response
Gemini jobs write `predictions.jsonl`; unmanaged-container jobs (custom_endpoint
deployments) write sharded `prediction.results-*` files into a directory Vertex names
`prediction-custom-unmanaged-model-<timestamp>`.
"""
output_info: Final = response.get("outputInfo") or OutputInfo()
output_file_id: str = output_info.get("gcsOutputDirectory", "")
output_directory: Final = output_info.get("gcsOutputDirectory", "")
results_filename: Final = (
"prediction.results-00000-of-00001"
if "prediction-custom-unmanaged-model" in output_directory
else "predictions.jsonl"
)
output_file_id: str = output_directory
if output_file_id:
output_file_id = output_file_id.rstrip("/") + "/predictions.jsonl"
if output_file_id and output_file_id != "/predictions.jsonl":
output_file_id = output_file_id.rstrip("/") + f"/{results_filename}"
if output_file_id and output_file_id != f"/{results_filename}":
return output_file_id
output_config: Final = response.get("outputConfig")
@ -209,17 +219,21 @@ class VertexAIBatchTransformation:
to its deployed tuned model (`projects/../locations/../models/<id>`) before sending the job.
"""
parsed_model: Final = cls._get_model_from_gcs_file(input_file_id)
if not parsed_model.startswith("endpoints/"):
if not parsed_model.startswith(("endpoints/", "custom-endpoints/")):
return parsed_model
if not vertex_project:
raise VertexAIError(
status_code=400,
message=(
f"Vertex AI batch jobs against a fine-tuned endpoint ('{parsed_model}') require "
f"Vertex AI batch jobs against an endpoint ('{parsed_model}') require "
"`vertex_project` to build the endpoint resource name"
),
)
return f"projects/{vertex_project}/locations/{vertex_location or 'us-central1'}/{parsed_model}"
location_segment: Final = f"projects/{vertex_project}/locations/{vertex_location or 'us-central1'}"
if parsed_model.startswith("custom-endpoints/"):
endpoint_id: Final = parsed_model.removeprefix("custom-endpoints/")
return f"{location_segment}/custom-endpoints/{endpoint_id}"
return f"{location_segment}/{parsed_model}"
@classmethod
def _get_model_from_gcs_file(cls, gcs_file_uri: str) -> str:
@ -260,12 +274,15 @@ class VertexAIBatchTransformation:
@classmethod
def _parse_model_from_gcs_file(cls, gcs_file_uri: str) -> str | None:
"""
Returns the `publishers/<publisher>/models/<model>` or `endpoints/<numeric id>` path from a
gcs uri, or None if the uri does not contain one.
Returns the `publishers/<publisher>/models/<model>`, `endpoints/<numeric id>`, or
`custom-endpoints/<numeric id>` path from a gcs uri, or None if the uri does not contain
one.
A publisher path wins over an `endpoints/` segment, and the last `endpoints/` occurrence is
used, so a user-configured bucket prefix that happens to contain `endpoints/<digits>` cannot
override the model path LiteLLM appended after it.
A publisher path wins over an endpoints segment, `custom-endpoints/` (a custom_endpoint
deployment's serving container run as an unmanaged-container batch) wins over a plain
`endpoints/` (a fine-tuned Gemini endpoint), and the last occurrence of each is used, so a
user-configured bucket prefix that happens to contain `endpoints/<digits>` cannot override
the model path LiteLLM appended after it.
"""
unquoted_uri: Final = unquote(gcs_file_uri)
_, separator, model_path = unquoted_uri.partition("publishers/")
@ -274,6 +291,11 @@ class VertexAIBatchTransformation:
if len(parts) >= 3 and parts[1] == "models" and parts[2]:
return f"publishers/{'/'.join(parts[:3])}"
_, custom_separator, custom_path = unquoted_uri.rpartition("custom-endpoints/")
custom_endpoint_id: Final = custom_path.split("/")[0] if custom_separator else ""
if custom_endpoint_id.isdigit():
return f"custom-endpoints/{custom_endpoint_id}"
_, endpoint_separator, endpoint_path = unquoted_uri.rpartition("endpoints/")
endpoint_id: Final = endpoint_path.split("/")[0] if endpoint_separator else ""
if endpoint_id.isdigit():

View file

@ -370,6 +370,9 @@ def get_vertex_base_model_name(model: str) -> str:
return model
VERTEX_CUSTOM_ENDPOINT_KEY_FIELD: Final = "litellm_custom_id"
def get_vertex_ai_fine_tuned_endpoint_id(model: str) -> str | None:
"""
Fine-tuned Gemini deployments are addressed by a numeric endpoint id,

View file

@ -7,7 +7,7 @@ import re
import time
from collections.abc import Callable, Iterable, Iterator, Mapping
from typing import Any, Final, TypedDict
from urllib.parse import quote, unquote
from urllib.parse import quote, unquote, urlparse
import httpx
from httpx import Headers, Response
@ -38,6 +38,7 @@ from litellm.llms.base_llm.files.transformation import (
LiteLLMLoggingObj,
)
from litellm.llms.vertex_ai.common_utils import (
VERTEX_CUSTOM_ENDPOINT_KEY_FIELD,
_convert_vertex_datetime_to_openai_datetime,
get_vertex_ai_fine_tuned_endpoint_id,
)
@ -645,6 +646,40 @@ def _parse_vertex_batch_output_row(line: str) -> _VertexBatchRow:
return row
def _is_custom_endpoint_batch_output_row(row: Mapping[str, object]) -> bool:
"""
An unmanaged-container (custom_endpoint) batch output row: Vertex echoes the instance (or the
`key` extracted from it) alongside a `prediction` wrapper, unlike Gemini rows which pair
`request`/`response`/`processed_time`.
"""
return "prediction" in row and ("instance" in row or "key" in row)
def _custom_endpoint_row_to_openai_batch_output_row(row: Mapping[str, object]) -> _OpenAIBatchOutputRow:
"""
Unwraps one unmanaged-container batch output row. The vLLM `@requestFormat: chatCompletions`
mode already produces a full OpenAI chat.completion under `prediction.predictions`, so the
transform is: recover the custom_id (the `key` field when `instanceConfig.keyField` was
honored, else the echoed instance's tag) and re-wrap in the OpenAI batch output envelope.
"""
key: Final = row.get("key")
instance: Final = row.get("instance")
instance_map: Final = instance if isinstance(instance, Mapping) else {}
custom_id: Final = str(key if key is not None else instance_map.get(VERTEX_CUSTOM_ENDPOINT_KEY_FIELD, ""))
prediction: Final = row.get("prediction")
prediction_map: Final = prediction if isinstance(prediction, Mapping) else {}
body: Final = prediction_map.get("predictions")
if not isinstance(body, Mapping):
error_text: Final = str(row.get("status") or prediction or "prediction carries no response body")
return _openai_batch_output_row(
custom_id=custom_id,
error_code="vertex_ai_error",
error_message=error_text,
)
return _openai_batch_output_row(custom_id=custom_id, body=body)
class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
"""Streams an OpenAI batch JSONL upload as Vertex-wrapped JSONL one row at a
time, so the transformed payload is never held in full.
@ -673,6 +708,59 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
return self._iter_vertex_jsonl_chunks()
VERTEX_CUSTOM_ENDPOINT_GCS_SEGMENT: Final = "custom-endpoints"
_VERTEX_CHAT_COMPLETIONS_REQUEST_FORMAT: Final = "chatCompletions"
def get_custom_endpoint_id_from_api_base(api_base: str | None) -> str | None:
"""
The Vertex endpoint a `custom_endpoint` deployment serves from is only recorded in its
api_base (`.../endpoints/<id>:rawPredict` or a dedicated-domain equivalent); batch jobs need
that id to read the endpoint's containerSpec, so extract it (verb suffix stripped).
"""
if not api_base:
return None
path_segments: Final = urlparse(api_base).path.split("/")
after_endpoints: Final = tuple(
segment for prior, segment in zip(path_segments, path_segments[1:]) if prior == "endpoints"
)
if not after_endpoints:
return None
return after_endpoints[-1].split(":")[0] or None
def _openai_batch_jsonl_entry_to_custom_endpoint_row(openai_entry: dict[str, Any]) -> Mapping[str, object]:
"""
One OpenAI batch JSONL line as the instance a vLLM-serving Vertex container consumes:
the OpenAI request body itself tagged `@requestFormat: chatCompletions` (the container
speaks OpenAI natively, so no Gemini translation), minus `model` (the batch replica
serves exactly one model) plus the custom_id under the job's `instanceConfig.keyField`.
"""
body: Final = openai_entry.get("body") or {}
row: Final = {k: v for k, v in body.items() if k != "model"}
return {
"@requestFormat": _VERTEX_CHAT_COMPLETIONS_REQUEST_FORMAT,
**row,
VERTEX_CUSTOM_ENDPOINT_KEY_FIELD: str(openai_entry.get("custom_id", "")),
}
class _OpenAIToCustomEndpointBatchUploadStream(BaseFileUploadStream):
"""Streams an OpenAI batch JSONL upload as `@requestFormat: chatCompletions` instances
for a custom_endpoint (OpenAI-compatible container) batch job, one row at a time."""
def __init__(self, openai_file_content: FileTypes) -> None:
self._openai_file_content = openai_file_content
def iter_bytes(self) -> Iterator[bytes]:
first = True
for entry in _iter_openai_jsonl_entries(self._openai_file_content):
row = _openai_batch_jsonl_entry_to_custom_endpoint_row(entry)
prefix = b"" if first else b"\n"
first = False
yield prefix + json.dumps(row).encode("utf-8")
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Config for VertexAI Files
@ -740,13 +828,25 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
object_name: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
def get_object_name(self, file_data: FileTypes, purpose: str, deployment_model: str | None = None) -> str:
def get_object_name(
self,
file_data: FileTypes,
purpose: str,
deployment_model: str | None = None,
custom_endpoint_id: str | None = None,
) -> str:
"""
Get the object name for the request.
Reads only the first JSONL entry (streamed) for batch files, so a large
upload is never materialized just to derive the GCS object name.
"""
if purpose == "batch" and custom_endpoint_id is not None:
safe_endpoint_id: Final = sanitize_cloud_object_path(custom_endpoint_id, fallback="endpoint")
return (
f"{VERTEX_AI_MANAGED_GCS_PREFIX}{VERTEX_CUSTOM_ENDPOINT_GCS_SEGMENT}/"
f"{safe_endpoint_id}/{uuid.uuid4()}"
)
if purpose == "batch":
## 1. If jsonl, derive the object name from the deployment model (or the first entry's)
first_entry: Final = next(_iter_openai_jsonl_entries(file_data), None)
@ -781,16 +881,6 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Get the complete url for the request
"""
if data.get("purpose") == "batch" and litellm_params.get("custom_endpoint"):
raise VertexAIError(
status_code=400,
message=(
"Vertex AI batch prediction is not supported for `custom_endpoint` deployments. "
"The OpenAI-compatible custom endpoint path has no batch surface in LiteLLM; "
"remove this deployment from the batch request (e.g. `target_model_names`) or "
"use a publisher model / fine-tuned Gemini endpoint instead."
),
)
bucket_name = self._get_configured_bucket_name(litellm_params)
bucket_name, object_prefix = split_configured_cloud_bucket_name(bucket_name)
file_data: Final = data.get("file")
@ -800,10 +890,29 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
if purpose is None:
raise ValueError("purpose is required")
configured_model: Final = litellm_params.get("model")
deployment_api_base: Final = litellm_params.get("api_base")
custom_endpoint_id: Final = (
get_custom_endpoint_id_from_api_base(
deployment_api_base if isinstance(deployment_api_base, str) else None
)
if litellm_params.get("custom_endpoint")
else None
)
if litellm_params.get("custom_endpoint") and purpose == "batch" and custom_endpoint_id is None:
raise VertexAIError(
status_code=400,
message=(
"Vertex AI batch prediction on a `custom_endpoint` deployment requires the "
"deployment's `api_base` to name its Vertex endpoint "
"(e.g. https://.../endpoints/<endpoint id>:rawPredict), so the batch job can "
"run replicas of that endpoint's serving container."
),
)
object_name = self.get_object_name(
file_data,
purpose,
deployment_model=configured_model if isinstance(configured_model, str) else None,
custom_endpoint_id=custom_endpoint_id,
)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
@ -867,12 +976,17 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
create_file_data=create_file_data,
content_type=content_type,
):
body_stream: Final[BaseFileUploadStream] = (
_OpenAIToCustomEndpointBatchUploadStream(file_data)
if litellm_params.get("custom_endpoint")
else _OpenAIToVertexBatchUploadStream(
file_data,
self._map_openai_to_vertex_params,
)
)
return {
"streaming_media_upload": StreamingMediaUploadConfig(
body_stream=_OpenAIToVertexBatchUploadStream(
file_data,
self._map_openai_to_vertex_params,
),
body_stream=body_stream,
content_type="application/json",
)
}
@ -1116,14 +1230,19 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
# first line is not valid UTF-8/JSON) raises and falls through to the
# passthrough below, leaving the content untouched.
first_row: Final = _parse_vertex_batch_output_row(first_line)
is_vertex_batch_output: Final = _is_vertex_embeddings_batch_output_row(first_row) or (
"request" in first_row
and "response" in first_row
and "processed_time" in first_row
and (
"candidates" in first_row.get("response", {})
or "promptFeedback" in first_row.get("response", {})
or bool(first_row.get("status"))
is_custom_endpoint_output: Final = _is_custom_endpoint_batch_output_row(first_row)
is_vertex_batch_output: Final = (
is_custom_endpoint_output
or _is_vertex_embeddings_batch_output_row(first_row)
or (
"request" in first_row
and "response" in first_row
and "processed_time" in first_row
and (
"candidates" in first_row.get("response", {})
or "promptFeedback" in first_row.get("response", {})
or bool(first_row.get("status"))
)
)
)
if not is_vertex_batch_output:
@ -1151,6 +1270,13 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
all_lines = itertools.chain((first_line,), lines)
if is_custom_endpoint_output:
return b"\n".join(
json.dumps(_custom_endpoint_row_to_openai_batch_output_row(json.loads(line))).encode("utf-8")
for line in all_lines
if line.strip()
)
# Embedding rows are grouped by `custom_id` rather than transformed one at a
# time, since an entry that asked for several embeddings comes back as
# several rows, in arbitrary order.

View file

@ -1,7 +1,9 @@
from collections.abc import Mapping
from enum import Enum
from typing import Any, Final, Literal, Protocol
from typing_extensions import (
ReadOnly,
Required,
TypedDict,
)
@ -678,11 +680,36 @@ class GcsBucketResponse(TypedDict):
timeFinalized: str
class VertexAIBatchPredictionJob(TypedDict):
displayName: str
model: str
inputConfig: InputConfig
outputConfig: OutputConfig
class BatchDedicatedResources(TypedDict, total=False):
"""Sizing for batch-owned replicas; machineSpec is copied verbatim from the online
deployment's dedicatedResources, hence the loose Mapping."""
machineSpec: ReadOnly[Mapping[str, object]]
startingReplicaCount: ReadOnly[int]
maxReplicaCount: ReadOnly[int]
class UnmanagedContainerModel(TypedDict, total=False):
"""The v1beta1 batch shape for running batch-owned replicas of a serving container. The
containerSpec is copied verbatim from the deployed model resource (hence the loose Mapping):
hand-building one loses model-source args/env and crash-loops the batch container."""
containerSpec: ReadOnly[Mapping[str, object]]
class BatchInstanceConfig(TypedDict, total=False):
instanceType: ReadOnly[str]
keyField: ReadOnly[str]
class VertexAIBatchPredictionJob(TypedDict, total=False):
displayName: ReadOnly[Required[str]]
model: ReadOnly[str]
unmanagedContainerModel: ReadOnly[UnmanagedContainerModel]
dedicatedResources: ReadOnly[BatchDedicatedResources]
instanceConfig: ReadOnly[BatchInstanceConfig]
inputConfig: ReadOnly[Required[InputConfig]]
outputConfig: ReadOnly[Required[OutputConfig]]
class VertexBatchPredictionResponse(TypedDict, total=False):

View file

@ -257,6 +257,123 @@ def test_create_batch_sync_resolves_fine_tuned_endpoint_to_tuned_model():
assert sent["model"] == TUNED_MODEL_RESOURCE
CUSTOM_ENDPOINT_ID = "4980511146650894336"
CUSTOM_ENDPOINT_CREATE_DATA = {
"input_file_id": (f"gs://bucket/litellm-vertex-files/custom-endpoints/{CUSTOM_ENDPOINT_ID}/file-uuid")
}
CONTAINER_MODEL_RESOURCE = f"projects/{PROJECT}/locations/{LOCATION}/models/google-gemma2-123"
CONTAINER_SPEC = {
"imageUri": "us-docker.pkg.dev/vertex-ai/pytorch-vllm-serve:x",
"args": ["python", "-m", "vllm.entrypoints.api_server"],
"predictRoute": "/generate",
"healthRoute": "/ping",
}
MACHINE_SPEC = {"machineType": "g2-standard-12", "acceleratorType": "NVIDIA_L4", "acceleratorCount": 1}
def _custom_endpoint_get_response() -> MagicMock:
resp = MagicMock()
resp.status_code = 200
resp.json.return_value = {
"name": f"projects/{PROJECT}/locations/{LOCATION}/endpoints/{CUSTOM_ENDPOINT_ID}",
"deployedModels": [
{
"model": CONTAINER_MODEL_RESOURCE,
"dedicatedResources": {"machineSpec": MACHINE_SPEC, "minReplicaCount": 1, "maxReplicaCount": 2},
}
],
}
return resp
def _container_model_get_response(container_spec: dict | None = CONTAINER_SPEC) -> MagicMock:
resp = MagicMock()
resp.status_code = 200
resp.json.return_value = (
{"name": CONTAINER_MODEL_RESOURCE, "containerSpec": container_spec}
if container_spec is not None
else {"name": CONTAINER_MODEL_RESOURCE}
)
return resp
def test_create_batch_sync_custom_endpoint_builds_unmanaged_container_job():
"""A custom_endpoint batch must run batch-owned replicas of the endpoint's own serving
container: the live API refuses both the v1beta1 BYOE `endpoint` field and Model-Garden model
resources, and a hand-built containerSpec crash-loops, so the job carries the deployed
model's containerSpec verbatim under `unmanagedContainerModel` on the v1beta1 route with the
custom_id extracted server-side via instanceConfig.keyField (LIT-7387)."""
h = _make_handler()
client = MagicMock()
client.post.return_value = _http_response()
with (
patch(f"{HMOD}._get_httpx_client", return_value=client),
patch(
f"{HMOD}.safe_get",
side_effect=[_custom_endpoint_get_response(), _container_model_get_response()],
) as safe_get,
):
out = h.create_batch(
_is_async=False,
create_batch_data=CUSTOM_ENDPOINT_CREATE_DATA,
api_base=None,
vertex_credentials=None,
vertex_project=PROJECT,
vertex_location=LOCATION,
timeout=600.0,
max_retries=None,
custom_endpoint=True,
)
assert isinstance(out, LiteLLMBatch)
endpoint_get_url = safe_get.call_args_list[0].args[1]
assert endpoint_get_url.endswith(f"/endpoints/{CUSTOM_ENDPOINT_ID}")
model_get_url = safe_get.call_args_list[1].args[1]
assert model_get_url.endswith(CONTAINER_MODEL_RESOURCE)
post_url = client.post.call_args.kwargs["url"]
assert "/v1beta1/" in post_url
sent = json.loads(client.post.call_args.kwargs["data"])
assert "model" not in sent
assert sent["unmanagedContainerModel"] == {"containerSpec": CONTAINER_SPEC}
assert sent["dedicatedResources"] == {
"machineSpec": MACHINE_SPEC,
"startingReplicaCount": 1,
"maxReplicaCount": 2,
}
assert sent["instanceConfig"] == {"instanceType": "object", "keyField": "litellm_custom_id"}
def test_create_batch_sync_custom_endpoint_without_container_spec_raises_400():
h = _make_handler()
client = MagicMock()
with (
patch(f"{HMOD}._get_httpx_client", return_value=client),
patch(
f"{HMOD}.safe_get",
side_effect=[_custom_endpoint_get_response(), _container_model_get_response(container_spec=None)],
),
):
with pytest.raises(VertexAIError) as exc_info:
h.create_batch(
_is_async=False,
create_batch_data=CUSTOM_ENDPOINT_CREATE_DATA,
api_base=None,
vertex_credentials=None,
vertex_project=PROJECT,
vertex_location=LOCATION,
timeout=600.0,
max_retries=None,
custom_endpoint=True,
)
assert exc_info.value.status_code == 400
assert "containerSpec" in str(exc_info.value)
client.post.assert_not_called()
def test_create_batch_sync_ignores_resource_shaped_api_base():
"""A deployment api_base like `.../endpoints/<id>:rawPredict` targets online inference, not
the Vertex API root; grafting batch urls onto it yields guaranteed 404s, so batch operations
@ -356,9 +473,9 @@ def test_create_batch_sync_endpoint_resolution_error_raises():
client.post.assert_not_called()
def test_create_batch_custom_endpoint_raises_400_without_io():
"""custom_endpoint deployments have no Vertex batch surface; creating a job would target a
nonexistent publisher model, so the handler must 400 before any auth or HTTP work (LIT-6899)."""
def test_create_batch_custom_endpoint_rejects_non_custom_endpoint_file():
"""A custom_endpoint batch create over a file staged for a publisher model would run the wrong
workload on batch replicas of the container; the handler must 400 before any HTTP work."""
h = _make_handler()
client = MagicMock()
@ -378,8 +495,8 @@ def test_create_batch_custom_endpoint_raises_400_without_io():
assert exc_info.value.status_code == 400
assert "custom_endpoint" in str(exc_info.value)
h._ensure_access_token.assert_not_called()
client.post.assert_not_called()
client.get.assert_not_called()
def test_create_batch_sync_endpoint_without_deployed_model_raises_400():

View file

@ -391,6 +391,29 @@ def test_get_bare_model_name_from_gcs_file_fine_tuned_endpoint():
assert T.get_bare_model_name_from_gcs_file(ENDPOINT_INPUT_FILE) == ENDPOINT_ID
CUSTOM_ENDPOINT_ID = "4980511146650894336"
CUSTOM_ENDPOINT_INPUT_FILE = (
f"gs://litellm-testing-bucket/litellm-vertex-files/custom-endpoints/{CUSTOM_ENDPOINT_ID}/"
"e9412502-2c91-42a6-8e61-f5c294cc0fc8"
)
def test_get_model_from_gcs_file_custom_endpoint():
"""`custom-endpoints/` contains `endpoints/` as a substring, so the custom marker must be
matched first or the id would be misread as a fine-tuned Gemini endpoint and the batch job
would target a nonexistent tuned model (LIT-7387)."""
assert T._get_model_from_gcs_file(CUSTOM_ENDPOINT_INPUT_FILE) == f"custom-endpoints/{CUSTOM_ENDPOINT_ID}"
def test_batch_job_model_custom_endpoint_builds_resource_path():
job = T.transform_openai_batch_request_to_vertex_ai_batch_request(
{"input_file_id": CUSTOM_ENDPOINT_INPUT_FILE},
vertex_project="my-project",
vertex_location="us-central1",
)
assert job["model"] == f"projects/my-project/locations/us-central1/custom-endpoints/{CUSTOM_ENDPOINT_ID}"
# =========================================================================== #
# is_unmanaged_gcs_batch_input_file_id
# =========================================================================== #

View file

@ -234,10 +234,41 @@ class TestBatchObjectNaming:
assert "9999999999999999999" not in object_name
CUSTOM_ENDPOINT_ID = "4980511146650894336"
CUSTOM_ENDPOINT_API_BASE = (
"https://us-central1-aiplatform.googleapis.com/v1/projects/my-project"
f"/locations/us-central1/endpoints/{CUSTOM_ENDPOINT_ID}:rawPredict"
)
class TestCustomEndpointBatchUpload:
def test_should_reject_batch_upload_for_custom_endpoint_deployment(self, config):
"""custom_endpoint deployments have no Vertex batch surface; the upload must 400 instead
of staging a file that can only produce a doomed batch job (LIT-6899)."""
def test_should_stage_batch_upload_under_custom_endpoints_path(self, config):
"""The GCS path is how the later batch create learns which serving container to
replicate, so a custom_endpoint upload must record the endpoint id from the api_base
under the custom-endpoints/ marker (LIT-7387)."""
url = config.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params={
"gcs_bucket_name": "my-bucket",
"custom_endpoint": True,
"api_base": CUSTOM_ENDPOINT_API_BASE,
"model": "vertex_ai/openai/gemma-2-2b-it",
},
data={
"file": ("batch.jsonl", b'{"body": {"model": "openai/gemma-2-2b-it"}}', "application/jsonl"),
"purpose": "batch",
},
)
assert url.startswith("https://storage.googleapis.com/")
object_name = parse_qs(urlparse(url).query)["name"][0]
assert object_name.startswith(f"litellm-vertex-files/custom-endpoints/{CUSTOM_ENDPOINT_ID}/")
def test_should_reject_batch_upload_when_api_base_names_no_endpoint(self, config):
"""Without an endpoint id in the api_base there is no container to run the batch with, so
the upload must fail with a clear 400 instead of staging a doomed file."""
from litellm.llms.vertex_ai.common_utils import VertexAIError
with pytest.raises(VertexAIError) as exc_info:
@ -246,14 +277,18 @@ class TestCustomEndpointBatchUpload:
api_key=None,
model="",
optional_params={},
litellm_params={"gcs_bucket_name": "my-bucket", "custom_endpoint": True},
litellm_params={
"gcs_bucket_name": "my-bucket",
"custom_endpoint": True,
"api_base": "https://my-gateway.internal/v1",
},
data={
"file": ("batch.jsonl", b'{"body": {"model": "openai/gemma-2-2b-it"}}', "application/jsonl"),
"purpose": "batch",
},
)
assert exc_info.value.status_code == 400
assert "custom_endpoint" in str(exc_info.value)
assert "api_base" in str(exc_info.value)
def test_should_allow_non_batch_upload_for_custom_endpoint_deployment(self, config):
url = config.get_complete_file_url(
@ -270,6 +305,63 @@ class TestCustomEndpointBatchUpload:
assert "/b/my-bucket/" in url
class TestCustomEndpointBatchRows:
def test_upload_stream_emits_chat_completions_instances(self):
"""Each OpenAI batch line must become a `@requestFormat: chatCompletions` instance the
vLLM container accepts natively, with `model` dropped (the batch replica serves exactly
one model) and the custom_id under the keyField name the batch job strips server-side."""
from litellm.llms.vertex_ai.files.transformation import (
_OpenAIToCustomEndpointBatchUploadStream,
)
openai_jsonl = (
b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions",'
b' "body": {"model": "gemma", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 5}}\n'
b'{"custom_id": "req-2", "method": "POST", "url": "/v1/chat/completions",'
b' "body": {"model": "gemma", "messages": [{"role": "user", "content": "yo"}]}}'
)
stream = _OpenAIToCustomEndpointBatchUploadStream(("batch.jsonl", openai_jsonl, "application/jsonl"))
rows = [json.loads(line) for line in b"".join(stream.iter_bytes()).split(b"\n")]
assert rows == [
{
"@requestFormat": "chatCompletions",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 5,
"litellm_custom_id": "req-1",
},
{
"@requestFormat": "chatCompletions",
"messages": [{"role": "user", "content": "yo"}],
"litellm_custom_id": "req-2",
},
]
def test_output_rows_unwrap_to_openai_batch_format(self, config):
"""An unmanaged-container output row already carries a full OpenAI chat.completion under
prediction.predictions; the transform must unwrap it and recover the custom_id from the
keyField echo, and a failed row must become an OpenAI batch error row."""
vertex_output = (
b'{"key": "req-1", "prediction": {"predictions": {"id": "chatcmpl-1", "object": "chat.completion",'
b' "model": "google/gemma2-2b-it", "choices": [{"index": 0, "message": {"role": "assistant",'
b' "content": "Hello"}, "finish_reason": "stop"}], "usage": {"prompt_tokens": 5,'
b' "completion_tokens": 2, "total_tokens": 7}}}}\n'
b'{"key": "req-2", "prediction": "Post request fails.", "status": "Post request fails."}'
)
transformed = config._try_transform_vertex_batch_output_to_openai(content=vertex_output)
rows = [json.loads(line) for line in transformed.split(b"\n")]
assert rows[0]["custom_id"] == "req-1"
assert rows[0]["error"] is None
assert rows[0]["response"]["status_code"] == 200
assert rows[0]["response"]["body"]["choices"][0]["message"]["content"] == "Hello"
assert rows[0]["response"]["body"]["usage"]["total_tokens"] == 7
assert rows[1]["custom_id"] == "req-2"
assert rows[1]["response"] is None
assert rows[1]["error"]["code"] == "vertex_ai_error"
assert "Post request fails." in rows[1]["error"]["message"]
class TestTransformRetrieveFile:
def test_should_build_correct_gcs_metadata_url(self, config):
file_id = "gs://my-bucket/litellm-vertex-files/path/to/file.jsonl"