Merge pull request #35141 from BerriAI/litellm_vertex_batch_create_error_propagation

fix(vertex_ai): surface real error/status on vertex batch create instead of IndexError 500
This commit is contained in:
Mateo Wang 2026-08-07 23:28:36 -07:00 • committed by GitHub
commit c28cbb804c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 142 additions and 103 deletions

View file

@ -14,7 +14,7 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.common_utils import VertexAIError, get_vertex_base_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
from litellm.types.llms.openai import CreateBatchRequest
from litellm.types.llms.vertex_ai import (
@ -98,9 +98,6 @@ class VertexAIBatchPrediction(VertexLLM):
data=json.dumps(vertex_batch_request),
)
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
_json_response: Final = response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
response=_json_response
@ -130,8 +127,6 @@ class VertexAIBatchPrediction(VertexLLM):
error_body[:1000],
)
raise
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
_json_response: Final = response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
@ -243,7 +238,9 @@ class VertexAIBatchPrediction(VertexLLM):
)
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
raise VertexAIError(
status_code=response.status_code, message=f"Error: {response.status_code} {response.text}"
)
_json_response: Final = response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
@ -293,7 +290,9 @@ class VertexAIBatchPrediction(VertexLLM):
headers=headers,
)
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
raise VertexAIError(
status_code=response.status_code, message=f"Error: {response.status_code} {response.text}"
)
_json_response: Final = response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
@ -366,7 +365,9 @@ class VertexAIBatchPrediction(VertexLLM):
)
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
raise VertexAIError(
status_code=response.status_code, message=f"Error: {response.status_code} {response.text}"
)
_json_response: Final = response.json()
vertex_batch_response: Final = (
@ -391,7 +392,9 @@ class VertexAIBatchPrediction(VertexLLM):
params=params,
)
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
raise VertexAIError(
status_code=response.status_code, message=f"Error: {response.status_code} {response.text}"
)
_json_response: Final = response.json()
vertex_batch_response: Final = (
@ -461,7 +464,7 @@ class VertexAIBatchPrediction(VertexLLM):
sync_handler: Final = _get_httpx_client()
try:
response: Final = sync_handler.post(
sync_handler.post(
url=api_base,
headers=headers,
data=json.dumps({}),
@ -475,9 +478,6 @@ class VertexAIBatchPrediction(VertexLLM):
)
raise
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
# HTTPHandler.get() does not accept a timeout parameter
retrieve_response: Final = sync_handler.get(
url=retrieve_api_base,
@ -489,7 +489,10 @@ class VertexAIBatchPrediction(VertexLLM):
retrieve_response.status_code,
retrieve_response.text[:1000],
)
raise Exception(f"Error: {retrieve_response.status_code} {retrieve_response.text}")
raise VertexAIError(
status_code=retrieve_response.status_code,
message=f"Error: {retrieve_response.status_code} {retrieve_response.text}",
)
_json_response: Final = retrieve_response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
@ -508,7 +511,7 @@ class VertexAIBatchPrediction(VertexLLM):
llm_provider=litellm.LlmProviders.VERTEX_AI,
)
try:
response: Final = await client.post(
await client.post(
url=api_base,
headers=headers,
data=json.dumps({}),
@ -521,8 +524,6 @@ class VertexAIBatchPrediction(VertexLLM):
e.response.text[:1000],
)
raise
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
# AsyncHTTPHandler.get() does not accept a timeout parameter
retrieve_response: Final = await client.get(
@ -535,7 +536,10 @@ class VertexAIBatchPrediction(VertexLLM):
retrieve_response.status_code,
retrieve_response.text[:1000],
)
raise Exception(f"Error: {retrieve_response.status_code} {retrieve_response.text}")
raise VertexAIError(
status_code=retrieve_response.status_code,
message=f"Error: {retrieve_response.status_code} {retrieve_response.text}",
)
_json_response: Final = retrieve_response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(

View file

@ -1,7 +1,9 @@
from typing import Any, Final
from urllib.parse import unquote
from litellm._uuid import uuid
from litellm.llms.vertex_ai.common_utils import (
VertexAIError,
_convert_vertex_datetime_to_openai_datetime,
)
from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest
@ -199,16 +201,40 @@ class VertexAIBatchTransformation:
gcs_file_uri format: gs://litellm-testing-bucket/litellm-vertex-files/publishers/google/models/gemini-1.5-flash-001/e9412502-2c91-42a6-8e61-f5c294cc0fc8
returns: "publishers/google/models/gemini-1.5-flash-001"
Raises a 400 `VertexAIError` when the uri carries no parseable model path.
"""
from urllib.parse import unquote
decoded_uri: Final = unquote(gcs_file_uri)
model_path: Final = decoded_uri.split("publishers/")[1]
parts: Final = model_path.split("/")
model: Final = f"publishers/{'/'.join(parts[:3])}"
model: Final = cls._parse_model_from_gcs_file(gcs_file_uri)
if model is None:
raise VertexAIError(
status_code=400,
message=(
"Vertex AI batch creation requires the model to be part of `input_file_id`, but "
f"'{gcs_file_uri}' contains no 'publishers/<publisher>/models/<model>' path segment. "
"Either upload the input file through LiteLLM (POST /v1/files with "
"custom_llm_provider=vertex_ai), which encodes the model into the returned file id, or "
"pass a uri of the form "
"gs://<bucket>/<prefix>/publishers/<publisher>/models/<model>/<file>"
),
)
return model
@classmethod
def _parse_model_from_gcs_file(cls, gcs_file_uri: str) -> str | None:
"""
Returns the `publishers/<publisher>/models/<model>` path from a gcs uri, or None if the uri
does not contain one.
"""
_, separator, model_path = unquote(gcs_file_uri).partition("publishers/")
if not separator:
return None
parts: Final = model_path.split("/")
if len(parts) < 3 or parts[1] != "models" or not parts[2]:
return None
return f"publishers/{'/'.join(parts[:3])}"
@classmethod
def is_unmanaged_gcs_batch_input_file_id(cls, input_file_id: str | None) -> bool:
"""
@ -216,7 +242,11 @@ class VertexAIBatchTransformation:
LiteLLM-managed unified file id) with a `publishers/` model path that
`_get_model_from_gcs_file` can parse.
"""
return input_file_id is not None and input_file_id.startswith("gs://") and "publishers/" in input_file_id
return (
input_file_id is not None
and input_file_id.startswith("gs://")
and cls._parse_model_from_gcs_file(input_file_id) is not None
)
@classmethod
def get_bare_model_name_from_gcs_file(cls, gcs_file_uri: str) -> str:

View file

@ -5,8 +5,10 @@ The handler is HTTP/auth glue around the (separately-tested) pure
``VertexAIBatchTransformation``. Each public method (create / retrieve / list /
cancel) resolves a Vertex access token + URL, branches on ``_is_async``
(returning the coroutine in the async case, doing the sync HTTP call otherwise),
checks the HTTP status, and parses the JSON into ``LiteLLMBatch`` (or the OpenAI
list shape).
and parses the JSON into ``LiteLLMBatch`` (or the OpenAI list shape). POST-backed
calls rely on the client's ``raise_for_status`` (non-2xx surfaces as
``httpx.HTTPStatusError``); GET-backed calls return without raising, so the
handler checks their status codes itself.
We mock only true I/O / auth seams:
* ``_ensure_access_token`` - the Vertex credential seam. Returns a fixed
@ -20,7 +22,7 @@ We mock only true I/O / auth seams:
what URL/headers/body, and that the response is parsed into the litellm
type. Sibling seams are asserted NOT called where relevant.
The ``_is_async`` branch, status-code error paths, and the cancel
The ``_is_async`` branch, the error paths, and the cancel
retrieve-after-cancel sequencing run for real.
"""
@ -40,6 +42,7 @@ sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.llms.vertex_ai.batches.handler import ( # noqa: E402
VertexAIBatchPrediction,
)
from litellm.llms.vertex_ai.common_utils import VertexAIError # noqa: E402
from litellm.types.utils import LiteLLMBatch # noqa: E402
HMOD = "litellm.llms.vertex_ai.batches.handler"
@ -178,13 +181,19 @@ def test_create_batch_async_returns_coroutine_and_uses_async_client():
sync_client.post.assert_not_called()
def test_create_batch_sync_non_200_raises():
def test_create_batch_sync_httpstatuserror_propagates():
"""``HTTPHandler.post`` raises for non-2xx via ``raise_for_status``; the
sync create path must surface that error, not swallow it."""
h = _make_handler()
client = MagicMock()
client.post.return_value = _http_response(status_code=500)
request = httpx.Request("POST", "https://x/batchPredictionJobs")
err_response = httpx.Response(status_code=500, request=request, text="boom")
client.post.side_effect = httpx.HTTPStatusError(
"boom", request=request, response=err_response
)
with patch(f"{HMOD}._get_httpx_client", return_value=client):
with pytest.raises(Exception, match="Error: 500"):
with pytest.raises(httpx.HTTPStatusError):
h.create_batch(
_is_async=False,
create_batch_data=CREATE_DATA,
@ -197,27 +206,27 @@ def test_create_batch_sync_non_200_raises():
)
def test_create_batch_async_non_200_raises():
def test_create_batch_input_file_id_without_model_raises_400_before_post():
"""A gs:// uri with no publishers/<publisher>/models/<model> path is a 400, not a bare 500."""
h = _make_handler()
async_client = MagicMock()
async_client.post = AsyncMock(return_value=_http_response(status_code=403))
client = MagicMock()
with (
patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()),
patch(f"{HMOD}.get_async_httpx_client", return_value=async_client),
):
coro = h.create_batch(
_is_async=True,
create_batch_data=CREATE_DATA,
api_base=None,
vertex_credentials=None,
vertex_project=PROJECT,
vertex_location=LOCATION,
timeout=600.0,
max_retries=None,
)
with pytest.raises(Exception, match="Error: 403"):
_run(coro)
with patch(f"{HMOD}._get_httpx_client", return_value=client):
with pytest.raises(VertexAIError) as exc_info:
h.create_batch(
_is_async=False,
create_batch_data={"input_file_id": "gs://bucket/batch-input.jsonl"},
api_base=None,
vertex_credentials=None,
vertex_project=PROJECT,
vertex_location=LOCATION,
timeout=600.0,
max_retries=None,
)
assert exc_info.value.status_code == 400
assert "gs://bucket/batch-input.jsonl" in str(exc_info.value)
client.post.assert_not_called()
# =========================================================================== #
@ -292,7 +301,7 @@ def test_retrieve_batch_sync_non_200_raises():
patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()),
patch(f"{HMOD}.safe_get", return_value=_http_response(status_code=404)),
):
with pytest.raises(Exception, match="Error: 404"):
with pytest.raises(VertexAIError, match="Error: 404"):
h.retrieve_batch(
_is_async=False,
batch_id=BATCH_ID,
@ -438,7 +447,7 @@ def test_list_batches_sync_non_200_raises():
client.get.return_value = _http_response(status_code=500)
with patch(f"{HMOD}._get_httpx_client", return_value=client):
with pytest.raises(Exception, match="Error: 500"):
with pytest.raises(VertexAIError, match="Error: 500"):
h.list_batches(
_is_async=False,
after=None,
@ -524,27 +533,6 @@ def test_cancel_batch_async_returns_coroutine_posts_then_retrieves():
assert post_kwargs["url"].endswith(":cancel")
def test_cancel_batch_sync_cancel_post_non_200_raises():
h = _make_handler()
client = MagicMock()
client.post.return_value = _http_response(status_code=500)
with patch(f"{HMOD}._get_httpx_client", return_value=client):
with pytest.raises(Exception, match="Error: 500"):
h.cancel_batch(
_is_async=False,
batch_id=BATCH_ID,
api_base=None,
vertex_credentials=None,
vertex_project=PROJECT,
vertex_location=LOCATION,
timeout=600.0,
max_retries=None,
)
# cancel POST failed -> retrieve GET must never fire
client.get.assert_not_called()
def test_cancel_batch_sync_retrieve_non_200_raises():
h = _make_handler()
client = MagicMock()
@ -552,7 +540,7 @@ def test_cancel_batch_sync_retrieve_non_200_raises():
client.get.return_value = _http_response(status_code=404)
with patch(f"{HMOD}._get_httpx_client", return_value=client):
with pytest.raises(Exception, match="Error: 404"):
with pytest.raises(VertexAIError, match="Error: 404"):
h.cancel_batch(
_is_async=False,
batch_id=BATCH_ID,
@ -672,7 +660,7 @@ def test_async_retrieve_batch_non_200_raises():
timeout=600.0,
max_retries=None,
)
with pytest.raises(Exception, match="Error: 500"):
with pytest.raises(VertexAIError, match="Error: 500"):
_run(coro)
@ -726,7 +714,7 @@ def test_async_list_batches_non_200_raises():
timeout=600.0,
max_retries=None,
)
with pytest.raises(Exception, match="Error: 500"):
with pytest.raises(VertexAIError, match="Error: 500"):
_run(coro)
@ -761,28 +749,6 @@ def test_async_cancel_batch_httpstatuserror_and_retrieve_non_200():
_run(coro)
async_client.get.assert_not_awaited()
# (a2) cancel POST returns a plain non-200 (no exception) -> raises
async_client_post500 = MagicMock()
async_client_post500.post = AsyncMock(return_value=_http_response(status_code=500))
async_client_post500.get = AsyncMock()
with (
patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()),
patch(f"{HMOD}.get_async_httpx_client", return_value=async_client_post500),
):
coro = h.cancel_batch(
_is_async=True,
batch_id=BATCH_ID,
api_base=None,
vertex_credentials=None,
vertex_project=PROJECT,
vertex_location=LOCATION,
timeout=600.0,
max_retries=None,
)
with pytest.raises(Exception, match="Error: 500"):
_run(coro)
async_client_post500.get.assert_not_awaited()
# (b) retrieve-after-cancel returns non-200
async_client2 = MagicMock()
async_client2.post = AsyncMock(return_value=_http_response(json_body={}))
@ -801,5 +767,5 @@ def test_async_cancel_batch_httpstatuserror_and_retrieve_non_200():
timeout=600.0,
max_retries=None,
)
with pytest.raises(Exception, match="Error: 404"):
with pytest.raises(VertexAIError, match="Error: 404"):
_run(coro)

