mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor(vertex_ai): drop unreachable post-path status checks in batches handler
HTTPHandler.post and AsyncHTTPHandler.post call raise_for_status before returning, so the status_code != 200 branches after the create and cancel POSTs could never run. Non-2xx already surfaces as httpx.HTTPStatusError from inside the client. The checks after GETs stay: the get helpers return without raising. Tests that faked a non-raising POST response are replaced by HTTPStatusError propagation coverage.
This commit is contained in:
parent
21df36ed09
commit
7bffbbd1f2
2 changed files with 16 additions and 98 deletions
|
|
@ -98,11 +98,6 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
data=json.dumps(vertex_batch_request),
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
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(
|
||||
response=_json_response
|
||||
|
|
@ -132,10 +127,6 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
error_body[:1000],
|
||||
)
|
||||
raise
|
||||
if response.status_code != 200:
|
||||
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(
|
||||
|
|
@ -473,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({}),
|
||||
|
|
@ -487,11 +478,6 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
)
|
||||
raise
|
||||
|
||||
if response.status_code != 200:
|
||||
raise VertexAIError(
|
||||
status_code=response.status_code, message=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,
|
||||
|
|
@ -525,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({}),
|
||||
|
|
@ -538,10 +524,6 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
e.response.text[:1000],
|
||||
)
|
||||
raise
|
||||
if response.status_code != 200:
|
||||
raise VertexAIError(
|
||||
status_code=response.status_code, message=f"Error: {response.status_code} {response.text}"
|
||||
)
|
||||
|
||||
# AsyncHTTPHandler.get() does not accept a timeout parameter
|
||||
retrieve_response: Final = await client.get(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
||||
|
|
@ -179,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(VertexAIError, match="Error: 500") as exc_info:
|
||||
with pytest.raises(httpx.HTTPStatusError):
|
||||
h.create_batch(
|
||||
_is_async=False,
|
||||
create_batch_data=CREATE_DATA,
|
||||
|
|
@ -197,9 +205,6 @@ def test_create_batch_sync_non_200_raises():
|
|||
max_retries=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "error text" in str(exc_info.value)
|
||||
|
||||
|
||||
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."""
|
||||
|
|
@ -224,32 +229,6 @@ def test_create_batch_input_file_id_without_model_raises_400_before_post():
|
|||
client.post.assert_not_called()
|
||||
|
||||
|
||||
def test_create_batch_async_non_200_raises():
|
||||
h = _make_handler()
|
||||
async_client = MagicMock()
|
||||
async_client.post = AsyncMock(return_value=_http_response(status_code=403))
|
||||
|
||||
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(VertexAIError, match="Error: 403") as exc_info:
|
||||
_run(coro)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "error text" in str(exc_info.value)
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# retrieve_batch
|
||||
# =========================================================================== #
|
||||
|
|
@ -554,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(VertexAIError, 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()
|
||||
|
|
@ -791,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(VertexAIError, 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={}))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue