diff --git a/litellm/llms/vertex_ai/batches/handler.py b/litellm/llms/vertex_ai/batches/handler.py index 8e36c2a0faa..6481b67fad7 100644 --- a/litellm/llms/vertex_ai/batches/handler.py +++ b/litellm/llms/vertex_ai/batches/handler.py @@ -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( diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_handler.py b/tests/test_litellm/llms/vertex_ai/batches/test_handler.py index b9fb5dfe3c5..9535bf17411 100644 --- a/tests/test_litellm/llms/vertex_ai/batches/test_handler.py +++ b/tests/test_litellm/llms/vertex_ai/batches/test_handler.py @@ -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//models/ 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={}))