View file

@ -25,6 +25,7 @@ from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402
VertexAIBatchTransformation,
)
from litellm.llms.vertex_ai.common_utils import ( # noqa: E402
VertexAIError,
_convert_vertex_datetime_to_openai_datetime,
)
from litellm.types.utils import LiteLLMBatch # noqa: E402
@ -69,6 +70,24 @@ def test_transform_openai_request_missing_input_file_id_raises():
T.transform_openai_batch_request_to_vertex_ai_batch_request({})
@pytest.mark.parametrize(
"input_file_id",
[
"gs://bucket/no-model-here.jsonl",
"gs://bucket/publishers/google/gemini-1.5-flash-001/file-uuid",
"gs://bucket/publishers/google/models",
"gs://bucket/publishers/google/models//file-uuid",
],
)
def test_transform_openai_request_unparseable_model_raises_400(input_file_id: str):
"""An input_file_id with no parseable model path is a client error, not an IndexError -> 500."""
with pytest.raises(VertexAIError) as exc_info:
T.transform_openai_batch_request_to_vertex_ai_batch_request({"input_file_id": input_file_id})
assert exc_info.value.status_code == 400
assert input_file_id in str(exc_info.value)
# =========================================================================== #
# transform_vertex_ai_batch_response_to_openai_batch_response
# =========================================================================== #
@ -299,9 +318,29 @@ def test_get_model_from_gcs_file_url_encoded():
assert T._get_model_from_gcs_file(encoded) == "publishers/google/models/gemini-1.5-flash-001"
def test_get_model_from_gcs_file_no_publishers_raises():
with pytest.raises(IndexError):
def test_get_model_from_gcs_file_no_publishers_raises_400():
with pytest.raises(VertexAIError) as exc_info:
T._get_model_from_gcs_file("gs://bucket/no-model-here.jsonl")
assert exc_info.value.status_code == 400
# =========================================================================== #
# is_unmanaged_gcs_batch_input_file_id
# =========================================================================== #
@pytest.mark.parametrize(
"input_file_id, expected",
[
(INPUT_FILE, True),
(None, False),
("file-abc123", False),
("gs://bucket/no-model-here.jsonl", False),
("gs://bucket/publishers/google/gemini-1.5-flash-001/file-uuid", False),
],
)
def test_is_unmanaged_gcs_batch_input_file_id(input_file_id, expected):
assert T.is_unmanaged_gcs_batch_input_file_id(input_file_id) is expected
# =========================================================================== #