mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
697da95bf4
commit
7c35312657
8 changed files with 617 additions and 63 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue