From b5ed9fbd3a81fee7618ff5c195c0ead1a2e42fdc Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 17 Jun 2026 11:50:59 +0530 Subject: [PATCH] Add handler and transformation tests for all providers --- tests/test_litellm/batches/__init__.py | 0 tests/test_litellm/llms/anthropic/__init__.py | 0 .../llms/anthropic/batches/__init__.py | 0 .../llms/anthropic/batches/test_handler.py | 286 +++++++ .../anthropic/batches/test_transformation.py | 650 ++++++++++++++ tests/test_litellm/llms/azure/__init__.py | 0 .../llms/azure/batches/__init__.py | 0 .../llms/azure/batches/test_handler.py | 491 +++++++++++ tests/test_litellm/llms/base_llm/__init__.py | 0 .../llms/base_llm/batches/__init__.py | 0 .../batches/base_batches_config_test.py | 128 +++ .../base_llm/batches/test_transformation.py | 231 +++++ tests/test_litellm/llms/bedrock/__init__.py | 0 .../llms/bedrock/batches/__init__.py | 0 .../bedrock/batches/test_transformation.py | 684 +++++++++++++++ .../llms/vertex_ai/batches/__init__.py | 0 .../llms/vertex_ai/batches/test_handler.py | 805 ++++++++++++++++++ .../vertex_ai/batches/test_transformation.py | 396 +++++++++ 18 files changed, 3671 insertions(+) create mode 100644 tests/test_litellm/batches/__init__.py create mode 100644 tests/test_litellm/llms/anthropic/__init__.py create mode 100644 tests/test_litellm/llms/anthropic/batches/__init__.py create mode 100644 tests/test_litellm/llms/azure/__init__.py create mode 100644 tests/test_litellm/llms/azure/batches/__init__.py create mode 100644 tests/test_litellm/llms/base_llm/__init__.py create mode 100644 tests/test_litellm/llms/base_llm/batches/__init__.py create mode 100644 tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py create mode 100644 tests/test_litellm/llms/bedrock/__init__.py create mode 100644 tests/test_litellm/llms/bedrock/batches/__init__.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/__init__.py diff --git a/tests/test_litellm/batches/__init__.py b/tests/test_litellm/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/anthropic/__init__.py b/tests/test_litellm/llms/anthropic/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/anthropic/batches/__init__.py b/tests/test_litellm/llms/anthropic/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/anthropic/batches/test_handler.py b/tests/test_litellm/llms/anthropic/batches/test_handler.py index e69de29bb2d..0a472d86257 100644 --- a/tests/test_litellm/llms/anthropic/batches/test_handler.py +++ b/tests/test_litellm/llms/anthropic/batches/test_handler.py @@ -0,0 +1,286 @@ +""" +Unit tests for litellm/llms/anthropic/batches/handler.py + +AnthropicBatchesHandler is the HTTP/auth glue for retrieving Anthropic Message +Batches. It resolves credentials, builds the retrieve URL + auth headers via the +provider config, fires a single GET against the async httpx client, and hands the +response to the config's transform. These tests mock ONLY the genuine I/O seams - +the async httpx client (network) and credential resolution (secret managers / +env) - and assert exactly which seam fired, with what URL/headers, and that the +parsed result is the LiteLLMBatch the transform produced. + +The sync ``retrieve_batch`` dispatch (``_is_async`` true -> coroutine, false -> +asyncio.run) is exercised directly, mirroring the dispatch-contract discipline in +tests/test_litellm/batches/test_main.py. +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler +from litellm.types.utils import LiteLLMBatch + + +def _ok_batch_response(): + """A real httpx.Response shaped like an Anthropic MessageBatch retrieval.""" + return httpx.Response( + status_code=200, + json={ + "id": "msgbatch_abc", + "processing_status": "ended", + "created_at": "2024-09-24T10:00:00Z", + "ended_at": "2024-09-24T11:00:00Z", + "request_counts": {"succeeded": 2, "errored": 0}, + }, + request=httpx.Request( + "GET", "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ), + ) + + +@pytest.fixture +def handler(): + return AnthropicBatchesHandler() + + +@pytest.fixture +def patched_client(): + """Patch the async httpx client seam; yield the (fake_client, factory).""" + fake_client = MagicMock() + fake_client.get = AsyncMock(return_value=_ok_batch_response()) + with patch( + "litellm.llms.anthropic.batches.handler.get_async_httpx_client", + return_value=fake_client, + ) as factory: + yield fake_client, factory + + +@pytest.mark.asyncio +async def test_aretrieve_batch_fires_get_with_correct_url_and_headers( + handler, patched_client +): + fake_client, factory = patched_client + + batch = await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + + # The single network seam fired exactly once. + fake_client.get.assert_awaited_once() + _, call_kwargs = fake_client.get.call_args + # Exact URL built by get_retrieve_batch_url. + assert call_kwargs["url"] == ( + "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ) + # Auth + version + beta headers built by validate_environment. + headers = call_kwargs["headers"] + assert headers["x-api-key"] == "sk-ant-test" + assert headers["anthropic-version"] == "2023-06-01" + assert headers["anthropic-beta"] == "message-batches-2024-09-24" + + # Response parsed through the config transform. + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "msgbatch_abc" + assert batch.status == "completed" + assert batch.request_counts.completed == 2 + + +@pytest.mark.asyncio +async def test_aretrieve_batch_uses_anthropic_provider_for_client( + handler, patched_client +): + from litellm.types.utils import LlmProviders + + _, factory = patched_client + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + _, kwargs = factory.call_args + assert kwargs["llm_provider"] == LlmProviders.ANTHROPIC + + +@pytest.mark.asyncio +async def test_aretrieve_batch_resolves_api_key_from_model_info( + handler, patched_client +): + fake_client, _ = patched_client + # api_key=None -> handler falls back to AnthropicModelInfo.get_api_key(). + with patch.object( + handler.anthropic_model_info, "get_api_key", return_value="sk-from-env" + ): + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key=None, + timeout=60.0, + max_retries=0, + ) + _, call_kwargs = fake_client.get.call_args + assert call_kwargs["headers"]["x-api-key"] == "sk-from-env" + + +@pytest.mark.asyncio +async def test_aretrieve_batch_missing_api_key_raises(handler, patched_client): + fake_client, _ = patched_client + # No api_key and resolver yields None -> hard error before any network call. + with patch.object( + handler.anthropic_model_info, "get_api_key", return_value=None + ): + with pytest.raises(ValueError, match="Missing Anthropic API Key"): + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key=None, + timeout=60.0, + max_retries=0, + ) + fake_client.get.assert_not_called() + + +@pytest.mark.asyncio +async def test_aretrieve_batch_resolves_default_api_base(handler, patched_client): + fake_client, _ = patched_client + # api_base=None -> resolved via get_api_base() default before URL build. + with patch.object( + handler.anthropic_model_info, + "get_api_base", + return_value="https://api.anthropic.com", + ): + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base=None, + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + _, call_kwargs = fake_client.get.call_args + assert call_kwargs["url"] == ( + "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_raises_for_status(handler): + # A non-2xx response must surface via raise_for_status (no silent parse). + error_response = httpx.Response( + status_code=404, + json={"error": "not found"}, + request=httpx.Request( + "GET", "https://api.anthropic.com/v1/messages/batches/missing" + ), + ) + fake_client = MagicMock() + fake_client.get = AsyncMock(return_value=error_response) + with patch( + "litellm.llms.anthropic.batches.handler.get_async_httpx_client", + return_value=fake_client, + ): + with pytest.raises(httpx.HTTPStatusError): + await handler.aretrieve_batch( + batch_id="missing", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_invokes_pre_call_logging(handler, patched_client): + fake_client, _ = patched_client + logging_obj = MagicMock() + await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + logging_obj=logging_obj, + ) + logging_obj.pre_call.assert_called_once() + pre_kwargs = logging_obj.pre_call.call_args.kwargs + assert pre_kwargs["input"] == "msgbatch_abc" + assert pre_kwargs["api_key"] == "sk-ant-test" + # The logged api_base is the full retrieve URL, not the bare base. + assert pre_kwargs["additional_args"]["api_base"] == ( + "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_builds_default_logging_obj_when_absent( + handler, patched_client +): + # logging_obj=None -> handler constructs a real Logging object; the call + # must still complete (no AttributeError on a missing logger). + _, _ = patched_client + with patch( + "litellm.litellm_core_utils.litellm_logging.Logging" + ) as logging_cls: + logging_cls.return_value = MagicMock() + batch = await handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + logging_obj=None, + ) + logging_cls.assert_called_once() + # call_type wires through to the constructed logging object. + assert logging_cls.call_args.kwargs["call_type"] == "batch_retrieve" + assert batch.id == "msgbatch_abc" + + +# =========================================================================== # +# retrieve_batch dispatch (sync wrapper) +# =========================================================================== # + + +async def test_retrieve_batch_async_returns_coroutine(handler, patched_client): + # _is_async=True -> returns the un-awaited coroutine (caller awaits it). + import asyncio + + coro = handler.retrieve_batch( + _is_async=True, + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + assert asyncio.iscoroutine(coro) + # Await directly - robust under asyncio_mode=auto's session-scoped loop + # (manually driving get_event_loop().run_until_complete() breaks when prior + # async tests in the suite have already used/closed that loop). + batch = await coro + assert batch.id == "msgbatch_abc" + + +def test_retrieve_batch_sync_runs_to_result(handler, patched_client): + # _is_async=False -> asyncio.run(...) returns the resolved LiteLLMBatch. + batch = handler.retrieve_batch( + _is_async=False, + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key="sk-ant-test", + timeout=60.0, + max_retries=0, + ) + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "msgbatch_abc" + assert batch.status == "completed" diff --git a/tests/test_litellm/llms/anthropic/batches/test_transformation.py b/tests/test_litellm/llms/anthropic/batches/test_transformation.py index e69de29bb2d..4a2adb01ea5 100644 --- a/tests/test_litellm/llms/anthropic/batches/test_transformation.py +++ b/tests/test_litellm/llms/anthropic/batches/test_transformation.py @@ -0,0 +1,650 @@ +""" +Unit tests for litellm/llms/anthropic/batches/transformation.py + +AnthropicBatchesConfig is the pure request/response mapping layer for Anthropic +Message Batches. It builds auth headers, constructs the batch create/retrieve +URLs, and (most importantly) maps an Anthropic MessageBatch JSON response into a +LiteLLM/OpenAI ``LiteLLMBatch`` (status mapping, timestamp parsing, request +counts). A silent bug here mis-reports batch status or counts to the caller, so +these tests assert EXACT output values rather than "ran without error". + +Pure transform code runs for real. The only mocked boundaries are the credential +resolvers on AnthropicModelInfo (get_api_base / get_auth_header), which would +otherwise read process env / secret managers - mocking them keeps the URL/header +assertions deterministic without touching production transform logic. +""" + +import os +import sys +import time +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig +from litellm.types.utils import LiteLLMBatch, LlmProviders + + +@pytest.fixture +def config(): + return AnthropicBatchesConfig() + + +def _response(payload): + """A real httpx.Response whose .json() yields ``payload``.""" + return httpx.Response( + status_code=200, + json=payload, + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + + +# =========================================================================== # +# custom_llm_provider +# =========================================================================== # + + +def test_custom_llm_provider_is_anthropic(config): + assert config.custom_llm_provider == LlmProviders.ANTHROPIC + + +# =========================================================================== # +# validate_environment (auth + fixed headers + beta header) +# =========================================================================== # + + +def test_validate_environment_builds_headers_with_api_key(config): + headers = config.validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-test", + ) + assert headers["accept"] == "application/json" + assert headers["anthropic-version"] == "2023-06-01" + assert headers["content-type"] == "application/json" + # Plain api key -> x-api-key auth header. + assert headers["x-api-key"] == "sk-ant-test" + # Beta header is injected when not already present. + assert headers["anthropic-beta"] == "message-batches-2024-09-24" + + +def test_validate_environment_preserves_existing_beta_header(config): + headers = config.validate_environment( + headers={"anthropic-beta": "custom-beta-value"}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-test", + ) + # Existing beta header must NOT be overwritten. + assert headers["anthropic-beta"] == "custom-beta-value" + + +def test_validate_environment_oauth_key_uses_bearer(config): + headers = config.validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-oat-abc123", + ) + # OAuth tokens map to Authorization: Bearer, not x-api-key. + assert headers["authorization"] == "Bearer sk-ant-oat-abc123" + assert "x-api-key" not in headers + + +def test_validate_environment_missing_key_raises(config): + # No api_key passed and no env credentials -> get_auth_header returns None. + with patch.object( + config.anthropic_model_info, "get_auth_header", return_value=None + ): + with pytest.raises(ValueError, match="Missing Anthropic API Key"): + config.validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + +# =========================================================================== # +# get_complete_batch_url (batch creation URL) +# =========================================================================== # + + +def test_get_complete_batch_url_appends_path(config): + url = config.get_complete_batch_url( + api_base="https://api.anthropic.com", + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == "https://api.anthropic.com/v1/messages/batches" + + +def test_get_complete_batch_url_strips_trailing_slash(config): + url = config.get_complete_batch_url( + api_base="https://api.anthropic.com/", + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == "https://api.anthropic.com/v1/messages/batches" + + +def test_get_complete_batch_url_already_complete_is_unchanged(config): + complete = "https://proxy.internal/v1/messages/batches" + url = config.get_complete_batch_url( + api_base=complete, + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == complete + + +def test_get_complete_batch_url_uses_default_api_base(config): + # api_base=None -> falls back to get_api_base() default. + with patch.object( + config.anthropic_model_info, + "get_api_base", + return_value="https://api.anthropic.com", + ): + url = config.get_complete_batch_url( + api_base=None, + api_key="sk", + model="claude-3", + optional_params={}, + litellm_params={}, + data={}, # type: ignore[arg-type] + ) + assert url == "https://api.anthropic.com/v1/messages/batches" + + +# =========================================================================== # +# get_retrieve_batch_url (batch retrieval URL + path encoding) +# =========================================================================== # + + +def test_get_retrieve_batch_url_happy_path(config): + url = config.get_retrieve_batch_url( + api_base="https://api.anthropic.com", + batch_id="msgbatch_123", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/msgbatch_123" + + +def test_get_retrieve_batch_url_strips_trailing_slash(config): + url = config.get_retrieve_batch_url( + api_base="https://api.anthropic.com/", + batch_id="msgbatch_123", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/msgbatch_123" + + +def test_get_retrieve_batch_url_encodes_batch_id(config): + # batch_id is user-controlled; a path-traversal attempt must be percent-encoded. + url = config.get_retrieve_batch_url( + api_base="https://api.anthropic.com", + batch_id="a/b id", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/a%2Fb%20id" + + +def test_get_retrieve_batch_url_rejects_dot_segment(config): + with pytest.raises(ValueError, match="dot path segment"): + config.get_retrieve_batch_url( + api_base="https://api.anthropic.com", + batch_id="..", + optional_params={}, + litellm_params={}, + ) + + +def test_get_retrieve_batch_url_uses_default_api_base(config): + with patch.object( + config.anthropic_model_info, + "get_api_base", + return_value="https://api.anthropic.com", + ): + url = config.get_retrieve_batch_url( + api_base=None, + batch_id="msgbatch_123", + optional_params={}, + litellm_params={}, + ) + assert url == "https://api.anthropic.com/v1/messages/batches/msgbatch_123" + + +# =========================================================================== # +# transform_retrieve_batch_request (no-op for Anthropic) +# =========================================================================== # + + +def test_transform_retrieve_batch_request_returns_empty_dict(config): + assert ( + config.transform_retrieve_batch_request( + batch_id="msgbatch_123", optional_params={}, litellm_params={} + ) + == {} + ) + + +# =========================================================================== # +# Unimplemented create-batch methods raise NotImplementedError +# =========================================================================== # + + +def test_transform_create_batch_request_not_implemented(config): + with pytest.raises(NotImplementedError, match="not yet implemented"): + config.transform_create_batch_request( + model="claude-3", + create_batch_data={}, # type: ignore[arg-type] + optional_params={}, + litellm_params={}, + ) + + +def test_transform_create_batch_response_not_implemented(config): + with pytest.raises(NotImplementedError, match="not yet implemented"): + config.transform_create_batch_response( + model="claude-3", + raw_response=_response({}), + logging_obj=MagicMock(), + litellm_params={}, + ) + + +# =========================================================================== # +# transform_retrieve_batch_response (the core mapping - exact values) +# =========================================================================== # + + +def test_transform_retrieve_response_in_progress(config): + raw = _response( + { + "id": "msgbatch_abc", + "processing_status": "in_progress", + "created_at": "2024-09-24T10:00:00Z", + "expires_at": "2024-09-25T10:00:00Z", + "request_counts": { + "processing": 3, + "succeeded": 2, + "errored": 1, + "canceled": 0, + "expired": 0, + }, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "msgbatch_abc" + assert batch.object == "batch" + assert batch.endpoint == "/v1/messages" + assert batch.status == "in_progress" + # output_file_id mirrors the batch id for Anthropic. + assert batch.output_file_id == "msgbatch_abc" + assert batch.input_file_id == "None" + assert batch.completion_window == "24h" + # created_at parsed from ISO8601 (UTC). + assert batch.created_at == 1727172000 + assert batch.expires_at == 1727258400 + # in_progress -> in_progress_at is set to created_at. + assert batch.in_progress_at == 1727172000 + assert batch.completed_at is None + assert batch.cancelling_at is None + assert batch.cancelled_at is None + # request_counts: total = processing+succeeded+errored+canceled+expired. + assert batch.request_counts.total == 6 + assert batch.request_counts.completed == 2 + assert batch.request_counts.failed == 1 + assert batch.metadata == {} + + +def test_transform_retrieve_response_ended_maps_to_completed(config): + raw = _response( + { + "id": "msgbatch_done", + "processing_status": "ended", + "created_at": "2024-09-24T10:00:00Z", + "ended_at": "2024-09-24T11:00:00Z", + "request_counts": {"succeeded": 5, "errored": 0}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # "ended" -> OpenAI "completed". + assert batch.status == "completed" + # completed_at populated only because processing_status == "ended". + assert batch.completed_at == 1727175600 + # not in_progress -> in_progress_at stays None. + assert batch.in_progress_at is None + assert batch.request_counts.total == 5 + assert batch.request_counts.completed == 5 + + +def test_transform_retrieve_response_canceling_maps_to_cancelling(config): + raw = _response( + { + "id": "msgbatch_cancel", + "processing_status": "canceling", + "created_at": "2024-09-24T10:00:00Z", + "cancel_initiated_at": "2024-09-24T10:30:00Z", + "ended_at": "2024-09-24T10:45:00Z", + "request_counts": {}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # "canceling" -> OpenAI "cancelling". + assert batch.status == "cancelling" + assert batch.cancelling_at == 1727173800 + # cancelled_at = ended_at when canceling and ended_at present. + assert batch.cancelled_at == 1727174700 + assert batch.completed_at is None + + +def test_transform_retrieve_response_unknown_status_defaults_in_progress(config): + raw = _response( + { + "id": "msgbatch_x", + "processing_status": "some_future_status", + "request_counts": {}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # Unmapped status falls back to in_progress (don't 500 on new enum values). + assert batch.status == "in_progress" + + +def test_transform_retrieve_response_missing_id_and_status_defaults(config): + # Empty body: id defaults to "", status defaults to "in_progress". + raw = _response({}) + before = int(time.time()) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + after = int(time.time()) + assert batch.id == "" + assert batch.status == "in_progress" + # No created_at -> created_at falls back to int(time.time()). + assert before <= batch.created_at <= after + # No created_at -> in_progress_at (which mirrors created_at) is None. + assert batch.in_progress_at is None + assert batch.request_counts.total == 0 + + +def test_transform_retrieve_response_archived_sets_expired_at(config): + raw = _response( + { + "id": "msgbatch_arch", + "processing_status": "ended", + "created_at": "2024-09-24T10:00:00Z", + "ended_at": "2024-09-24T11:00:00Z", + "archived_at": "2024-09-26T10:00:00Z", + "request_counts": {}, + } + ) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + # archived_at present -> expired_at populated. + assert batch.expired_at == 1727344800 + + +def test_transform_retrieve_response_bad_timestamp_is_none(config): + raw = _response( + { + "id": "msgbatch_bad", + "processing_status": "in_progress", + "created_at": "not-a-real-timestamp", + "request_counts": {}, + } + ) + before = int(time.time()) + batch = config.transform_retrieve_batch_response( + model=None, raw_response=raw, logging_obj=MagicMock(), litellm_params={} + ) + after = int(time.time()) + # Unparseable created_at -> parse_timestamp returns None, created_at falls + # back to time.time(). + assert before <= batch.created_at <= after + + +def test_transform_retrieve_response_unparseable_json_raises(config): + bad = httpx.Response( + status_code=200, + content=b"not json", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + with pytest.raises(ValueError, match="Failed to parse Anthropic batch response"): + config.transform_retrieve_batch_response( + model=None, raw_response=bad, logging_obj=MagicMock(), litellm_params={} + ) + + +# =========================================================================== # +# get_error_class +# =========================================================================== # + + +def test_get_error_class_with_dict_headers(config): + err = config.get_error_class( + error_message="rate limited", status_code=429, headers={"x-ratelimit": "0"} + ) + from litellm.llms.anthropic.common_utils import AnthropicError + + assert isinstance(err, AnthropicError) + assert err.status_code == 429 + assert err.message == "rate limited" + + +def test_get_error_class_with_httpx_headers(config): + hdrs = httpx.Headers({"retry-after": "5"}) + err = config.get_error_class( + error_message="server error", status_code=500, headers=hdrs + ) + assert err.status_code == 500 + assert err.message == "server error" + + +# =========================================================================== # +# transform_response (batch results JSONL -> summed usage on ModelResponse) +# =========================================================================== # + + +def test_transform_response_sums_usage_across_lines(config): + from litellm.types.utils import ModelResponse, Usage + + # Two result lines; transform_parsed_response is stubbed to attach a fixed + # Usage per line so we can assert the SUM is what lands on model_response. + line1 = '{"result": {"message": {"content": [{"type": "text", "text": "a"}]}}}' + line2 = '{"result": {"message": {"content": [{"type": "text", "text": "b"}]}}}' + raw = httpx.Response( + status_code=200, + text=f"{line1}\n{line2}\n", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + + model_response = ModelResponse() + + def fake_transform_parsed(*, completion_response, raw_response, model_response): + mr = ModelResponse() + setattr( + mr, + "usage", + Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + return mr + + with patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ): + out = config.transform_response( + model="claude-3", + raw_response=raw, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert out is model_response + usage = getattr(out, "usage") + # Two lines * (10 prompt, 5 completion) summed. + assert usage.prompt_tokens == 20 + assert usage.completion_tokens == 10 + assert usage.total_tokens == 30 + + +def test_transform_response_skips_malformed_lines(config): + from litellm.types.utils import ModelResponse, Usage + + valid = '{"result": {"message": {"content": [{"type": "text", "text": "a"}]}}}' + # Interior blank line (survives the outer strip) exercises the empty-line + # `continue`; leading not-json exercises the JSONDecodeError `continue`. + raw = httpx.Response( + status_code=200, + text=f"not-json\n\n{valid}\n", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + model_response = ModelResponse() + + def fake_transform_parsed(*, completion_response, raw_response, model_response): + mr = ModelResponse() + setattr( + mr, "usage", Usage(prompt_tokens=7, completion_tokens=3, total_tokens=10) + ) + return mr + + with patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ): + out = config.transform_response( + model="claude-3", + raw_response=raw, + model_response=model_response, + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + # Only the single valid line contributed usage; malformed/empty skipped. + usage = getattr(out, "usage") + assert usage.prompt_tokens == 7 + assert usage.completion_tokens == 3 + + +def test_transform_response_reraises_unexpected_error(config): + from litellm.types.utils import ModelResponse, Usage + + valid = '{"result": {"message": {"content": [{"type": "text", "text": "a"}]}}}' + raw = httpx.Response( + status_code=200, + text=f"{valid}\n", + request=httpx.Request("GET", "https://api.anthropic.com"), + ) + + def fake_transform_parsed(*, completion_response, raw_response, model_response): + mr = ModelResponse() + setattr(mr, "usage", Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)) + return mr + + # A non-JSONDecodeError raised during usage aggregation must propagate + # (the outer `except Exception: raise e`), not be swallowed. + with patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ), patch( + "litellm.cost_calculator.BaseTokenUsageProcessor.combine_usage_objects", + side_effect=RuntimeError("boom"), + ): + with pytest.raises(RuntimeError, match="boom"): + config.transform_response( + model="claude-3", + raw_response=raw, + model_response=ModelResponse(), + logging_obj=MagicMock(), + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +# --------------------------------------------------------------------------- # +# Shared BaseBatchesConfig contract suite (consistency net across providers). +# This subclass supplies anthropic fixtures; the inherited contract tests run +# automatically. See base_batches_config_test.py. +# --------------------------------------------------------------------------- # + +from litellm.types.utils import LlmProviders # noqa: E402 +from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 + BatchesConfigContractTests, +) + + +class TestAnthropicBatchesContract(BatchesConfigContractTests): + def make_config(self): + from litellm.llms.anthropic.batches.transformation import ( + AnthropicBatchesConfig, + ) + + return AnthropicBatchesConfig() + + expected_provider = LlmProviders.ANTHROPIC + supports_create = False # anthropic raises NotImplementedError on create + supports_retrieve_response = True + + def sample_retrieve_response_body(self) -> dict: + return { + "id": "msgbatch_123", + "processing_status": "ended", + "created_at": "2024-01-01T00:00:00Z", + "ended_at": "2024-01-02T00:00:00Z", + "request_counts": {"succeeded": 2, "errored": 1}, + } + + expected_retrieve_batch_id = "msgbatch_123" + expected_retrieve_status = "completed" # "ended" -> "completed" diff --git a/tests/test_litellm/llms/azure/__init__.py b/tests/test_litellm/llms/azure/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/azure/batches/__init__.py b/tests/test_litellm/llms/azure/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/azure/batches/test_handler.py b/tests/test_litellm/llms/azure/batches/test_handler.py index e69de29bb2d..f2332a7de7c 100644 --- a/tests/test_litellm/llms/azure/batches/test_handler.py +++ b/tests/test_litellm/llms/azure/batches/test_handler.py @@ -0,0 +1,491 @@ +"""Unit tests for ``AzureBatchesAPI`` (litellm/llms/azure/batches/handler.py). + +The Azure batches handler is HTTP/auth glue: each public method +(create/retrieve/cancel/list) resolves an Azure OpenAI client via the inherited +``get_azure_openai_client`` seam, branches on ``_is_async`` (returning the +``a*`` coroutine in the async case, calling the sync client otherwise), validates +the client type, and parses the SDK response into ``LiteLLMBatch``. + +We mock only true boundaries: + * ``get_azure_openai_client`` - the credential/client-construction seam. We + assert the EXACT auth args (api_key / api_base / api_version / client / + _is_async / litellm_params) forwarded to it. + * the returned Azure OpenAI client's ``batches.*`` methods - the network call. + We assert the request data forwarded and that the SDK response is parsed + into ``LiteLLMBatch`` (sibling SDK methods asserted NOT called). + +Pure logic (the _is_async branch, the isinstance guards, the model_dump parse) +runs for real. +""" + +from __future__ import annotations + +import asyncio +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from openai import AsyncOpenAI, OpenAI # noqa: E402 + +from litellm.llms.azure.azure import AsyncAzureOpenAI, AzureOpenAI # noqa: E402 +from litellm.llms.azure.batches.handler import AzureBatchesAPI # noqa: E402 +from litellm.types.utils import LiteLLMBatch # noqa: E402 + +GET_CLIENT = "litellm.llms.azure.batches.handler.AzureBatchesAPI.get_azure_openai_client" + +AUTH_KW = dict( + api_key="sk-azure-test", + api_base="https://my-azure.openai.azure.com", + api_version="2024-12-01", + timeout=600.0, + max_retries=3, +) + +CREATE_DATA = { + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc", +} +RETRIEVE_DATA = {"batch_id": "batch-123"} +CANCEL_DATA = {"batch_id": "batch-123"} + + +def _batch_dict(batch_id: str = "batch-123", status: str = "completed") -> dict: + """A minimal-but-valid dict for ``LiteLLMBatch(**response.model_dump())``.""" + return { + "id": batch_id, + "completion_window": "24h", + "created_at": 1700000000, + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc", + "object": "batch", + "status": status, + "output_file_id": "file-out-xyz", + } + + +def _sdk_response(batch_dict: dict) -> MagicMock: + """An object that mimics the OpenAI SDK Batch: only ``.model_dump()`` is used.""" + resp = MagicMock() + resp.model_dump.return_value = batch_dict + return resp + + +def _sync_client() -> MagicMock: + """A sync Azure client (passes ``isinstance(.., AzureOpenAI)``).""" + return MagicMock(spec=AzureOpenAI) + + +def _async_client() -> MagicMock: + """An async Azure client (passes ``isinstance(.., AsyncAzureOpenAI)``). + + The ``batches.*`` SDK methods are awaited by the handler, so they must be + AsyncMocks. + """ + client = MagicMock(spec=AsyncAzureOpenAI) + client.batches.create = AsyncMock() + client.batches.retrieve = AsyncMock() + client.batches.cancel = AsyncMock() + client.batches.list = AsyncMock() + return client + + +@pytest.fixture +def handler() -> AzureBatchesAPI: + return AzureBatchesAPI() + + +# =========================================================================== # +# create_batch - sync path +# =========================================================================== # + + +def test_create_sync_forwards_auth_to_client_seam(handler): + client = _sync_client() + client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.create_batch( + _is_async=False, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + # EXACT auth args forwarded to the client-construction seam. + assert get_client.call_count == 1 + kw = get_client.call_args.kwargs + assert kw["api_key"] == "sk-azure-test" + assert kw["api_base"] == "https://my-azure.openai.azure.com" + assert kw["api_version"] == "2024-12-01" + assert kw["_is_async"] is False + assert kw["client"] is None + # litellm_params defaults to {} (not None) when not supplied. + assert kw["litellm_params"] == {} + + # PAYLOAD: request data forwarded verbatim to the SDK as kwargs. + client.batches.create.assert_called_once_with(**CREATE_DATA) + # sibling SDK seams untouched. + client.batches.retrieve.assert_not_called() + client.batches.cancel.assert_not_called() + + # RESULT: parsed into LiteLLMBatch from the SDK response's model_dump. + assert isinstance(result, LiteLLMBatch) + assert result.id == "batch-123" + assert result.status == "completed" + assert result.output_file_id == "file-out-xyz" + + +def test_create_sync_passes_litellm_params_through(handler): + client = _sync_client() + client.batches.create.return_value = _sdk_response(_batch_dict()) + lp = {"azure_ad_token": "tok", "tenant_id": "t1"} + + with patch(GET_CLIENT, return_value=client) as get_client: + handler.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + litellm_params=lp, + **AUTH_KW, + ) + + assert get_client.call_args.kwargs["litellm_params"] == lp + + +def test_create_sync_explicit_client_forwarded_to_seam(handler): + sentinel_client = _sync_client() + sentinel_client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=sentinel_client) as get_client: + handler.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + client=sentinel_client, + **AUTH_KW, + ) + + assert get_client.call_args.kwargs["client"] is sentinel_client + + +def test_create_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.create_batch( + _is_async=False, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + +# =========================================================================== # +# create_batch - async path +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create_async_returns_coroutine_and_awaits_async_client(handler): + client = _async_client() + client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.create_batch( + _is_async=True, create_batch_data=CREATE_DATA, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.create.assert_awaited_once_with(**CREATE_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.id == "batch-123" + + +@pytest.mark.asyncio +async def test_create_async_rejects_sync_client(handler): + """_is_async=True but seam returns a sync client -> ValueError, no network.""" + sync_client = _sync_client() + + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="not an instance of AsyncOpenAI"): + handler.create_batch( + _is_async=True, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + sync_client.batches.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_acreate_batch_parses_response(handler): + client = _async_client() + client.batches.create.return_value = _sdk_response(_batch_dict(status="validating")) + + result = await handler.acreate_batch( + create_batch_data=CREATE_DATA, azure_client=client + ) + + client.batches.create.assert_awaited_once_with(**CREATE_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.status == "validating" + + +# =========================================================================== # +# retrieve_batch +# =========================================================================== # + + +def test_retrieve_sync_dispatch_payload_and_result(handler): + client = _sync_client() + client.batches.retrieve.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.retrieve_batch( + _is_async=False, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + + assert get_client.call_args.kwargs["_is_async"] is False + client.batches.retrieve.assert_called_once_with(**RETRIEVE_DATA) + client.batches.create.assert_not_called() + client.batches.cancel.assert_not_called() + assert isinstance(result, LiteLLMBatch) + assert result.id == "batch-123" + + +def test_retrieve_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.retrieve_batch( + _is_async=False, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + + +@pytest.mark.asyncio +async def test_retrieve_async_returns_coroutine_and_awaits(handler): + client = _async_client() + client.batches.retrieve.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.retrieve_batch( + _is_async=True, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.retrieve.assert_awaited_once_with(**RETRIEVE_DATA) + assert isinstance(result, LiteLLMBatch) + + +@pytest.mark.asyncio +async def test_retrieve_async_rejects_sync_client(handler): + sync_client = _sync_client() + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="not an instance of AsyncOpenAI"): + handler.retrieve_batch( + _is_async=True, retrieve_batch_data=RETRIEVE_DATA, **AUTH_KW + ) + sync_client.batches.retrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_aretrieve_batch_parses_response(handler): + client = _async_client() + client.batches.retrieve.return_value = _sdk_response(_batch_dict()) + + result = await handler.aretrieve_batch( + retrieve_batch_data=RETRIEVE_DATA, client=client + ) + + client.batches.retrieve.assert_awaited_once_with(**RETRIEVE_DATA) + assert isinstance(result, LiteLLMBatch) + + +# =========================================================================== # +# cancel_batch (has an EXTRA sync-side isinstance guard the others lack) +# =========================================================================== # + + +def test_cancel_sync_dispatch_payload_and_result(handler): + client = _sync_client() + client.batches.cancel.return_value = _sdk_response(_batch_dict(status="cancelled")) + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.cancel_batch( + _is_async=False, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + + assert get_client.call_args.kwargs["_is_async"] is False + client.batches.cancel.assert_called_once_with(**CANCEL_DATA) + client.batches.create.assert_not_called() + client.batches.retrieve.assert_not_called() + assert isinstance(result, LiteLLMBatch) + assert result.status == "cancelled" + + +def test_cancel_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.cancel_batch( + _is_async=False, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + + +def test_cancel_sync_rejects_non_sync_client(handler): + """cancel_batch has a unique sync-side guard: if _is_async is False but the + resolved client is async (neither AzureOpenAI nor OpenAI), it must raise + rather than call .cancel().""" + async_client = _async_client() + + with patch(GET_CLIENT, return_value=async_client): + with pytest.raises(ValueError, match="sync client"): + handler.cancel_batch( + _is_async=False, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + + async_client.batches.cancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_cancel_async_returns_coroutine_and_awaits(handler): + client = _async_client() + client.batches.cancel.return_value = _sdk_response(_batch_dict(status="cancelled")) + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.cancel_batch( + _is_async=True, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.cancel.assert_awaited_once_with(**CANCEL_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.status == "cancelled" + + +@pytest.mark.asyncio +async def test_cancel_async_rejects_sync_client(handler): + sync_client = _sync_client() + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="async client"): + handler.cancel_batch( + _is_async=True, cancel_batch_data=CANCEL_DATA, **AUTH_KW + ) + sync_client.batches.cancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_acancel_batch_parses_response(handler): + client = _async_client() + client.batches.cancel.return_value = _sdk_response(_batch_dict(status="cancelled")) + + result = await handler.acancel_batch(cancel_batch_data=CANCEL_DATA, client=client) + + client.batches.cancel.assert_awaited_once_with(**CANCEL_DATA) + assert isinstance(result, LiteLLMBatch) + assert result.status == "cancelled" + + +# =========================================================================== # +# list_batches (returns the raw SDK response, NOT a LiteLLMBatch) +# =========================================================================== # + + +def test_list_sync_forwards_after_limit_and_returns_raw_response(handler): + client = _sync_client() + raw = MagicMock(name="raw_list_response") + client.batches.list.return_value = raw + + with patch(GET_CLIENT, return_value=client) as get_client: + result = handler.list_batches( + _is_async=False, after="cur-1", limit=20, **AUTH_KW + ) + + assert get_client.call_args.kwargs["_is_async"] is False + client.batches.list.assert_called_once_with(after="cur-1", limit=20) + # list returns the SDK response untouched (no LiteLLMBatch parsing). + assert result is raw + + +def test_list_sync_defaults_after_and_limit_to_none(handler): + client = _sync_client() + client.batches.list.return_value = MagicMock() + + with patch(GET_CLIENT, return_value=client): + handler.list_batches(_is_async=False, **AUTH_KW) + + client.batches.list.assert_called_once_with(after=None, limit=None) + + +def test_list_raises_when_client_is_none(handler): + with patch(GET_CLIENT, return_value=None): + with pytest.raises(ValueError, match="client is not initialized"): + handler.list_batches(_is_async=False, **AUTH_KW) + + +@pytest.mark.asyncio +async def test_list_async_returns_coroutine_and_awaits(handler): + client = _async_client() + raw = MagicMock(name="raw_async_list_response") + client.batches.list.return_value = raw + + with patch(GET_CLIENT, return_value=client) as get_client: + coro = handler.list_batches( + _is_async=True, after="cur-2", limit=7, **AUTH_KW + ) + assert asyncio.iscoroutine(coro) + result = await coro + + assert get_client.call_args.kwargs["_is_async"] is True + client.batches.list.assert_awaited_once_with(after="cur-2", limit=7) + assert result is raw + + +@pytest.mark.asyncio +async def test_list_async_rejects_sync_client(handler): + sync_client = _sync_client() + with patch(GET_CLIENT, return_value=sync_client): + with pytest.raises(ValueError, match="not an instance of AsyncOpenAI"): + handler.list_batches(_is_async=True, **AUTH_KW) + sync_client.batches.list.assert_not_called() + + +@pytest.mark.asyncio +async def test_alist_batches_returns_raw_response(handler): + client = _async_client() + raw = MagicMock(name="raw") + client.batches.list.return_value = raw + + result = await handler.alist_batches(client=client, after="a", limit=2) + + client.batches.list.assert_awaited_once_with(after="a", limit=2) + assert result is raw + + +# =========================================================================== # +# Cross-cutting: an OpenAI (non-Azure) client also satisfies the type guards, +# since the Union allows OpenAI / AsyncOpenAI (Azure-v1 path returns these). +# =========================================================================== # + + +def test_create_sync_accepts_plain_openai_client(handler): + client = MagicMock(spec=OpenAI) + client.batches.create.return_value = _sdk_response(_batch_dict()) + + with patch(GET_CLIENT, return_value=client): + result = handler.create_batch( + _is_async=False, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + assert isinstance(result, LiteLLMBatch) + + +@pytest.mark.asyncio +async def test_create_async_accepts_plain_async_openai_client(handler): + client = MagicMock(spec=AsyncOpenAI) + client.batches.create = AsyncMock(return_value=_sdk_response(_batch_dict())) + + with patch(GET_CLIENT, return_value=client): + result = await handler.create_batch( + _is_async=True, create_batch_data=CREATE_DATA, **AUTH_KW + ) + + assert isinstance(result, LiteLLMBatch) diff --git a/tests/test_litellm/llms/base_llm/__init__.py b/tests/test_litellm/llms/base_llm/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/base_llm/batches/__init__.py b/tests/test_litellm/llms/base_llm/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py b/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py new file mode 100644 index 00000000000..fd526c55de4 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/batches/base_batches_config_test.py @@ -0,0 +1,128 @@ +""" +Reusable contract test suite for BaseBatchesConfig implementations. + +Any provider whose batch transformation subclasses +`litellm.llms.base_llm.batches.transformation.BaseBatchesConfig` gets a shared +consistency net by subclassing `BatchesConfigContractTests` in its own +`test_transformation.py` (as a `Test*`-named class) and overriding the hooks +below. pytest then runs every contract test against that provider, guaranteeing +all provider batch transformations honour the same BaseBatchesConfig contract - +e.g. every `transform_retrieve_batch_response` returns a real `LiteLLMBatch` +with `object == "batch"`, a valid status, and an int `created_at`. + +This module is intentionally NOT named `test_*`: it holds no standalone tests +and must not be collected on its own. It mirrors the established repo pattern in +`tests/llm_translation/base_*_unit_tests.py`. + +Providers that do NOT implement BaseBatchesConfig (e.g. vertex_ai, whose +transformation is a standalone class with a different shape) cannot use this and +keep fully standalone tests. +""" + +import os +import sys +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.types.utils import LiteLLMBatch, LlmProviders + +# The OpenAI BatchJobStatus literal set - every provider must map into this. +VALID_BATCH_STATUSES = { + "validating", + "failed", + "in_progress", + "finalizing", + "completed", + "expired", + "cancelling", + "cancelled", +} + + +def make_raw_response(body: dict, status_code: int = 200) -> httpx.Response: + """Build an httpx.Response whose .json() yields `body` - the input shape the + transform_*_response methods consume.""" + return httpx.Response(status_code=status_code, json=body) + + +class BatchesConfigContractTests: + """Contract every BaseBatchesConfig implementation must satisfy. + + Subclass this with a `Test`-prefixed class and override the hooks. Do NOT + add a `Test` prefix here - this base must not be collected directly. + """ + + # ----------------------------------------------------------------------- # + # Hooks - providers MUST override these. + # ----------------------------------------------------------------------- # + + def make_config(self): + """Return a fresh instance of the provider's BaseBatchesConfig.""" + raise NotImplementedError("override make_config()") + + # The LlmProviders value this config reports. + expected_provider: LlmProviders = None # type: ignore[assignment] + + # Does this provider implement batch CREATE via the transformation? + # (anthropic raises NotImplementedError; bedrock/others may support it.) + supports_create: bool = False + + # Does this provider parse retrieve responses in the transformation layer? + supports_retrieve_response: bool = True + + def sample_retrieve_response_body(self) -> dict: + """A representative raw provider retrieve-batch response body.""" + raise NotImplementedError("override sample_retrieve_response_body()") + + # Expected mapped values for the sample above. + expected_retrieve_batch_id: str = None # type: ignore[assignment] + expected_retrieve_status: str = None # type: ignore[assignment] + + # ----------------------------------------------------------------------- # + # Contract tests - run for every provider subclass. + # ----------------------------------------------------------------------- # + + def test_contract__custom_llm_provider(self): + assert self.make_config().custom_llm_provider == self.expected_provider + + def test_contract__get_error_class_is_exception_with_status(self): + err = self.make_config().get_error_class( + error_message="boom", status_code=429, headers={} + ) + assert isinstance(err, Exception) + assert getattr(err, "status_code", None) == 429 + + def test_contract__create_unsupported_raises(self): + if self.supports_create: + pytest.skip("provider supports batch create; see provider-specific tests") + with pytest.raises(NotImplementedError): + self.make_config().transform_create_batch_request( + model="m", + create_batch_data={ + "input_file_id": "f", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + optional_params={}, + litellm_params={}, + ) + + def test_contract__retrieve_response_is_valid_litellm_batch(self): + if not self.supports_retrieve_response: + pytest.skip("provider handles retrieve outside the transformation layer") + out = self.make_config().transform_retrieve_batch_response( + model=None, + raw_response=make_raw_response(self.sample_retrieve_response_body()), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert isinstance(out, LiteLLMBatch) + assert out.object == "batch" + assert out.status in VALID_BATCH_STATUSES + assert isinstance(out.created_at, int) + assert out.id == self.expected_retrieve_batch_id + assert out.status == self.expected_retrieve_status diff --git a/tests/test_litellm/llms/base_llm/batches/test_transformation.py b/tests/test_litellm/llms/base_llm/batches/test_transformation.py index e69de29bb2d..cfb9f278f80 100644 --- a/tests/test_litellm/llms/base_llm/batches/test_transformation.py +++ b/tests/test_litellm/llms/base_llm/batches/test_transformation.py @@ -0,0 +1,231 @@ +""" +Unit tests for litellm/llms/base_llm/batches/transformation.py + +BaseBatchesConfig is the abstract base class that every provider-specific +batches config subclasses. It is almost entirely interface (abstractmethods + +one abstract property), so the only concrete behavior to regression-lock is: + + - the abstractness contract: the base class cannot be instantiated, and a + subclass missing any abstract member also cannot be instantiated; a + subclass implementing all of them can. + - get_config(): a classmethod that reflects over ``cls.__dict__`` and returns + the class-level config attributes, filtering out dunders, ``_abc`` internals, + callables (function/builtin/classmethod/staticmethod), and ``None`` values. + +These tests assert the exact dict get_config() produces for hand-built +subclasses, so a change to the filter predicate (e.g. dropping the ``None`` +filter, dropping the staticmethod/classmethod filter, or widening the prefix +filter to all single-underscore names) makes a test fail. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig +from litellm.types.utils import LlmProviders + + +# --------------------------------------------------------------------------- # +# A fully-concrete subclass: implements every abstract member with trivial +# bodies so it can be instantiated and so get_config() has a real cls to +# reflect over. Class-level attributes here are the get_config() fixtures. +# --------------------------------------------------------------------------- # + + +class _ConcreteBatchesConfig(BaseBatchesConfig): + string_attr = "hello" + int_attr = 42 + list_attr = [1, 2, 3] + none_attr = None + _single_underscore = "kept" + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.OPENAI + + def validate_environment( + self, + headers, + model, + messages, + optional_params, + litellm_params, + api_key=None, + api_base=None, + ) -> dict: + return headers + + def get_complete_batch_url( + self, api_base, api_key, model, optional_params, litellm_params, data + ) -> str: + return "https://example.com/batch" + + def transform_create_batch_request( + self, model, create_batch_data, optional_params, litellm_params + ): + return {"created": True} + + def transform_create_batch_response( + self, model, raw_response, logging_obj, litellm_params + ): + return raw_response + + def transform_retrieve_batch_request( + self, batch_id, optional_params, litellm_params + ): + return {"batch_id": batch_id} + + def transform_retrieve_batch_response( + self, model, raw_response, logging_obj, litellm_params + ): + return raw_response + + def get_error_class(self, error_message, status_code, headers): + return Exception(error_message) + + +# =========================================================================== # +# Abstractness contract +# =========================================================================== # + + +def test_base_class_cannot_be_instantiated(): + """The base class has unimplemented abstractmethods, so direct + instantiation must raise TypeError.""" + with pytest.raises(TypeError): + BaseBatchesConfig() + + +def test_fully_concrete_subclass_can_be_instantiated(): + instance = _ConcreteBatchesConfig() + assert isinstance(instance, BaseBatchesConfig) + + +@pytest.mark.parametrize( + "missing_member", + [ + "custom_llm_provider", + "validate_environment", + "get_complete_batch_url", + "transform_create_batch_request", + "transform_create_batch_response", + "transform_retrieve_batch_request", + "transform_retrieve_batch_response", + "get_error_class", + ], +) +def test_subclass_missing_any_abstract_member_cannot_instantiate(missing_member): + """Every abstract member is part of the contract: dropping any one of them + leaves the subclass abstract and uninstantiable.""" + namespace = { + k: v + for k, v in _ConcreteBatchesConfig.__dict__.items() + if not k.startswith("__") + } + namespace.pop(missing_member) + Incomplete = type("Incomplete", (BaseBatchesConfig,), namespace) + with pytest.raises(TypeError): + Incomplete() + + +def test_concrete_instance_methods_run(): + """Sanity: the trivial overrides actually execute through the base contract.""" + instance = _ConcreteBatchesConfig() + assert instance.custom_llm_provider == LlmProviders.OPENAI + assert instance.validate_environment( + headers={"x": "1"}, + model="m", + messages=[], + optional_params={}, + litellm_params={}, + ) == {"x": "1"} + assert instance.transform_retrieve_batch_request( + batch_id="b-1", optional_params={}, litellm_params={} + ) == {"batch_id": "b-1"} + + +# =========================================================================== # +# get_config() +# =========================================================================== # + + +def test_get_config_returns_class_level_non_none_data_attrs(): + """Exact contents: only class-level data attributes that are not None, + not dunders, not callables. Single-underscore names ARE kept (only ``__`` + and ``_abc`` prefixes are filtered). The ``custom_llm_provider`` property + object also survives the filter (a property is neither a function nor None), + matching how real provider subclasses define it.""" + config = _ConcreteBatchesConfig.get_config() + custom_llm_provider = config.pop("custom_llm_provider") + assert isinstance(custom_llm_provider, property) + assert config == { + "string_attr": "hello", + "int_attr": 42, + "list_attr": [1, 2, 3], + "_single_underscore": "kept", + } + + +def test_get_config_excludes_none_valued_attrs(): + assert "none_attr" not in _ConcreteBatchesConfig.get_config() + + +def test_get_config_excludes_methods_and_property(): + config = _ConcreteBatchesConfig.get_config() + for method_name in ( + "validate_environment", + "get_complete_batch_url", + "transform_create_batch_request", + "transform_create_batch_response", + "transform_retrieve_batch_request", + "transform_retrieve_batch_response", + "get_error_class", + "get_config", + ): + assert method_name not in config + + +def test_get_config_excludes_classmethod_and_staticmethod(): + """classmethod and staticmethod objects are filtered even though they are + not plain FunctionType.""" + + class WithCallables(_ConcreteBatchesConfig): + keep_me = "yes" + + @staticmethod + def a_static(): + return 1 + + @classmethod + def a_class(cls): + return 2 + + config = WithCallables.get_config() + assert config == {"keep_me": "yes"} + + +def test_get_config_only_reflects_own_dict_not_inherited(): + """get_config reflects cls.__dict__ only, so attributes defined on a parent + do not leak into a child's config.""" + + class Parent(_ConcreteBatchesConfig): + parent_attr = "parent" + + class Child(Parent): + child_attr = "child" + + assert Parent.get_config() == {"parent_attr": "parent"} + assert Child.get_config() == {"child_attr": "child"} + + +def test_get_config_on_base_class_exposes_only_the_abstract_property(): + """On the base class itself, the only ``__dict__`` member that survives the + filter is the ``custom_llm_provider`` property object (a property is neither + a function nor None and its name has no filtered prefix).""" + config = BaseBatchesConfig.get_config() + assert list(config.keys()) == ["custom_llm_provider"] + assert isinstance(config["custom_llm_provider"], property) diff --git a/tests/test_litellm/llms/bedrock/__init__.py b/tests/test_litellm/llms/bedrock/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/bedrock/batches/__init__.py b/tests/test_litellm/llms/bedrock/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/test_litellm/llms/bedrock/batches/test_transformation.py index e69de29bb2d..d1ad5943ae6 100644 --- a/tests/test_litellm/llms/bedrock/batches/test_transformation.py +++ b/tests/test_litellm/llms/bedrock/batches/test_transformation.py @@ -0,0 +1,684 @@ +""" +Regression tests for ``BedrockBatchesConfig`` (the BaseBatchesConfig +implementation for Bedrock model-invocation-job batches). + +This file complements (does not duplicate): + - ``test_batch_metadata_sanitization.py`` (covers + ``_get_openai_compatible_batch_metadata`` exhaustively) + - ``test_handler.py`` (covers the boto3-backed handler, not this transform) + +Here we lock the pure transform logic in ``transformation.py``: request +construction (S3 input/output config, model id, job name, role ARN), the +AWS-JobStatus -> OpenAI-status mapping, timestamp parsing, retrieve-request +URL/ARN handling, and the error class. AWS auth/sigv4 is the only external seam +we mock; everything else runs for real. +""" + +import os +import sys +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig +from litellm.types.utils import LiteLLMBatch, LlmProviders + +# AWS JobStatus -> OpenAI BatchJobStatus, exactly as encoded in transformation.py +# (both transform_create_batch_response and transform_retrieve_batch_response). +STATUS_MAP = { + "Submitted": "validating", + "Validating": "validating", + "Scheduled": "in_progress", + "InProgress": "in_progress", + "PartiallyCompleted": "completed", + "Completed": "completed", + "Failed": "failed", + "Stopping": "cancelling", + "Stopped": "cancelled", + "Expired": "expired", +} + +ARN = "arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/abc1234567" + + +@pytest.fixture +def config(): + return BedrockBatchesConfig() + + +def _raw(body: dict, status_code: int = 200) -> httpx.Response: + return httpx.Response(status_code=status_code, json=body) + + +# --------------------------------------------------------------------------- # +# get_complete_batch_url +# --------------------------------------------------------------------------- # + + +def test_get_complete_batch_url_uses_region(config): + url = config.get_complete_batch_url( + api_base=None, + api_key=None, + model="anthropic.claude-3", + optional_params={"aws_region_name": "eu-central-1"}, + litellm_params={}, + data={"input_file_id": "s3://b/k"}, + ) + assert url == "https://bedrock.eu-central-1.amazonaws.com/model-invocation-job" + + +# --------------------------------------------------------------------------- # +# transform_create_batch_request - request construction (sign_aws_request mocked) +# --------------------------------------------------------------------------- # + + +def test_create_request_builds_s3_input_output_and_arn(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-deadbeef", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({"Authorization": "signed"}, b'{"x": 1}') + result = config.transform_create_batch_request( + model="anthropic.claude-3-5-sonnet", + create_batch_data={ + "input_file_id": "s3://in-bucket/path/to/input.jsonl", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + optional_params={"aws_region_name": "us-west-2"}, + litellm_params={ + "s3_output_bucket_name": "out-bucket", + "aws_batch_role_arn": "arn:aws:iam::123:role/my-batch-role", + }, + ) + + bedrock_request = mock_sign.call_args.kwargs["data"] + assert bedrock_request["modelId"] == "anthropic.claude-3-5-sonnet" + assert bedrock_request["jobName"] == "litellm-batch-deadbeef" + assert bedrock_request["roleArn"] == "arn:aws:iam::123:role/my-batch-role" + assert ( + bedrock_request["inputDataConfig"]["s3InputDataConfig"]["s3Uri"] + == "s3://in-bucket/path/to/input.jsonl" + ) + assert ( + bedrock_request["outputDataConfig"]["s3OutputDataConfig"]["s3Uri"] + == "s3://out-bucket/litellm-batch-outputs/litellm-batch-deadbeef/" + ) + # 24h completion window -> 24 hour timeout + assert bedrock_request["timeoutDurationInHours"] == 24 + # signing was over the bedrock endpoint via POST + assert mock_sign.call_args.kwargs["service_name"] == "bedrock" + assert mock_sign.call_args.kwargs["method"] == "POST" + assert mock_sign.call_args.kwargs["endpoint_url"] == ( + "https://bedrock.us-west-2.amazonaws.com/model-invocation-job" + ) + # the transform returns the pre-signed envelope + assert result["method"] == "POST" + assert result["url"] == ( + "https://bedrock.us-west-2.amazonaws.com/model-invocation-job" + ) + assert result["headers"] == {"Authorization": "signed"} + + +def test_create_request_defaults_output_bucket_to_input_bucket(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-cafef00d", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://same-bucket/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + bedrock_request = mock_sign.call_args.kwargs["data"] + assert ( + bedrock_request["outputDataConfig"]["s3OutputDataConfig"]["s3Uri"] + == "s3://same-bucket/litellm-batch-outputs/litellm-batch-cafef00d/" + ) + + +def test_create_request_adds_kms_encryption_key_when_provided(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={ + "aws_batch_role_arn": "arn:aws:iam::1:role/r", + "s3_encryption_key_id": "kms-key-123", + }, + ) + s3out = mock_sign.call_args.kwargs["data"]["outputDataConfig"][ + "s3OutputDataConfig" + ] + assert s3out["s3EncryptionKeyId"] == "kms-key-123" + + +def test_create_request_omits_kms_key_when_absent(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign, patch( + "litellm.llms.bedrock.batches.transformation.get_secret_str", + return_value=None, + ): + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + s3out = mock_sign.call_args.kwargs["data"]["outputDataConfig"][ + "s3OutputDataConfig" + ] + assert "s3EncryptionKeyId" not in s3out + + +def test_create_request_missing_input_file_id_raises(config): + with pytest.raises(ValueError, match="input_file_id is required"): + config.transform_create_batch_request( + model="m", + create_batch_data={}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + + +def test_create_request_missing_role_arn_raises(config, monkeypatch): + monkeypatch.delenv("AWS_BATCH_ROLE_ARN", raising=False) + with pytest.raises(ValueError, match="IAM role ARN is required"): + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={}, + ) + + +def test_create_request_role_arn_from_env(config, monkeypatch): + monkeypatch.setenv("AWS_BATCH_ROLE_ARN", "arn:aws:iam::9:role/env-role") + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={}, + ) + assert ( + mock_sign.call_args.kwargs["data"]["roleArn"] + == "arn:aws:iam::9:role/env-role" + ) + + +def test_create_request_missing_model_raises(config): + with pytest.raises(ValueError, match="Could not determine Bedrock model ID"): + config.transform_create_batch_request( + model="", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + + +def test_create_request_no_timeout_for_non_24h_window(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={ + "input_file_id": "s3://b/in.jsonl", + "completion_window": "48h", + }, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + assert "timeoutDurationInHours" not in mock_sign.call_args.kwargs["data"] + + +# --------------------------------------------------------------------------- # +# transform_create_batch_response - status mapping + LiteLLMBatch shape +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize("bedrock_status,openai_status", list(STATUS_MAP.items())) +def test_create_response_status_mapping(config, bedrock_status, openai_status): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": bedrock_status}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == openai_status + assert out.id == ARN + assert out.object == "batch" + + +def test_create_response_unknown_status_falls_back_to_validating(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "SomeFutureStatus"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "validating" + + +def test_create_response_default_status_when_missing(config): + # status defaults to "Submitted" -> "validating" + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "validating" + + +def test_create_response_in_progress_sets_in_progress_at(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "InProgress"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "in_progress" + assert isinstance(out.in_progress_at, int) + + +def test_create_response_non_in_progress_leaves_in_progress_at_none(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Submitted"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.in_progress_at is None + + +def test_create_response_uses_original_request_fields(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Submitted"}), + logging_obj=MagicMock(), + litellm_params={ + "original_batch_request": { + "endpoint": "/v1/embeddings", + "input_file_id": "s3://b/in.jsonl", + "completion_window": "24h", + "metadata": {"user": "alice"}, + } + }, + ) + assert out.endpoint == "/v1/embeddings" + assert out.input_file_id == "s3://b/in.jsonl" + assert out.metadata == {"user": "alice"} + + +def test_create_response_default_endpoint_when_no_original_request(config): + out = config.transform_create_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Submitted"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.endpoint == "/v1/chat/completions" + assert out.completion_window == "24h" + + +def test_create_response_raises_on_unparseable_body(config): + bad = httpx.Response(status_code=200, text="not-json") + with pytest.raises(ValueError, match="Failed to parse Bedrock batch response"): + config.transform_create_batch_response( + model=None, + raw_response=bad, + logging_obj=MagicMock(), + litellm_params={}, + ) + + +# --------------------------------------------------------------------------- # +# transform_retrieve_batch_request - ARN validation + URL construction +# --------------------------------------------------------------------------- # + + +def test_retrieve_request_builds_encoded_arn_url(config): + with patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({"Authorization": "signed"}, b"") + result = config.transform_retrieve_batch_request( + batch_id=ARN, optional_params={}, litellm_params={} + ) + # ARN is URL-encoded (colons and slashes escaped) into the path + assert result["method"] == "GET" + assert result["data"] is None + assert result["headers"] == {"Authorization": "signed"} + assert result["url"].startswith( + "https://bedrock.us-west-2.amazonaws.com/model-invocation-job/" + ) + assert "%3A" in result["url"] # colon encoded + assert "%2F" in result["url"] # slash encoded + assert mock_sign.call_args.kwargs["method"] == "GET" + assert mock_sign.call_args.kwargs["data"] == {} + + +def test_retrieve_request_rejects_non_arn(config): + with pytest.raises(ValueError, match="Expected ARN"): + config.transform_retrieve_batch_request( + batch_id="abc1234567", optional_params={}, litellm_params={} + ) + + +def test_retrieve_request_rejects_short_arn(config): + with pytest.raises(ValueError, match="Invalid ARN format"): + config.transform_retrieve_batch_request( + batch_id="arn:aws:bedrock:us-west-2", optional_params={}, litellm_params={} + ) + + +def test_retrieve_request_rejects_bad_region(config): + bad = "arn:aws:bedrock:US_WEST:123:model-invocation-job/x" + with pytest.raises(ValueError, match="Invalid region in ARN"): + config.transform_retrieve_batch_request( + batch_id=bad, optional_params={}, litellm_params={} + ) + + +# --------------------------------------------------------------------------- # +# transform_retrieve_batch_response - status, timestamps, files, errors +# --------------------------------------------------------------------------- # + + +@pytest.mark.parametrize("bedrock_status,openai_status", list(STATUS_MAP.items())) +def test_retrieve_response_status_mapping(config, bedrock_status, openai_status): + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": bedrock_status}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == openai_status + + +def test_retrieve_response_unknown_status_falls_back_to_validating(config): + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "NewStatus"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.status == "validating" + + +def test_retrieve_response_extracts_file_configs(config): + body = { + "jobArn": ARN, + "status": "Completed", + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": "s3://b/in.jsonl"}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": "s3://b/out/"}}, + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.input_file_id == "s3://b/in.jsonl" + assert out.output_file_id == "s3://b/out/" + + +def test_retrieve_response_parses_timestamps_for_completed(config): + body = { + "jobArn": ARN, + "status": "Completed", + "submitTime": "2026-04-28T12:00:00Z", + "endTime": "2026-04-28T12:30:00Z", + "jobExpirationTime": "2026-05-28T12:00:00Z", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + expect_created = int( + datetime.datetime.fromisoformat("2026-04-28T12:00:00+00:00").timestamp() + ) + expect_completed = int( + datetime.datetime.fromisoformat("2026-04-28T12:30:00+00:00").timestamp() + ) + expect_expires = int( + datetime.datetime.fromisoformat("2026-05-28T12:00:00+00:00").timestamp() + ) + assert out.created_at == expect_created + assert out.completed_at == expect_completed + assert out.expires_at == expect_expires + # completed -> not failed/cancelled, no in_progress timestamp + assert out.failed_at is None + assert out.cancelled_at is None + assert out.in_progress_at is None + + +def test_retrieve_response_failed_sets_failed_at_from_end_time(config): + body = { + "jobArn": ARN, + "status": "Failed", + "endTime": "2026-04-28T12:30:00Z", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + assert out.failed_at == int( + datetime.datetime.fromisoformat("2026-04-28T12:30:00+00:00").timestamp() + ) + assert out.completed_at is None + assert out.cancelled_at is None + + +def test_retrieve_response_stopped_sets_cancelled_at(config): + body = {"jobArn": ARN, "status": "Stopped", "endTime": "2026-04-28T12:30:00Z"} + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + assert out.cancelled_at == int( + datetime.datetime.fromisoformat("2026-04-28T12:30:00+00:00").timestamp() + ) + assert out.completed_at is None + assert out.failed_at is None + + +def test_retrieve_response_in_progress_sets_in_progress_at_from_last_modified(config): + body = { + "jobArn": ARN, + "status": "InProgress", + "lastModifiedTime": "2026-04-28T12:15:00Z", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + import datetime + + assert out.in_progress_at == int( + datetime.datetime.fromisoformat("2026-04-28T12:15:00+00:00").timestamp() + ) + + +def test_retrieve_response_invalid_timestamp_becomes_none(config): + body = { + "jobArn": ARN, + "status": "Completed", + "submitTime": "not-a-timestamp", + "endTime": "also-bad", + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + # created_at falls back to int(time.time()) when submitTime unparseable + assert isinstance(out.created_at, int) + assert out.completed_at is None + + +def test_retrieve_response_builds_errors_from_message(config): + body = {"jobArn": ARN, "status": "Failed", "message": "validation failed"} + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body, status_code=400), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.errors is not None + assert out.errors.data[0].message == "validation failed" + assert out.errors.data[0].code == "400" + + +def test_retrieve_response_no_errors_when_no_message(config): + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw({"jobArn": ARN, "status": "Completed"}), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.errors is None + + +def test_retrieve_response_enriches_metadata(config): + body = { + "jobArn": ARN, + "status": "Completed", + "jobName": "litellm-batch-1", + "modelId": "anthropic.claude-3", + "roleArn": "arn:aws:iam::1:role/r", + "timeoutDurationInHours": 24, + "vpcConfig": {"subnetIds": ["subnet-1"]}, + "clientRequestToken": None, + } + out = config.transform_retrieve_batch_response( + model=None, + raw_response=_raw(body), + logging_obj=MagicMock(), + litellm_params={}, + ) + assert out.metadata["jobName"] == "litellm-batch-1" + assert out.metadata["modelId"] == "anthropic.claude-3" + assert out.metadata["roleArn"] == "arn:aws:iam::1:role/r" + # non-string scalar is stringified + assert out.metadata["timeoutDurationInHours"] == "24" + # dict/list serialized to JSON string + assert out.metadata["vpcConfig"] == '{"subnetIds": ["subnet-1"]}' + # None-valued fields dropped + assert "clientRequestToken" not in out.metadata + + +def test_retrieve_response_raises_on_unparseable_body(config): + bad = httpx.Response(status_code=200, text="<<>>") + with pytest.raises(ValueError, match="Failed to parse Bedrock batch response"): + config.transform_retrieve_batch_response( + model=None, + raw_response=bad, + logging_obj=MagicMock(), + litellm_params={}, + ) + + +# --------------------------------------------------------------------------- # +# get_error_class + custom_llm_provider +# --------------------------------------------------------------------------- # + + +def test_get_error_class_returns_bedrock_error(config): + err = config.get_error_class( + error_message="throttled", status_code=429, headers={} + ) + assert isinstance(err, Exception) + assert err.status_code == 429 + assert "throttled" in str(err) + + +def test_custom_llm_provider_is_bedrock(config): + assert config.custom_llm_provider == LlmProviders.BEDROCK + + +def test_validate_environment_passes_headers_through(config): + headers = {"X-Custom": "v"} + out = config.validate_environment( + headers=headers, + model="m", + messages=[], + optional_params={}, + litellm_params={}, + ) + assert out == headers + + +# --------------------------------------------------------------------------- # +# Shared BaseBatchesConfig contract suite. +# --------------------------------------------------------------------------- # + +from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402 + BatchesConfigContractTests, +) + + +class TestBedrockBatchesContract(BatchesConfigContractTests): + def make_config(self): + return BedrockBatchesConfig() + + expected_provider = LlmProviders.BEDROCK + # Bedrock builds a real create request (no NotImplementedError); the + # bedrock-specific create tests above cover the request/response shape. + supports_create = True + # Retrieve responses are parsed in this transformation layer + # (transform_retrieve_batch_response). + supports_retrieve_response = True + + def sample_retrieve_response_body(self) -> dict: + return { + "jobArn": ARN, + "status": "Completed", + "submitTime": "2026-04-28T12:00:00Z", + "endTime": "2026-04-28T12:30:00Z", + "inputDataConfig": {"s3InputDataConfig": {"s3Uri": "s3://b/in.jsonl"}}, + "outputDataConfig": {"s3OutputDataConfig": {"s3Uri": "s3://b/out/"}}, + } + + expected_retrieve_batch_id = ARN + expected_retrieve_status = "completed" diff --git a/tests/test_litellm/llms/vertex_ai/batches/__init__.py b/tests/test_litellm/llms/vertex_ai/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d 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 e69de29bb2d..cacea234777 100644 --- a/tests/test_litellm/llms/vertex_ai/batches/test_handler.py +++ b/tests/test_litellm/llms/vertex_ai/batches/test_handler.py @@ -0,0 +1,805 @@ +""" +Unit tests for ``VertexAIBatchPrediction`` (litellm/llms/vertex_ai/batches/handler.py). + +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). + +We mock only true I/O / auth seams: + * ``_ensure_access_token`` - the Vertex credential seam. Returns a fixed + (token, project) so we can assert the ``Authorization: Bearer `` + header is forwarded. + * ``_check_custom_proxy`` - returns ``(None, url)``; we let it pass the + computed default url straight through so we can assert the request URL. + * the httpx client factories (``_get_httpx_client`` / + ``get_async_httpx_client``) and the SSRF wrappers (``safe_get`` / + ``async_safe_get``) - the network calls. We assert which seam fired with + 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 +retrieve-after-cancel sequencing run for real. +""" + +from __future__ import annotations + +import asyncio +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.vertex_ai.batches.handler import ( # noqa: E402 + VertexAIBatchPrediction, +) +from litellm.types.utils import LiteLLMBatch # noqa: E402 + +HMOD = "litellm.llms.vertex_ai.batches.handler" +TOKEN = "ya29.fake-access-token" +PROJECT = "my-project" +LOCATION = "us-central1" +BATCH_ID = "3814889423749775360" + +CREATE_DATA = { + "input_file_id": ( + "gs://bucket/publishers/google/models/gemini-1.5-flash-001/file-uuid" + ) +} + + +def _vertex_job_response(state: str = "JOB_STATE_SUCCEEDED") -> dict: + return { + "name": f"projects/p/locations/{LOCATION}/batchPredictionJobs/{BATCH_ID}", + "state": state, + "createTime": "2024-12-04T21:53:12.120184Z", + "inputConfig": { + "instancesFormat": "jsonl", + "gcsSource": {"uris": ["gs://bucket/in.jsonl"]}, + }, + "outputInfo": {"gcsOutputDirectory": "gs://bucket/out"}, + } + + +def _http_response(status_code: int = 200, json_body: dict | None = None) -> MagicMock: + resp = MagicMock() + resp.status_code = status_code + resp.text = "error text" + resp.json.return_value = json_body if json_body is not None else _vertex_job_response() + return resp + + +def _make_handler() -> VertexAIBatchPrediction: + """Construct the handler with auth + proxy seams patched at the instance level. + + ``_ensure_access_token`` and ``_check_custom_proxy`` are inherited from + ``VertexLLM``; we patch them on the instance (DI-style) so the URL/auth + plumbing is deterministic and we can assert what got forwarded downstream. + """ + h = VertexAIBatchPrediction(gcs_bucket_name="litellm-testing-bucket") + h._ensure_access_token = MagicMock(return_value=(TOKEN, PROJECT)) # type: ignore[method-assign] + # pass the computed default url straight through (no custom proxy) + h._check_custom_proxy = MagicMock( # type: ignore[method-assign] + side_effect=lambda **kw: (None, kw["url"]) + ) + return h + + +def _run(coro): + return asyncio.run(coro) + + +# =========================================================================== # +# create_vertex_batch_url +# =========================================================================== # + + +def test_create_vertex_batch_url(): + h = _make_handler() + url = h.create_vertex_batch_url(vertex_location=LOCATION, vertex_project=PROJECT) + assert url == ( + f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{PROJECT}" + f"/locations/{LOCATION}/batchPredictionJobs" + ) + + +# =========================================================================== # +# create_batch +# =========================================================================== # + + +def test_create_batch_sync_posts_and_parses(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response() + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert isinstance(out, LiteLLMBatch) + assert out.id == BATCH_ID + assert out.status == "completed" + + # auth seam fired + h._ensure_access_token.assert_called_once() + # the POST hit the batchPredictionJobs collection url with bearer auth + _, kwargs = client.post.call_args + assert kwargs["url"].endswith(f"/projects/{PROJECT}/locations/{LOCATION}/batchPredictionJobs") + assert kwargs["headers"]["Authorization"] == f"Bearer {TOKEN}" + # body is the transformed vertex job (json-serialized) + sent = json.loads(kwargs["data"]) + assert sent["model"] == "publishers/google/models/gemini-1.5-flash-001" + assert sent["inputConfig"]["gcsSource"]["uris"] == [CREATE_DATA["input_file_id"]] + + +def test_create_batch_async_returns_coroutine_and_uses_async_client(): + h = _make_handler() + async_client = MagicMock() + async_client.post = AsyncMock(return_value=_http_response()) + sync_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + 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, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert isinstance(out, LiteLLMBatch) + assert out.id == BATCH_ID + async_client.post.assert_awaited_once() + # the async branch must NOT use the sync client for the request + sync_client.post.assert_not_called() + + +def test_create_batch_sync_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.create_batch( + _is_async=False, + create_batch_data=CREATE_DATA, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + +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(Exception, match="Error: 403"): + _run(coro) + + +# =========================================================================== # +# retrieve_batch +# =========================================================================== # + + +def test_retrieve_batch_sync_uses_safe_get_with_batch_id_url(): + h = _make_handler() + sync_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.safe_get", return_value=_http_response()) as safe_get, + ): + out = h.retrieve_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, + ) + + assert isinstance(out, LiteLLMBatch) + assert out.id == BATCH_ID + # SSRF-wrapped fetch fired with the batch-id-appended url + bearer header + args, kwargs = safe_get.call_args + assert args[0] is sync_client + assert args[1].endswith(f"/batchPredictionJobs/{BATCH_ID}") + assert kwargs["headers"]["Authorization"] == f"Bearer {TOKEN}" + # plain client.get must NOT be used (SSRF wrapper is the seam) + sync_client.get.assert_not_called() + + +def test_retrieve_batch_async_returns_coroutine_uses_async_safe_get(): + h = _make_handler() + async_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + patch( + f"{HMOD}.async_safe_get", + new=AsyncMock(return_value=_http_response()), + ) as async_safe_get, + ): + coro = h.retrieve_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, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert isinstance(out, LiteLLMBatch) + async_safe_get.assert_awaited_once() + args, _ = async_safe_get.await_args + assert args[1].endswith(f"/batchPredictionJobs/{BATCH_ID}") + + +def test_retrieve_batch_sync_non_200_raises(): + h = _make_handler() + with ( + 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"): + h.retrieve_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, + ) + + +def test_retrieve_batch_sync_invokes_logging_pre_call(): + """When a real ``Logging`` obj is passed, ``pre_call`` is invoked with the + request url + headers (the curl-redaction branch).""" + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + h = _make_handler() + logging_obj = MagicMock(spec=LiteLLMLogging) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.safe_get", return_value=_http_response()), + ): + h.retrieve_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, + logging_obj=logging_obj, + ) + + logging_obj.pre_call.assert_called_once() + _, kwargs = logging_obj.pre_call.call_args + assert kwargs["additional_args"]["api_base"].endswith( + f"/batchPredictionJobs/{BATCH_ID}" + ) + + +# =========================================================================== # +# list_batches +# =========================================================================== # + + +def _list_response() -> dict: + return { + "batchPredictionJobs": [_vertex_job_response()], + "nextPageToken": "next-tok", + } + + +def test_list_batches_sync_passes_pagination_params(): + h = _make_handler() + client = MagicMock() + client.get.return_value = _http_response(json_body=_list_response()) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.list_batches( + _is_async=False, + after="cursor-xyz", + limit=7, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert out["object"] == "list" + assert out["data"][0].id == BATCH_ID + assert out["has_more"] is True + assert out["next_page_token"] == "next-tok" + + _, kwargs = client.get.call_args + # limit -> pageSize (stringified), after -> pageToken + assert kwargs["params"] == {"pageSize": "7", "pageToken": "cursor-xyz"} + assert kwargs["headers"]["Authorization"] == f"Bearer {TOKEN}" + + +def test_list_batches_sync_omits_unset_pagination_params(): + h = _make_handler() + client = MagicMock() + client.get.return_value = _http_response(json_body={"batchPredictionJobs": []}) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.list_batches( + _is_async=False, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + _, kwargs = client.get.call_args + assert kwargs["params"] == {} + assert out["data"] == [] + assert out["has_more"] is False + + +def test_list_batches_async_returns_coroutine(): + h = _make_handler() + async_client = MagicMock() + async_client.get = AsyncMock( + return_value=_http_response(json_body=_list_response()) + ) + sync_client = MagicMock() + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.list_batches( + _is_async=True, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert out["data"][0].id == BATCH_ID + async_client.get.assert_awaited_once() + sync_client.get.assert_not_called() + + +def test_list_batches_sync_non_200_raises(): + h = _make_handler() + client = MagicMock() + 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"): + h.list_batches( + _is_async=False, + after=None, + limit=None, + api_base=None, + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + +# =========================================================================== # +# cancel_batch +# =========================================================================== # + + +def test_cancel_batch_sync_posts_cancel_then_retrieves(): + h = _make_handler() + client = MagicMock() + client.post.return_value = _http_response(json_body={}) + client.get.return_value = _http_response( + json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") + ) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = 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, + ) + + assert isinstance(out, LiteLLMBatch) + assert out.status == "cancelled" + + # POST hit the :cancel url + _, post_kwargs = client.post.call_args + assert post_kwargs["url"].endswith(f"/batchPredictionJobs/{BATCH_ID}:cancel") + assert post_kwargs["data"] == json.dumps({}) + # then GET hit the plain retrieve url (no :cancel suffix) + _, get_kwargs = client.get.call_args + assert get_kwargs["url"].endswith(f"/batchPredictionJobs/{BATCH_ID}") + assert not get_kwargs["url"].endswith(":cancel") + + +def test_cancel_batch_async_returns_coroutine_posts_then_retrieves(): + h = _make_handler() + async_client = MagicMock() + async_client.post = AsyncMock(return_value=_http_response(json_body={})) + async_client.get = AsyncMock( + return_value=_http_response( + json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") + ) + ) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + 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, + ) + assert asyncio.iscoroutine(coro) + out = _run(coro) + + assert out.status == "cancelled" + async_client.post.assert_awaited_once() + async_client.get.assert_awaited_once() + _, post_kwargs = async_client.post.await_args + 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() + client.post.return_value = _http_response(json_body={}) + 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"): + 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, + ) + + +def test_cancel_batch_sync_proxy_url_without_cancel_suffix_uses_rsplit_branch(): + """If ``_check_custom_proxy`` hands back a url that does NOT end in + ``:cancel`` (e.g. a custom proxy rewrote it), the retrieve url is derived + via the ``rsplit(':cancel')`` else-branch rather than ``removesuffix``.""" + h = _make_handler() + # override the proxy seam to return a non-:cancel-suffixed url + h._check_custom_proxy = MagicMock( # type: ignore[method-assign] + return_value=(None, "https://proxy.internal/vertex/batch") + ) + client = MagicMock() + client.post.return_value = _http_response(json_body={}) + client.get.return_value = _http_response( + json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") + ) + + with patch(f"{HMOD}._get_httpx_client", return_value=client): + out = h.cancel_batch( + _is_async=False, + batch_id=BATCH_ID, + api_base="https://proxy.internal", + vertex_credentials=None, + vertex_project=PROJECT, + vertex_location=LOCATION, + timeout=600.0, + max_retries=None, + ) + + assert out.status == "cancelled" + _, get_kwargs = client.get.call_args + # rsplit(":cancel")[0].rstrip("/") of a url with no :cancel -> url unchanged + assert get_kwargs["url"] == "https://proxy.internal/vertex/batch" + + +def test_cancel_batch_sync_httpstatuserror_logged_and_reraised(): + """The cancel POST ``httpx.HTTPStatusError`` except-branch logs + re-raises.""" + h = _make_handler() + client = MagicMock() + request = httpx.Request("POST", "https://x/batchPredictionJobs/1:cancel") + err_response = httpx.Response(status_code=502, request=request, text="bad gw") + 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(httpx.HTTPStatusError): + 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, + ) + client.get.assert_not_called() + + +def test_create_batch_async_httpstatuserror_logged_and_reraised(): + h = _make_handler() + async_client = MagicMock() + request = httpx.Request("POST", "https://x/batchPredictionJobs") + err_response = httpx.Response(status_code=500, request=request, text="boom") + async_client.post = AsyncMock( + side_effect=httpx.HTTPStatusError( + "boom", request=request, response=err_response + ) + ) + + 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(httpx.HTTPStatusError): + _run(coro) + + +def test_async_retrieve_batch_non_200_raises(): + h = _make_handler() + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=MagicMock()), + patch( + f"{HMOD}.async_safe_get", + new=AsyncMock(return_value=_http_response(status_code=500)), + ), + ): + coro = h.retrieve_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) + + +def test_async_retrieve_batch_invokes_logging_pre_call(): + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging + + h = _make_handler() + logging_obj = MagicMock(spec=LiteLLMLogging) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=MagicMock()), + patch( + f"{HMOD}.async_safe_get", + new=AsyncMock(return_value=_http_response()), + ), + ): + coro = h.retrieve_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, + logging_obj=logging_obj, + ) + _run(coro) + + logging_obj.pre_call.assert_called_once() + + +def test_async_list_batches_non_200_raises(): + h = _make_handler() + async_client = MagicMock() + async_client.get = AsyncMock(return_value=_http_response(status_code=500)) + + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + coro = h.list_batches( + _is_async=True, + after=None, + limit=None, + 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) + + +def test_async_cancel_batch_httpstatuserror_and_retrieve_non_200(): + """Async cancel: POST HTTPStatusError re-raises; and separately the + retrieve-after-cancel non-200 raises.""" + h = _make_handler() + + # (a) POST raises HTTPStatusError + async_client = MagicMock() + request = httpx.Request("POST", "https://x/batchPredictionJobs/1:cancel") + err_response = httpx.Response(status_code=502, request=request, text="bad") + async_client.post = AsyncMock( + side_effect=httpx.HTTPStatusError("boom", request=request, response=err_response) + ) + async_client.get = AsyncMock() + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), + ): + 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(httpx.HTTPStatusError): + _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={})) + async_client2.get = AsyncMock(return_value=_http_response(status_code=404)) + with ( + patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_async_httpx_client", return_value=async_client2), + ): + 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: 404"): + _run(coro) diff --git a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py index e69de29bb2d..37084d43441 100644 --- a/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/batches/test_transformation.py @@ -0,0 +1,396 @@ +""" +Unit tests for ``VertexAIBatchTransformation`` +(litellm/llms/vertex_ai/batches/transformation.py). + +This module is pure transformation logic: it maps OpenAI-shaped batch requests +into Vertex AI ``VertexAIBatchPredictionJob`` payloads, and maps Vertex AI batch +responses back into ``LiteLLMBatch`` / OpenAI list shapes. Unlike anthropic / +bedrock, this class does NOT subclass ``BaseBatchesConfig`` - it's a standalone +set of classmethods with a Vertex-specific shape, so these tests are fully +standalone and assert exact values rather than "ran without error". + +There are no real I/O seams here; ``uuid.uuid4`` is the only nondeterministic +dependency and is patched where the displayName is asserted. +""" + +import os +import sys +from unittest.mock import patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402 + VertexAIBatchTransformation, +) +from litellm.llms.vertex_ai.common_utils import ( # noqa: E402 + _convert_vertex_datetime_to_openai_datetime, +) +from litellm.types.utils import LiteLLMBatch # noqa: E402 + +T = VertexAIBatchTransformation + +INPUT_FILE = ( + "gs://litellm-testing-bucket/litellm-vertex-files/publishers/google/" + "models/gemini-1.5-flash-001/e9412502-2c91-42a6-8e61-f5c294cc0fc8" +) + + +# =========================================================================== # +# transform_openai_batch_request_to_vertex_ai_batch_request +# =========================================================================== # + + +def test_transform_openai_request_builds_full_vertex_job(): + with patch( + "litellm.llms.vertex_ai.batches.transformation.uuid.uuid4", + return_value="fixed-uuid", + ): + job = T.transform_openai_batch_request_to_vertex_ai_batch_request( + {"input_file_id": INPUT_FILE} + ) + + assert job["displayName"] == "litellm-vertex-batch-fixed-uuid" + assert job["model"] == "publishers/google/models/gemini-1.5-flash-001" + + assert job["inputConfig"]["instancesFormat"] == "jsonl" + assert job["inputConfig"]["gcsSource"]["uris"] == [INPUT_FILE] + + assert job["outputConfig"]["predictionsFormat"] == "jsonl" + # gcs uri prefix == file path with the filename stripped + assert ( + job["outputConfig"]["gcsDestination"]["outputUriPrefix"] + == "gs://litellm-testing-bucket/litellm-vertex-files/publishers/google/" + "models/gemini-1.5-flash-001" + ) + + +def test_transform_openai_request_missing_input_file_id_raises(): + with pytest.raises(ValueError, match="input_file_id is required"): + T.transform_openai_batch_request_to_vertex_ai_batch_request({}) + + +# =========================================================================== # +# transform_vertex_ai_batch_response_to_openai_batch_response +# =========================================================================== # + + +def test_transform_vertex_response_full_mapping(): + response = { + "name": "projects/510528649030/locations/us-central1/batchPredictionJobs/3814889423749775360", + "state": "JOB_STATE_SUCCEEDED", + "createTime": "2024-12-04T21:53:12.120184Z", + "inputConfig": { + "instancesFormat": "jsonl", + "gcsSource": {"uris": ["gs://bucket/in.jsonl"]}, + }, + "outputInfo": {"gcsOutputDirectory": "gs://bucket/out"}, + } + batch = T.transform_vertex_ai_batch_response_to_openai_batch_response(response) + + assert isinstance(batch, LiteLLMBatch) + assert batch.id == "3814889423749775360" + assert batch.completion_window == "24hrs" + # created_at is parsed via the shared helper (uses local tz); assert the + # transform forwards createTime through that helper rather than a hardcoded + # epoch that would be tz-dependent + assert batch.created_at == _convert_vertex_datetime_to_openai_datetime( + "2024-12-04T21:53:12.120184Z" + ) + assert batch.endpoint == "" + assert batch.object == "batch" + assert batch.input_file_id == "gs://bucket/in.jsonl" + assert batch.status == "completed" + assert batch.error_file_id is None + assert batch.output_file_id == "gs://bucket/out/predictions.jsonl" + + +def test_transform_vertex_response_error_file_id_always_none(): + batch = T.transform_vertex_ai_batch_response_to_openai_batch_response( + { + "name": "x/y/123", + "state": "JOB_STATE_FAILED", + "createTime": "2024-12-04T21:53:12.120184Z", + } + ) + assert batch.error_file_id is None + + +# =========================================================================== # +# _get_batch_job_status_from_vertex_ai_batch_response (test EVERY entry) +# =========================================================================== # + + +@pytest.mark.parametrize( + "vertex_state,expected", + [ + ("JOB_STATE_UNSPECIFIED", "failed"), + ("JOB_STATE_QUEUED", "validating"), + ("JOB_STATE_PENDING", "validating"), + ("JOB_STATE_RUNNING", "in_progress"), + ("JOB_STATE_SUCCEEDED", "completed"), + ("JOB_STATE_FAILED", "failed"), + ("JOB_STATE_CANCELLING", "cancelling"), + ("JOB_STATE_CANCELLED", "cancelled"), + ("JOB_STATE_PAUSED", "in_progress"), + ("JOB_STATE_EXPIRED", "expired"), + ("JOB_STATE_UPDATING", "in_progress"), + ("JOB_STATE_PARTIALLY_SUCCEEDED", "completed"), + ], +) +def test_status_mapping_every_entry(vertex_state, expected): + assert ( + T._get_batch_job_status_from_vertex_ai_batch_response({"state": vertex_state}) + == expected + ) + + +def test_status_mapping_defaults_to_unspecified_when_missing(): + # No "state" key -> defaults to JOB_STATE_UNSPECIFIED -> "failed" + assert T._get_batch_job_status_from_vertex_ai_batch_response({}) == "failed" + + +def test_status_mapping_unknown_state_raises_keyerror(): + with pytest.raises(KeyError): + T._get_batch_job_status_from_vertex_ai_batch_response({"state": "NOPE"}) + + +# =========================================================================== # +# _get_batch_id_from_vertex_ai_batch_response +# =========================================================================== # + + +def test_get_batch_id_splits_path(): + assert ( + T._get_batch_id_from_vertex_ai_batch_response( + {"name": "projects/p/locations/l/batchPredictionJobs/999"} + ) + == "999" + ) + + +def test_get_batch_id_no_slash_returns_name(): + assert T._get_batch_id_from_vertex_ai_batch_response({"name": "abc"}) == "abc" + + +def test_get_batch_id_empty_name_returns_empty(): + assert T._get_batch_id_from_vertex_ai_batch_response({"name": ""}) == "" + assert T._get_batch_id_from_vertex_ai_batch_response({}) == "" + + +# =========================================================================== # +# _get_input_file_id_from_vertex_ai_batch_response +# =========================================================================== # + + +def test_get_input_file_id_happy_path(): + assert ( + T._get_input_file_id_from_vertex_ai_batch_response( + {"inputConfig": {"gcsSource": {"uris": ["gs://b/a.jsonl", "gs://b/c.jsonl"]}}} + ) + == "gs://b/a.jsonl" + ) + + +def test_get_input_file_id_missing_input_config(): + assert T._get_input_file_id_from_vertex_ai_batch_response({}) == "" + + +def test_get_input_file_id_missing_gcs_source(): + assert ( + T._get_input_file_id_from_vertex_ai_batch_response({"inputConfig": {}}) == "" + ) + + +def test_get_input_file_id_empty_uris(): + assert ( + T._get_input_file_id_from_vertex_ai_batch_response( + {"inputConfig": {"gcsSource": {"uris": []}}} + ) + == "" + ) + + +# =========================================================================== # +# _get_output_file_id_from_vertex_ai_batch_response +# =========================================================================== # + + +def test_get_output_file_id_from_output_info(): + # outputInfo branch: rstrip trailing slash, append predictions.jsonl + assert ( + T._get_output_file_id_from_vertex_ai_batch_response( + {"outputInfo": {"gcsOutputDirectory": "gs://bucket/out/"}} + ) + == "gs://bucket/out/predictions.jsonl" + ) + + +def test_get_output_file_id_output_info_no_trailing_slash(): + assert ( + T._get_output_file_id_from_vertex_ai_batch_response( + {"outputInfo": {"gcsOutputDirectory": "gs://bucket/out"}} + ) + == "gs://bucket/out/predictions.jsonl" + ) + + +def test_get_output_file_id_empty_output_info_falls_through_to_output_config(): + # gcsOutputDirectory missing -> "" -> the "/predictions.jsonl" guard skips + # the outputInfo branch, falls through to outputConfig + resp = { + "outputInfo": {}, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}}, + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://b/cfg/predictions.jsonl" + ) + + +def test_get_output_file_id_no_output_info_and_no_output_config(): + assert T._get_output_file_id_from_vertex_ai_batch_response({}) == "" + + +def test_get_output_file_id_output_config_missing_gcs_destination(): + # outputConfig present but no gcsDestination -> returns the running "" value + assert ( + T._get_output_file_id_from_vertex_ai_batch_response({"outputConfig": {}}) == "" + ) + + +def test_get_output_file_id_output_config_already_has_suffix(): + # outputUriPrefix already ends in /predictions.jsonl -> returned as-is (no double append) + resp = { + "outputConfig": { + "gcsDestination": {"outputUriPrefix": "gs://b/cfg/predictions.jsonl"} + } + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://b/cfg/predictions.jsonl" + ) + + +def test_get_output_file_id_output_config_strips_trailing_slash(): + resp = { + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg/"}} + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://b/cfg/predictions.jsonl" + ) + + +def test_get_output_file_id_output_info_takes_precedence_over_output_config(): + resp = { + "outputInfo": {"gcsOutputDirectory": "gs://from-info"}, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://from-config"}}, + } + assert ( + T._get_output_file_id_from_vertex_ai_batch_response(resp) + == "gs://from-info/predictions.jsonl" + ) + + +# =========================================================================== # +# _get_gcs_uri_prefix_from_file +# =========================================================================== # + + +def test_get_gcs_uri_prefix_root(): + assert ( + T._get_gcs_uri_prefix_from_file("gs://litellm-testing-bucket/vtx_batch.jsonl") + == "gs://litellm-testing-bucket" + ) + + +def test_get_gcs_uri_prefix_nested(): + assert ( + T._get_gcs_uri_prefix_from_file( + "gs://litellm-testing-bucket/batches/vtx_batch.jsonl" + ) + == "gs://litellm-testing-bucket/batches" + ) + + +# =========================================================================== # +# _get_model_from_gcs_file +# =========================================================================== # + + +def test_get_model_from_gcs_file_plain(): + assert ( + T._get_model_from_gcs_file(INPUT_FILE) + == "publishers/google/models/gemini-1.5-flash-001" + ) + + +def test_get_model_from_gcs_file_url_encoded(): + # %2F decodes to "/" via urllib.unquote before splitting + encoded = ( + "gs://bucket/publishers%2Fgoogle%2Fmodels%2Fgemini-1.5-flash-001%2Fuuid" + ) + 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): + T._get_model_from_gcs_file("gs://bucket/no-model-here.jsonl") + + +# =========================================================================== # +# transform_vertex_ai_batch_list_response_to_openai_list_response +# =========================================================================== # + + +def _job(batch_id: str) -> dict: + return { + "name": f"projects/p/locations/l/batchPredictionJobs/{batch_id}", + "state": "JOB_STATE_SUCCEEDED", + "createTime": "2024-12-04T21:53:12.120184Z", + } + + +def test_list_response_multiple_jobs(): + response = { + "batchPredictionJobs": [_job("111"), _job("222"), _job("333")], + "nextPageToken": "tok-abc", + } + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response(response) + + assert out["object"] == "list" + assert [b.id for b in out["data"]] == ["111", "222", "333"] + assert out["first_id"] == "111" + assert out["last_id"] == "333" + assert out["has_more"] is True + assert out["next_page_token"] == "tok-abc" + + +def test_list_response_no_next_page_token(): + response = {"batchPredictionJobs": [_job("111")]} + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response(response) + assert out["has_more"] is False + assert out["next_page_token"] is None + assert out["first_id"] == "111" + assert out["last_id"] == "111" + + +def test_list_response_empty(): + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response({}) + assert out["data"] == [] + assert out["first_id"] is None + assert out["last_id"] is None + assert out["has_more"] is False + + +def test_list_response_none_jobs_treated_as_empty(): + out = T.transform_vertex_ai_batch_list_response_to_openai_list_response( + {"batchPredictionJobs": None} + ) + assert out["data"] == [] + assert out["first_id"] is None