From 2cf565ae282584a74d64d92bd2d59fe9a12a5484 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 29 Jun 2026 09:22:58 +0530 Subject: [PATCH] test(batches): add 1:1 test file scaffold for batches component paths (#30529) * test(batches): add 1:1 test file scaffold for batches component paths Co-authored-by: Cursor * Add harness test for create batch endpoint * Add retrieve endpoint harness tests * Add list endpoint harness tests * Add cancel endpoint harness tests * Add cancel endpoint harness tests * Add test for litellm/batches/main.py * Add test for litellm/tests/test_litellm/batches/test_batch_utils.py * Add handler and transformation tests for all providers * Fix: run batches tests in cicd * fix(tests): remove azure/__init__.py that shadowed azure namespace package Adding __init__.py to tests/test_litellm/llms/azure/ caused pytest to insert tests/test_litellm/llms/ into sys.path[0], making our empty azure/ dir shadow the real azure-identity namespace package. Any test that patched azure.identity.* would then fail with AttributeError. * style(tests): apply ruff format to test_batch_utils.py Base migrated the formatter from black to ruff format (#31317); reformat the batches scaffold test file to match. --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .github/workflows/test-unit-misc.yml | 1 + .../workflows/test-unit-proxy-endpoints.yml | 1 + tests/test_litellm/batches/__init__.py | 0 .../test_litellm/batches/test_batch_utils.py | 735 ++++++ tests/test_litellm/batches/test_main.py | 744 ++++++ 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 ++++++ .../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 ++++ .../proxy/batches_endpoints/__init__.py | 0 .../proxy/batches_endpoints/test_endpoints.py | 2026 +++++++++++++++++ 23 files changed, 7178 insertions(+) create mode 100644 tests/test_litellm/batches/__init__.py create mode 100644 tests/test_litellm/batches/test_batch_utils.py create mode 100644 tests/test_litellm/batches/test_main.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/anthropic/batches/test_handler.py create mode 100644 tests/test_litellm/llms/anthropic/batches/test_transformation.py create mode 100644 tests/test_litellm/llms/azure/batches/__init__.py create mode 100644 tests/test_litellm/llms/azure/batches/test_handler.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/base_llm/batches/test_transformation.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/bedrock/batches/test_transformation.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/__init__.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/test_handler.py create mode 100644 tests/test_litellm/llms/vertex_ai/batches/test_transformation.py create mode 100644 tests/test_litellm/proxy/batches_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/batches_endpoints/test_endpoints.py diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index d411d996d8e..7c3b195f0ad 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -22,6 +22,7 @@ jobs: uses: ./.github/workflows/_test-unit-base.yml with: test-path: >- + tests/test_litellm/batches tests/test_litellm/secret_managers tests/test_litellm/a2a_protocol tests/test_litellm/anthropic_interface diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index 23ca7a8f2f9..cbb36eebdb9 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -31,6 +31,7 @@ jobs: tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/openai_files_endpoint + tests/test_litellm/proxy/batches_endpoints tests/test_litellm/proxy/video_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/image_endpoints 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/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py new file mode 100644 index 00000000000..3aebfcb911e --- /dev/null +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -0,0 +1,735 @@ +""" +Unit tests for litellm/batches/batch_utils.py + +batch_utils.py is the batch cost/usage/parsing layer: it turns a batch output +JSONL into spend (cost), token usage, and the list of models seen, and counts +tokens in batch *input* files for rate limiting. A silent bug here mis-bills +real money or lets callers slip past TPM limits, so these tests assert exact +numeric results rather than "ran without error". + +Pure functions (parsing, token math, credential extraction, success checks) run +for real with exact-value assertions. The few true external seams - the cost +maps (litellm.completion_cost, batch_cost_calculator), the tokenizer +(token_counter), and remote file fetch (afile_content) - are mocked with +deterministic stand-ins so the arithmetic under test is the only variable. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +import litellm.batches.batch_utils as bu +from litellm.types.utils import Usage + +# --------------------------------------------------------------------------- # +# Builders for batch OUTPUT file rows. +# Shape: {"response": {"status_code": 200, "body": {... "usage": {...}}}} +# --------------------------------------------------------------------------- # + + +def _usage(p, c, t=None): + return { + "prompt_tokens": p, + "completion_tokens": c, + "total_tokens": t if t is not None else p + c, + } + + +def _success_row(model="gpt-4o", usage=None, **body_extra): + body = {"model": model, **body_extra} + if usage is not None: + body["usage"] = usage + return {"response": {"status_code": 200, "body": body}} + + +def _failed_row(status_code=500, model="gpt-4o"): + return {"response": {"status_code": status_code, "body": {"model": model}}} + + +# =========================================================================== # +# _batch_response_was_successful +# =========================================================================== # + + +@pytest.mark.parametrize( + "row,expected", + [ + ({"response": {"status_code": 200}}, True), + ({"response": {"status_code": 500}}, False), + ({"response": {"status_code": 429}}, False), + ({"response": {}}, False), # no status_code + ({}, False), # no response + ({"response": None}, False), # null response + ], +) +def test_batch_response_was_successful(row, expected): + assert bu._batch_response_was_successful(row) is expected + + +# =========================================================================== # +# _get_response_from_batch_job_output_file +# =========================================================================== # + + +def test_get_response_body_present(): + row = {"response": {"body": {"model": "gpt-4o", "usage": {"x": 1}}}} + assert bu._get_response_from_batch_job_output_file(row) == { + "model": "gpt-4o", + "usage": {"x": 1}, + } + + +@pytest.mark.parametrize( + "row", + [ + {}, # no response + {"response": {}}, # no body + {"response": None}, # null response + {"response": {"body": None}}, # null body + ], +) +def test_get_response_body_missing_returns_empty(row): + assert bu._get_response_from_batch_job_output_file(row) == {} + + +# =========================================================================== # +# _get_batch_job_usage_from_response_body +# =========================================================================== # + + +def test_get_usage_from_response_body(): + usage = bu._get_batch_job_usage_from_response_body( + {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}} + ) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 10, + 5, + 15, + ) + + +def test_get_usage_from_response_body_missing_is_zero(): + usage = bu._get_batch_job_usage_from_response_body({}) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 0, + 0, + 0, + ) + + +# =========================================================================== # +# _get_file_content_as_dictionary (JSONL parsing) +# =========================================================================== # + + +def test_parse_jsonl_multiple_lines(): + content = b'{"a": 1}\n{"b": 2}\n{"c": 3}' + assert bu._get_file_content_as_dictionary(content) == [ + {"a": 1}, + {"b": 2}, + {"c": 3}, + ] + + +def test_parse_jsonl_trailing_newline_skipped(): + # outer content is stripped; the trailing-newline empty line is dropped. + content = b'{"a": 1}\n{"b": 2}\n' + assert bu._get_file_content_as_dictionary(content) == [{"a": 1}, {"b": 2}] + + +def test_parse_jsonl_empty_content_is_empty_list(): + assert bu._get_file_content_as_dictionary(b"") == [] + + +def test_parse_jsonl_malformed_raises(): + with pytest.raises(Exception): + bu._get_file_content_as_dictionary(b"not valid json") + + +# =========================================================================== # +# _iter_batch_input_lines / _iter_batch_input_entries (JSONL parsing) +# =========================================================================== # + + +def test_iter_input_lines_skips_blank_and_strips(): + content = b'{"a":1}\n\n \n{"b":2}\n' + assert list(bu._iter_batch_input_lines(content)) == [b'{"a":1}', b'{"b":2}'] + + +def test_iter_input_lines_handles_missing_trailing_newline(): + assert list(bu._iter_batch_input_lines(b'{"a":1}')) == [b'{"a":1}'] + + +def test_iter_input_lines_empty(): + assert list(bu._iter_batch_input_lines(b"")) == [] + + +def test_iter_input_entries_parses_each_row(): + content = b'{"body": {"model": "gpt-4o"}}\n{"body": {"model": "claude-3"}}\n' + assert list(bu._iter_batch_input_entries(content)) == [ + {"body": {"model": "gpt-4o"}}, + {"body": {"model": "claude-3"}}, + ] + + +def test_iter_input_entries_raises_on_malformed_line(): + # _iter_batch_input_entries raises on a bad row; callers that must survive + # bad rows iterate _iter_batch_input_lines and parse per-row instead. + with pytest.raises(Exception): + list(bu._iter_batch_input_entries(b'{"ok":1}\nnot-json\n')) + + +# =========================================================================== # +# _estimate_batch_entry_tokens (regression: an uncountable/malformed row must +# never contribute zero tokens, or a crafted batch could evade the TPM limit) +# =========================================================================== # + + +def test_estimate_tokens_scales_with_size(): + # 4 bytes per token, floored, with a minimum of 1. + assert bu._estimate_batch_entry_tokens(b"a" * 40) == 10 + + +def test_estimate_tokens_never_zero_for_short_rows(): + assert bu._estimate_batch_entry_tokens(b"") == 1 + assert bu._estimate_batch_entry_tokens(b"abc") == 1 + + +# =========================================================================== # +# _get_batch_models_from_file_content (output file) +# =========================================================================== # + + +def test_output_models_uses_model_name_override(): + # model_name short-circuits: content is ignored entirely. + assert bu._get_batch_models_from_file_content([_success_row(model="ignored")], model_name="forced-model") == [ + "forced-model" + ] + + +def test_output_models_collects_from_successful_only(): + rows = [ + _success_row(model="gpt-4o"), + _failed_row(model="should-be-skipped"), + _success_row(model="claude-3"), + ] + assert bu._get_batch_models_from_file_content(rows) == ["gpt-4o", "claude-3"] + + +def test_output_models_skips_successful_without_model(): + rows = [{"response": {"status_code": 200, "body": {}}}] + assert bu._get_batch_models_from_file_content(rows) == [] + + +# =========================================================================== # +# _extract_file_access_credentials +# =========================================================================== # + + +def test_extract_credentials_only_known_keys(): + params = { + "api_key": "sk-1", + "api_base": "https://b", + "vertex_project": "proj", + "model": "gpt-4o", # not a credential key + "unrelated": "x", + } + assert bu._extract_file_access_credentials(params) == { + "api_key": "sk-1", + "api_base": "https://b", + "vertex_project": "proj", + } + + +@pytest.mark.parametrize("params", [None, {}]) +def test_extract_credentials_empty(params): + assert bu._extract_file_access_credentials(params) == {} + + +def test_extract_credentials_all_supported_keys(): + keys = { + "api_key", + "api_base", + "api_version", + "organization", + "azure_ad_token", + "azure_ad_token_provider", + "vertex_project", + "vertex_location", + "vertex_credentials", + "timeout", + "max_retries", + } + params = {k: f"val-{k}" for k in keys} + assert bu._extract_file_access_credentials(params) == params + + +# =========================================================================== # +# _count_prompt_or_input_tokens (regression-critical: list[list[int]] used to +# count as zero and let callers slip past TPM limits). token_counter stubbed to +# len(text) so every shape has an exact expected value. +# =========================================================================== # + + +@pytest.fixture +def fake_token_counter(monkeypatch): + def _tc(model=None, text=None, messages=None, **kw): + if messages is not None: + return len(messages) + if text is not None: + return len(text) + return 0 + + monkeypatch.setattr(bu, "token_counter", _tc) + return _tc + + +def test_count_tokens_str(fake_token_counter): + assert bu._count_prompt_or_input_tokens("m", "hello") == 5 # len("hello") + + +def test_count_tokens_list_of_str(fake_token_counter): + assert bu._count_prompt_or_input_tokens("m", ["ab", "cde"]) == 5 # 2 + 3 + + +def test_count_tokens_list_of_int(fake_token_counter): + # pre-tokenized prompt: each int counts as one token. + assert bu._count_prompt_or_input_tokens("m", [1, 2, 3, 4]) == 4 + + +def test_count_tokens_list_of_list_of_int(fake_token_counter): + # the bug-fix shape: nested pre-tokenized prompts, each int = 1 token. + assert bu._count_prompt_or_input_tokens("m", [[1, 2, 3], [4, 5]]) == 5 + + +def test_count_tokens_mixed_nested(fake_token_counter): + # nested list with ints + a string: 2 ints (=2) + len("xyz")=3 -> 5 + assert bu._count_prompt_or_input_tokens("m", [[1, 2, "xyz"]]) == 5 + + +def test_count_tokens_unsupported_shape_is_zero(fake_token_counter): + assert bu._count_prompt_or_input_tokens("m", 12345) == 0 + assert bu._count_prompt_or_input_tokens("m", {"a": 1}) == 0 + + +# =========================================================================== # +# _count_entry_tokens (per-entry rate-limit token counting). The individual +# prompt/input/embedding shapes are covered in test_batch_file_validation.py; +# here we pin the body-field precedence and the empty/fallback behavior. +# =========================================================================== # + + +def test_count_entry_messages_path(fake_token_counter): + entry = {"body": {"model": "gpt-4o", "messages": [{"role": "user"}, {"role": "x"}]}} + assert bu._count_entry_tokens(entry) == 2 # len(messages) + + +def test_count_entry_prompt_path(fake_token_counter): + assert bu._count_entry_tokens({"body": {"model": "gpt-4o", "prompt": "abcd"}}) == 4 + + +def test_count_entry_input_path(fake_token_counter): + assert bu._count_entry_tokens({"body": {"model": "gpt-4o", "input": "ab"}}) == 2 + + +def test_count_entry_messages_beats_prompt(fake_token_counter): + # messages present -> prompt/input are ignored (messages is checked first). + entry = { + "body": { + "model": "gpt-4o", + "messages": [{"role": "user"}], + "prompt": "this-should-be-ignored", + } + } + assert bu._count_entry_tokens(entry) == 1 + + +def test_count_entry_prompt_beats_input(fake_token_counter): + entry = {"body": {"model": "gpt-4o", "prompt": "abc", "input": "this-is-longer"}} + assert bu._count_entry_tokens(entry) == 3 + + +def test_count_entry_empty_body_is_zero(fake_token_counter): + assert bu._count_entry_tokens({"body": {}}) == 0 + assert bu._count_entry_tokens({}) == 0 + + +def test_count_entry_uses_model_name_fallback(monkeypatch): + # No body.model -> the model_name argument is forwarded to the token counter. + captured = {} + + def _tc(model=None, text=None, messages=None, **kw): + captured["model"] = model + return len(text or "") + + monkeypatch.setattr(bu, "token_counter", _tc) + bu._count_entry_tokens({"body": {"prompt": "ab"}}, model_name="fallback-model") + assert captured["model"] == "fallback-model" + + +# =========================================================================== # +# _get_batch_job_total_usage_from_file_content (output usage aggregation) +# =========================================================================== # + + +def test_total_usage_sums_successful_only(): + rows = [ + _success_row(usage=_usage(10, 5)), # 15 + _failed_row(), # excluded + _success_row(usage=_usage(20, 10)), # 30 + ] + usage = bu._get_batch_job_total_usage_from_file_content(rows) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 30, + 15, + 45, + ) + + +def test_total_usage_empty_is_zero(): + usage = bu._get_batch_job_total_usage_from_file_content([]) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 0, + 0, + 0, + ) + + +# =========================================================================== # +# _get_batch_job_cost_from_file_content (cost maps mocked) +# =========================================================================== # + + +def test_cost_from_content_completion_cost_path(monkeypatch): + # model_info is None -> litellm.completion_cost per successful row. + calls = [] + + def _completion_cost(**kw): + calls.append(kw) + return 0.5 + + monkeypatch.setattr(litellm, "completion_cost", _completion_cost) + rows = [ + _success_row(usage=_usage(10, 5)), + _failed_row(), # excluded -> not costed + _success_row(usage=_usage(20, 10)), + ] + + total = bu._get_batch_job_cost_from_file_content(rows, custom_llm_provider="openai") + + assert total == 1.0 # 2 successful * 0.5 + assert len(calls) == 2 # failed row not costed + + +def test_cost_from_content_model_info_path(monkeypatch): + # model_info set -> batch_cost_calculator(prompt_cost, completion_cost). + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.1, 0.2)) + rows = [ + _success_row(usage=_usage(10, 5)), + _success_row(usage=_usage(20, 10)), + ] + + total = bu._get_batch_job_cost_from_file_content( + rows, + custom_llm_provider="openai", + model_info={"input_cost_per_token": 0.0}, # type: ignore[arg-type] # truthy -> model_info path + ) + + assert total == pytest.approx(0.6) # 2 * (0.1 + 0.2) + + +# =========================================================================== # +# _batch_cost_calculator (dispatch: vertex-disable-transform vs generic) +# =========================================================================== # + + +def test_batch_cost_calculator_generic_path(monkeypatch): + monkeypatch.setattr(bu, "_get_batch_job_cost_from_file_content", lambda **kw: 4.2) + assert bu._batch_cost_calculator([], custom_llm_provider="openai", model_name="gpt-4o") == 4.2 + + +def test_batch_cost_calculator_vertex_disable_transform_path(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) + monkeypatch.setattr( + bu, + "calculate_vertex_ai_batch_cost_and_usage", + lambda content, model: (9.9, Usage()), + ) + # generic path must NOT be taken + monkeypatch.setattr( + bu, + "_get_batch_job_cost_from_file_content", + lambda **kw: pytest.fail("generic path should not run"), + ) + + cost = bu._batch_cost_calculator([], custom_llm_provider="vertex_ai", model_name="gemini-2.0-flash-001") + assert cost == 9.9 + + +# =========================================================================== # +# calculate_vertex_ai_batch_cost_and_usage (usageMetadata aggregation) +# =========================================================================== # + + +def test_vertex_cost_and_usage_aggregation(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.1, 0.2)) + responses = [ + { + "response": { + "usageMetadata": { + "promptTokenCount": 10, + "candidatesTokenCount": 5, + "totalTokenCount": 15, + } + } + }, + { + "response": { + "usageMetadata": { + "promptTokenCount": 20, + "candidatesTokenCount": 10, + "totalTokenCount": 30, + } + } + }, + ] + + cost, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + + assert cost == pytest.approx(0.6) # 2 * (0.1 + 0.2) + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + 30, + 15, + 45, + ) + + +def test_vertex_cost_skips_none_response_body(monkeypatch): + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (1.0, 0.0)) + responses = [ + {"response": None}, # skipped + { + "response": { + "usageMetadata": { + "promptTokenCount": 7, + "candidatesTokenCount": 3, + "totalTokenCount": 10, + } + } + }, + ] + + cost, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + + assert cost == pytest.approx(1.0) # only one line costed + assert usage.total_tokens == 10 + + +def test_vertex_usage_total_token_fallback(monkeypatch): + # no totalTokenCount -> falls back to prompt + completion. + import litellm.cost_calculator as cc + + monkeypatch.setattr(cc, "batch_cost_calculator", lambda **kw: (0.0, 0.0)) + responses = [{"response": {"usageMetadata": {"promptTokenCount": 8, "candidatesTokenCount": 4}}}] + + _, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + assert usage.total_tokens == 12 + + +def test_vertex_cost_error_in_line_is_swallowed(monkeypatch): + # a cost error on one line must not abort aggregation; usage still tallies. + import litellm.cost_calculator as cc + + def _boom(**kw): + raise RuntimeError("price map miss") + + monkeypatch.setattr(cc, "batch_cost_calculator", _boom) + responses = [ + { + "response": { + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 5, + "totalTokenCount": 10, + } + } + } + ] + + cost, usage = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-x") + assert cost == 0.0 + assert usage.total_tokens == 10 + + +# =========================================================================== # +# calculate_batch_cost_and_usage (async orchestrator) +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_calculate_batch_cost_and_usage_orchestration(monkeypatch): + rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] + monkeypatch.setattr(bu, "_batch_cost_calculator", lambda **kw: 2.5) + monkeypatch.setattr( + bu, + "_get_batch_job_total_usage_from_file_content", + lambda **kw: Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + + cost, usage, models = await bu.calculate_batch_cost_and_usage( + file_content_dictionary=rows, custom_llm_provider="openai" + ) + + assert cost == 2.5 + assert usage.total_tokens == 15 + assert models == ["gpt-4o"] # real _get_batch_models_from_file_content + + +# =========================================================================== # +# _get_batch_output_file_content_as_dictionary (file fetch + credential merge) +# =========================================================================== # + + +def _batch(output_file_id): + from litellm.types.llms.openai import Batch + + return Batch( + id="b", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="f", + object="batch", + status="completed", + output_file_id=output_file_id, + ) + + +@pytest.mark.asyncio +async def test_output_file_content_vertex_raises(): + with pytest.raises(ValueError, match="Vertex AI does not support"): + await bu._get_batch_output_file_content_as_dictionary(_batch("of"), custom_llm_provider="vertex_ai") + + +@pytest.mark.asyncio +async def test_output_file_content_no_output_file_id_raises(): + with pytest.raises(ValueError, match="Output file id is None"): + await bu._get_batch_output_file_content_as_dictionary(_batch(None), custom_llm_provider="openai") + + +@pytest.mark.asyncio +async def test_output_file_content_fetches_and_parses(monkeypatch): + import litellm.files.main as files_main + import litellm.proxy.openai_files_endpoints.common_utils as cu + + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": b'{"a": 1}\n{"b": 2}'})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + monkeypatch.setattr(cu, "_is_base64_encoded_unified_file_id", lambda fid: False) + + result = await bu._get_batch_output_file_content_as_dictionary( + _batch("file-out"), + custom_llm_provider="azure", + litellm_params={"api_key": "sk-az", "api_base": "https://az", "model": "x"}, + ) + + assert result == [{"a": 1}, {"b": 2}] + # afile_content received the file id + extracted credentials (not "model"). + assert captured["file_id"] == "file-out" + assert captured["custom_llm_provider"] == "azure" + assert captured["api_key"] == "sk-az" + assert captured["api_base"] == "https://az" + assert "model" not in captured + + +@pytest.mark.asyncio +async def test_output_file_content_unified_file_id_extraction(monkeypatch): + # a base64 unified id carries the real provider file id inside + # "llm_output_file_id,;" - it must be unwrapped before the fetch. + import litellm.files.main as files_main + import litellm.proxy.openai_files_endpoints.common_utils as cu + + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": b'{"a": 1}'})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + monkeypatch.setattr( + cu, + "_is_base64_encoded_unified_file_id", + lambda fid: "litellm_proxy;llm_output_file_id,real-file-99;rest", + ) + + await bu._get_batch_output_file_content_as_dictionary(_batch("encoded-blob"), custom_llm_provider="openai") + + assert captured["file_id"] == "real-file-99" + + +# =========================================================================== # +# _handle_completed_batch (async orchestrator: fetch -> cost/usage/models) +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_handle_completed_batch_orchestration(monkeypatch): + rows = [_success_row(model="gpt-4o", usage=_usage(10, 5))] + + async def fake_get_content(batch, custom_llm_provider, litellm_params=None): + return rows + + monkeypatch.setattr(bu, "_get_batch_output_file_content_as_dictionary", fake_get_content) + monkeypatch.setattr(bu, "_batch_cost_calculator", lambda **kw: 3.3) + monkeypatch.setattr( + bu, + "_get_batch_job_total_usage_from_file_content", + lambda **kw: Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15), + ) + + cost, usage, models = await bu._handle_completed_batch(_batch("of"), custom_llm_provider="openai") + + assert cost == 3.3 + assert usage.total_tokens == 15 + assert models == ["gpt-4o"] + + +# =========================================================================== # +# Remaining branch: vertex usage disable-transform path. +# +# NOTE: the error path of _get_batch_job_cost_from_file_content (its `raise e`) +# is intentionally NOT tested: the preceding line logs via +# `verbose_logger.error("...", e)`, which passes the exception as a logging +# format-arg with no placeholder and itself raises TypeError under +# logging.raiseExceptions, masking the original error. Asserting that masked +# behavior would lock a source bug; left uncovered on purpose. +# =========================================================================== # + + +def test_total_usage_vertex_disable_transform_path(monkeypatch): + monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False) + monkeypatch.setattr( + bu, + "calculate_vertex_ai_batch_cost_and_usage", + lambda content, model: ( + 0.0, + Usage(prompt_tokens=1, completion_tokens=2, total_tokens=3), + ), + ) + + usage = bu._get_batch_job_total_usage_from_file_content([], custom_llm_provider="vertex_ai", model_name="gemini-x") + assert usage.total_tokens == 3 diff --git a/tests/test_litellm/batches/test_main.py b/tests/test_litellm/batches/test_main.py new file mode 100644 index 00000000000..1f7a91a5511 --- /dev/null +++ b/tests/test_litellm/batches/test_main.py @@ -0,0 +1,744 @@ +""" +Provider-dispatch contract tests for litellm/batches/main.py + +main.py is the SDK layer beneath the proxy batch endpoints: each of +create/retrieve/list/cancel_batch is a switch on `custom_llm_provider` (and, for +create/retrieve, on whether a provider-config + model is present) that hands off +to exactly one provider handler. These tests lock that dispatch: + + 1. DISPATCH - exactly which provider seam fired (openai_batches_instance vs + azure vs vertex vs anthropic vs base_llm_http_handler vs the + Bedrock ARN handlers), with every sibling seam asserted NOT + called. A reordered/negated branch flips this. + 2. PAYLOAD - the request object (CreateBatchRequest/RetrieveBatchRequest/...) + and the _is_async flag forwarded to the handler. + 3. RESULT - the handler's return value is what the function returns. + 4. DELEGATION - the async wrappers (a*_batch) forward to the sync function in an + executor with the right "_is_async" flag, and pass the result + back untouched. + +Only the provider handler instances are mocked (true network boundaries). The +real public functions run (including the @client decorator) so dispatch reflects +production. Provider env vars are not required: missing creds resolve to None and +flow through harmlessly because the handler is mocked. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +import litellm.batches.main as bm + + +# --------------------------------------------------------------------------- # +# Seam harness - one mock per provider handler instance + the Bedrock ARN +# handler. Each handler method auto-returns a unique sentinel (its +# return_value), so "result is seam..return_value" verifies dispatch. +# --------------------------------------------------------------------------- # + + +@dataclass +class Seams: + openai: MagicMock + azure: MagicMock + vertex: MagicMock + anthropic: MagicMock + base_http: MagicMock + bedrock_arn: MagicMock + + +@pytest.fixture +def seams(): + openai_i = MagicMock(name="openai_batches_instance") + azure_i = MagicMock(name="azure_batches_instance") + vertex_i = MagicMock(name="vertex_ai_batches_instance") + anthropic_i = MagicMock(name="anthropic_batches_instance") + base_http = MagicMock(name="base_llm_http_handler") + bedrock_arn = MagicMock(name="BedrockBatchesHandler") + + with ExitStack() as stack: + stack.enter_context(patch.object(bm, "openai_batches_instance", openai_i)) + stack.enter_context(patch.object(bm, "azure_batches_instance", azure_i)) + stack.enter_context(patch.object(bm, "vertex_ai_batches_instance", vertex_i)) + stack.enter_context( + patch.object(bm, "anthropic_batches_instance", anthropic_i) + ) + stack.enter_context(patch.object(bm, "base_llm_http_handler", base_http)) + stack.enter_context(patch.object(bm, "BedrockBatchesHandler", bedrock_arn)) + yield Seams( + openai=openai_i, + azure=azure_i, + vertex=vertex_i, + anthropic=anthropic_i, + base_http=base_http, + bedrock_arn=bedrock_arn, + ) + + +# Every handler method across all provider instances - used to assert +# "no sibling seam fired" exhaustively. +def _all_seam_methods(seams: Seams, op: str): + return [ + getattr(seams.openai, op), + getattr(seams.azure, op), + getattr(seams.vertex, op), + getattr(seams.anthropic, op), + getattr(seams.base_http, op), + ] + + +def _assert_only(fired, seams: Seams, op: str): + """Assert `fired` was called exactly once and every other op seam was not.""" + assert fired.call_count == 1 + for m in _all_seam_methods(seams, op): + if m is not fired: + m.assert_not_called() + + +CREATE_KW: Dict[str, Any] = dict( + completion_window="24h", + endpoint="/v1/chat/completions", + input_file_id="file-abc", +) + + +# =========================================================================== # +# create_batch +# =========================================================================== # + + +def test_create__openai_dispatch_and_payload(seams): + result = bm.create_batch(**CREATE_KW, custom_llm_provider="openai") + + # DISPATCH + RESULT + assert result is seams.openai.create_batch.return_value + _assert_only(seams.openai.create_batch, seams, "create_batch") + seams.bedrock_arn._handle_async_invoke_status.assert_not_called() + + # PAYLOAD - request object built from the call, sync flag off. + kw = seams.openai.create_batch.call_args.kwargs + assert kw["create_batch_data"] == { + "completion_window": "24h", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-abc", + "metadata": None, + "extra_headers": None, + "extra_body": None, + } + assert kw["_is_async"] is False + assert kw["timeout"] == 600.0 + + +def test_create__hosted_vllm_routes_to_openai_instance(seams): + """hosted_vllm is in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, so it shares + the openai handler. Locks that set membership.""" + result = bm.create_batch(**CREATE_KW, custom_llm_provider="hosted_vllm") + + assert result is seams.openai.create_batch.return_value + _assert_only(seams.openai.create_batch, seams, "create_batch") + + +def test_create__azure_dispatch(seams): + result = bm.create_batch(**CREATE_KW, custom_llm_provider="azure") + + assert result is seams.azure.create_batch.return_value + _assert_only(seams.azure.create_batch, seams, "create_batch") + + +def test_create__vertex_ai_dispatch(seams): + result = bm.create_batch(**CREATE_KW, custom_llm_provider="vertex_ai") + + assert result is seams.vertex.create_batch.return_value + _assert_only(seams.vertex.create_batch, seams, "create_batch") + + +def test_create__provider_config_routes_to_base_http_handler(seams): + """model + a provider batches config (bedrock-style) routes to the generic + base_llm_http_handler, NOT the per-provider instance.""" + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + result = bm.create_batch( + **CREATE_KW, custom_llm_provider="bedrock", model="bedrock/my-batch-model" + ) + + assert result is seams.base_http.create_batch.return_value + _assert_only(seams.base_http.create_batch, seams, "create_batch") + + +def test_create__unsupported_provider_raises_badrequest(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.create_batch(**CREATE_KW, custom_llm_provider="cohere") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "create_batch"): + m.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__async_path_propagates_is_async(seams): + """Through the real async wrapper, the handler is invoked with _is_async=True. + (Calling the @client sync create_batch with acreate_batch=True directly is not + a real code path - logging-obj setup only happens on the async wrapper path.)""" + await bm.acreate_batch(**CREATE_KW, custom_llm_provider="openai") + + assert seams.openai.create_batch.call_args.kwargs["_is_async"] is True + + +# =========================================================================== # +# retrieve_batch +# =========================================================================== # + + +def test_retrieve__openai_dispatch_and_payload(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="openai") + + assert result is seams.openai.retrieve_batch.return_value + _assert_only(seams.openai.retrieve_batch, seams, "retrieve_batch") + + kw = seams.openai.retrieve_batch.call_args.kwargs + assert kw["retrieve_batch_data"] == { + "batch_id": "batch-1", + "extra_headers": None, + "extra_body": None, + } + assert kw["_is_async"] is False + + +def test_retrieve__hosted_vllm_routes_to_openai_instance(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="hosted_vllm") + + assert result is seams.openai.retrieve_batch.return_value + _assert_only(seams.openai.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__azure_dispatch(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="azure") + + assert result is seams.azure.retrieve_batch.return_value + _assert_only(seams.azure.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__vertex_ai_dispatch(seams): + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="vertex_ai") + + assert result is seams.vertex.retrieve_batch.return_value + _assert_only(seams.vertex.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__anthropic_dispatch(seams): + """anthropic is retrieve-capable (not in create's provider set).""" + result = bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="anthropic") + + assert result is seams.anthropic.retrieve_batch.return_value + _assert_only(seams.anthropic.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__provider_config_routes_to_base_http_handler(seams): + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + result = bm.retrieve_batch( + batch_id="batch-1", + custom_llm_provider="bedrock", + model="bedrock/my-batch-model", + ) + + assert result is seams.base_http.retrieve_batch.return_value + _assert_only(seams.base_http.retrieve_batch, seams, "retrieve_batch") + + +def test_retrieve__bedrock_async_invoke_arn(seams): + arn = "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123" + result = bm.retrieve_batch(batch_id=arn, custom_llm_provider="bedrock") + + seams.bedrock_arn._handle_async_invoke_status.assert_called_once() + assert result is seams.bedrock_arn._handle_async_invoke_status.return_value + # provider instances untouched. + for m in _all_seam_methods(seams, "retrieve_batch"): + m.assert_not_called() + + +def test_retrieve__bedrock_model_invocation_job_arn(seams): + arn = "arn:aws:bedrock:us-east-1:123456789012:model-invocation-job/xyz789" + result = bm.retrieve_batch(batch_id=arn, custom_llm_provider="bedrock") + + seams.bedrock_arn._handle_model_invocation_job_status.assert_called_once() + assert ( + result is seams.bedrock_arn._handle_model_invocation_job_status.return_value + ) + seams.bedrock_arn._handle_async_invoke_status.assert_not_called() + + +def test_retrieve__unsupported_provider_raises_badrequest(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="cohere") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "retrieve_batch"): + m.assert_not_called() + + +# =========================================================================== # +# list_batches (supported: openai, hosted_vllm, azure, vertex_ai) +# =========================================================================== # + + +def test_list__openai_dispatch_and_payload(seams): + result = bm.list_batches(custom_llm_provider="openai", after="cur", limit=5) + + assert result is seams.openai.list_batches.return_value + _assert_only(seams.openai.list_batches, seams, "list_batches") + + kw = seams.openai.list_batches.call_args.kwargs + assert kw["after"] == "cur" + assert kw["limit"] == 5 + assert kw["_is_async"] is False + + +def test_list__hosted_vllm_routes_to_openai_instance(seams): + result = bm.list_batches(custom_llm_provider="hosted_vllm") + + assert result is seams.openai.list_batches.return_value + _assert_only(seams.openai.list_batches, seams, "list_batches") + + +def test_list__azure_dispatch(seams): + result = bm.list_batches(custom_llm_provider="azure") + + assert result is seams.azure.list_batches.return_value + _assert_only(seams.azure.list_batches, seams, "list_batches") + + +def test_list__vertex_ai_dispatch(seams): + result = bm.list_batches(custom_llm_provider="vertex_ai") + + assert result is seams.vertex.list_batches.return_value + _assert_only(seams.vertex.list_batches, seams, "list_batches") + + +def test_list__unsupported_provider_raises_badrequest(seams): + # anthropic supports retrieve but NOT list - good negative case. + with pytest.raises(litellm.exceptions.BadRequestError): + bm.list_batches(custom_llm_provider="anthropic") # type: ignore[arg-type] + + for m in _all_seam_methods(seams, "list_batches"): + m.assert_not_called() + + +# =========================================================================== # +# cancel_batch (supported: openai, hosted_vllm, azure, vertex_ai; no @client) +# =========================================================================== # + + +def test_cancel__openai_dispatch_and_payload(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="openai") + + assert result is seams.openai.cancel_batch.return_value + _assert_only(seams.openai.cancel_batch, seams, "cancel_batch") + + kw = seams.openai.cancel_batch.call_args.kwargs + assert kw["cancel_batch_data"] == { + "batch_id": "batch-1", + "extra_headers": None, + "extra_body": None, + } + assert kw["_is_async"] is False + + +def test_cancel__azure_dispatch(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="azure") + + assert result is seams.azure.cancel_batch.return_value + _assert_only(seams.azure.cancel_batch, seams, "cancel_batch") + + +def test_cancel__vertex_ai_dispatch(seams): + result = bm.cancel_batch(batch_id="batch-1", custom_llm_provider="vertex_ai") + + assert result is seams.vertex.cancel_batch.return_value + _assert_only(seams.vertex.cancel_batch, seams, "cancel_batch") + + +def test_cancel__unsupported_provider_raises_badrequest(seams): + with pytest.raises(litellm.exceptions.BadRequestError): + bm.cancel_batch(batch_id="batch-1", custom_llm_provider="cohere") + + for m in _all_seam_methods(seams, "cancel_batch"): + m.assert_not_called() + + +def test_cancel__async_flag_propagates_is_async(seams): + bm.cancel_batch( + batch_id="batch-1", custom_llm_provider="openai", acancel_batch=True + ) + + assert seams.openai.cancel_batch.call_args.kwargs["_is_async"] is True + + +# =========================================================================== # +# Async wrappers - delegate to the sync function in an executor, set the right +# "_is_async" flag, and return the result untouched. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_acreate_batch_delegates_to_create_batch(): + with patch.object(bm, "create_batch", MagicMock(return_value="SENTINEL")) as m: + result = await bm.acreate_batch(**CREATE_KW, custom_llm_provider="openai") + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("acreate_batch") is True + # positional handoff: (completion_window, endpoint, input_file_id, provider, ...) + assert m.call_args.args[0] == "24h" + assert m.call_args.args[2] == "file-abc" + assert m.call_args.args[3] == "openai" + + +@pytest.mark.asyncio +async def test_aretrieve_batch_delegates_to_retrieve_batch(): + with patch.object(bm, "retrieve_batch", MagicMock(return_value="SENTINEL")) as m: + result = await bm.aretrieve_batch( + batch_id="batch-1", custom_llm_provider="azure" + ) + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("aretrieve_batch") is True + assert m.call_args.args[0] == "batch-1" + assert m.call_args.args[1] == "azure" + + +@pytest.mark.asyncio +async def test_alist_batches_delegates_to_list_batches(): + with patch.object(bm, "list_batches", MagicMock(return_value="SENTINEL")) as m: + result = await bm.alist_batches( + after="cur", limit=3, custom_llm_provider="vertex_ai" + ) + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("alist_batches") is True + assert m.call_args.args[0] == "cur" + assert m.call_args.args[1] == 3 + assert m.call_args.args[2] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_acancel_batch_delegates_to_cancel_batch(): + with patch.object(bm, "cancel_batch", MagicMock(return_value="SENTINEL")) as m: + result = await bm.acancel_batch( + batch_id="batch-1", custom_llm_provider="openai" + ) + + assert result == "SENTINEL" + assert m.call_count == 1 + assert m.call_args.kwargs.get("acancel_batch") is True + assert m.call_args.args[0] == "batch-1" + + +# =========================================================================== # +# Credential passthrough - when the caller supplies credentials in kwargs, they +# must reach the provider handler. Explicit kwargs win over litellm.* globals and +# env vars (they are first in each `optional_params.x or litellm.x or env` chain), +# so these assertions are deterministic regardless of the test environment. +# +# The credential-resolution blocks are copy-pasted per provider in EACH of +# create/retrieve/list/cancel, so a regression can land in any one independently; +# every function is checked. +# =========================================================================== # + + +# Distinct values so a cross-wired field (e.g. api_key forwarded as api_base) is +# impossible to miss. +OPENAI_CREDS: Dict[str, Any] = dict( + api_key="sk-user-openai", + api_base="https://openai.user.test", + organization="org-user-123", + max_retries=7, +) +AZURE_CREDS: Dict[str, Any] = dict( + api_key="sk-user-azure", + api_base="https://azure.user.test", + api_version="2024-12-99", +) +VERTEX_CREDS: Dict[str, Any] = dict( + vertex_project="proj-user", + vertex_location="loc-user", + vertex_credentials="cred-user", + api_base="https://vertex.user.test", +) + + +def _sent(mock_method, *keys): + """Subset of the call kwargs limited to `keys`, for exact comparison.""" + kw = mock_method.call_args.kwargs + return {k: kw.get(k) for k in keys} + + +# ---- create_batch ---------------------------------------------------------- # + + +def test_create__openai_credentials_passthrough(seams): + bm.create_batch(**CREATE_KW, custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.create_batch, "api_key", "api_base", "organization", "max_retries" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + "max_retries": 7, + } + + +def test_create__azure_credentials_passthrough(seams): + bm.create_batch(**CREATE_KW, custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.create_batch, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_create__vertex_credentials_passthrough(seams): + bm.create_batch(**CREATE_KW, custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.create_batch, + "vertex_project", + "vertex_location", + "vertex_credentials", + "api_base", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + "api_base": "https://vertex.user.test", + } + + +def test_create__provider_config_credentials_passthrough(seams): + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + bm.create_batch( + **CREATE_KW, + custom_llm_provider="bedrock", + model="bedrock/my-batch-model", + api_key="sk-user-bedrock", + api_base="https://bedrock.user.test", + ) + + assert _sent(seams.base_http.create_batch, "api_key", "api_base") == { + "api_key": "sk-user-bedrock", + "api_base": "https://bedrock.user.test", + } + + +# ---- retrieve_batch -------------------------------------------------------- # + + +def test_retrieve__openai_credentials_passthrough(seams): + bm.retrieve_batch(batch_id="b1", custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.retrieve_batch, "api_key", "api_base", "organization" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + } + + +def test_retrieve__azure_credentials_passthrough(seams): + bm.retrieve_batch(batch_id="b1", custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.retrieve_batch, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_retrieve__vertex_credentials_passthrough(seams): + bm.retrieve_batch(batch_id="b1", custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.retrieve_batch, + "vertex_project", + "vertex_location", + "vertex_credentials", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + } + + +def test_retrieve__anthropic_credentials_passthrough(seams): + bm.retrieve_batch( + batch_id="b1", + custom_llm_provider="anthropic", + api_key="sk-user-anthropic", + api_base="https://anthropic.user.test", + ) + + assert _sent(seams.anthropic.retrieve_batch, "api_key", "api_base") == { + "api_key": "sk-user-anthropic", + "api_base": "https://anthropic.user.test", + } + + +def test_retrieve__provider_config_credentials_passthrough(seams): + with patch.object( + bm.ProviderConfigManager, + "get_provider_batches_config", + return_value=MagicMock(name="provider_config"), + ): + bm.retrieve_batch( + batch_id="b1", + custom_llm_provider="bedrock", + model="bedrock/my-batch-model", + api_key="sk-user-bedrock", + api_base="https://bedrock.user.test", + ) + + assert _sent(seams.base_http.retrieve_batch, "api_key", "api_base") == { + "api_key": "sk-user-bedrock", + "api_base": "https://bedrock.user.test", + } + + +# ---- list_batches ---------------------------------------------------------- # + + +def test_list__openai_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.list_batches, "api_key", "api_base", "organization" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + } + + +def test_list__azure_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.list_batches, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_list__vertex_credentials_passthrough(seams): + bm.list_batches(custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.list_batches, + "vertex_project", + "vertex_location", + "vertex_credentials", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + } + + +# ---- cancel_batch ---------------------------------------------------------- # + + +def test_cancel__openai_credentials_passthrough(seams): + bm.cancel_batch(batch_id="b1", custom_llm_provider="openai", **OPENAI_CREDS) + + assert _sent( + seams.openai.cancel_batch, "api_key", "api_base", "organization" + ) == { + "api_key": "sk-user-openai", + "api_base": "https://openai.user.test", + "organization": "org-user-123", + } + + +def test_cancel__azure_credentials_passthrough(seams): + bm.cancel_batch(batch_id="b1", custom_llm_provider="azure", **AZURE_CREDS) + + assert _sent( + seams.azure.cancel_batch, "api_key", "api_base", "api_version" + ) == { + "api_key": "sk-user-azure", + "api_base": "https://azure.user.test", + "api_version": "2024-12-99", + } + + +def test_cancel__vertex_credentials_passthrough(seams): + bm.cancel_batch(batch_id="b1", custom_llm_provider="vertex_ai", **VERTEX_CREDS) + + assert _sent( + seams.vertex.cancel_batch, + "vertex_project", + "vertex_location", + "vertex_credentials", + ) == { + "vertex_project": "proj-user", + "vertex_location": "loc-user", + "vertex_credentials": "cred-user", + } + + +# =========================================================================== # +# _resolve_timeout - pure helper (used by create_batch). +# =========================================================================== # + + +def _params(**kw): + from litellm.types.router import GenericLiteLLMParams + + return GenericLiteLLMParams(**kw) + + +def test_resolve_timeout__explicit_numeric(): + assert bm._resolve_timeout(_params(timeout=30), {}, "openai") == 30.0 + + +def test_resolve_timeout__default_when_unset(): + assert bm._resolve_timeout(_params(), {}, "openai") == 600.0 + + +def test_resolve_timeout__request_timeout_kwarg_fallback(): + assert bm._resolve_timeout(_params(), {"request_timeout": 45}, "openai") == 45.0 + + +def test_resolve_timeout__httpx_timeout_returns_float_read(): + import httpx + + t = httpx.Timeout(99.0, connect=5.0) + resolved = bm._resolve_timeout(_params(timeout=t), {}, "openai") + assert isinstance(resolved, float) + assert resolved == 99.0 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 new file mode 100644 index 00000000000..0a472d86257 --- /dev/null +++ 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 new file mode 100644 index 00000000000..4a2adb01ea5 --- /dev/null +++ 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/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 new file mode 100644 index 00000000000..f2332a7de7c --- /dev/null +++ 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 new file mode 100644 index 00000000000..cfb9f278f80 --- /dev/null +++ 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 new file mode 100644 index 00000000000..d1ad5943ae6 --- /dev/null +++ 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 new file mode 100644 index 00000000000..cacea234777 --- /dev/null +++ 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 new file mode 100644 index 00000000000..37084d43441 --- /dev/null +++ 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 diff --git a/tests/test_litellm/proxy/batches_endpoints/__init__.py b/tests/test_litellm/proxy/batches_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py new file mode 100644 index 00000000000..26c654cd154 --- /dev/null +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -0,0 +1,2026 @@ +""" +Routing-contract tests for litellm/proxy/batches_endpoints/endpoints.py + +These are not happy-path smoke tests. Each row of the matrix locks the full +contract of a single routing branch so that *any* behavior change in this layer +fails loudly: + + 1. DISPATCH - exactly which downstream seam fired (litellm.acreate_batch + vs llm_router.acreate_batch), and every sibling seam is + asserted NOT called. A reordered/negated branch flips this. + 2. CREDENTIALS - the credential resolver receives the model derived from the + request, not a hardcoded value. The router's + get_deployment_credentials_with_provider is input-locked. + 3. SEAM PAYLOAD - the *entire* kwargs dict forwarded to the provider call is + exact-matched. Because LiteLLMBatchCreateRequest is a + TypedDict (zero runtime filtering), nothing else stops a + newly-added param from silently reaching every provider. + This exact-match is that missing guard: a new key fails the + test and forces a reviewer to ask "does this work for all + providers, or just openai". + 4. OUTPUT SHAPE - the id encode/decode round-trip clients depend on. + +Only true I/O boundaries are mocked (provider call, router, proxy logging, +request parsing, pre-call enrichment). The pure encode/decode/credential-merge +helpers run for real so the payload assertions reflect production exactly. + +The object mocks are spec'd to their real classes, so a brand-new method call +added to this layer raises instead of silently passing - the inventory of seams +cannot drift without a test failure. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +import litellm.proxy.batches_endpoints.endpoints as endpoints +import litellm.proxy.proxy_server as proxy_server +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.openai_files_endpoints.common_utils import ( + encode_file_id_with_model, +) +from litellm.proxy.utils import ProxyLogging +from litellm.router import Router +from litellm.types.llms.openai import BatchJobStatus +from litellm.types.utils import LiteLLMBatch + +from fastapi import Response + +# --------------------------------------------------------------------------- # +# Fixtures: distinguishable credentials per model so a wrong/hardcoded model_id +# produces wrong creds (or KeyError) and is impossible to hide. +# --------------------------------------------------------------------------- # + +CREDS: Dict[str, Dict[str, str]] = { + "azure/gpt-4o": { + "custom_llm_provider": "azure", + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o-deployment", + }, + "vertex-model": { + "custom_llm_provider": "vertex_ai", + "api_key": "sk-vertex", + "api_base": "https://vertex.test", + "model": "vertex_ai/gemini-2.0", + }, +} + +# A real model-encoded file id: decodes to "azure/gpt-4o", strips to "file-original123". +AZURE_FILE_ID = encode_file_id_with_model( + "file-original123", "azure/gpt-4o", id_type="file" +) + + +def make_batch( + *, + id: str = "batch-provider-id", + output_file_id: Optional[str] = None, + error_file_id: Optional[str] = None, + input_file_id: Optional[str] = None, + status: BatchJobStatus = "validating", +) -> LiteLLMBatch: + batch = LiteLLMBatch( + id=id, + completion_window="24h", + created_at=1234567890, + endpoint="/v1/chat/completions", + input_file_id=input_file_id or "file-provider-input", + object="batch", + status=status, + ) + if output_file_id is not None: + batch.output_file_id = output_file_id + if error_file_id is not None: + batch.error_file_id = error_file_id + batch._hidden_params = {} + return batch + + +class FakeRequest: + """Minimal stand-in. The request is only read via .headers/.query_params on + the model-param fallback path; everything else that touches it is mocked.""" + + def __init__( + self, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + ): + self.headers = headers or {} + self.query_params = query or {} + + +@dataclass +class Harness: + """Holds every mocked seam so a test can configure inputs and assert calls.""" + + body: Dict[str, Any] + read_body: AsyncMock + pre_call: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + is_known_model: MagicMock + litellm_acreate: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + + @property + def router_acreate(self) -> AsyncMock: + return self.router.acreate_batch + + def acreate_kwargs(self) -> Dict[str, Any]: + """Exact kwargs forwarded to litellm.acreate_batch.""" + assert self.litellm_acreate.call_count == 1 + return dict(self.litellm_acreate.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_acreate.call_count == 1 + return dict(self.router_acreate.call_args.kwargs) + + +def _creds_lookup(*, model_id: str) -> Dict[str, str]: + # KeyError on an unknown/hardcoded model_id - the bug cannot hide. + return dict(CREDS[model_id]) + + +@pytest.fixture +def harness(): + """Seam harness. Patches only true I/O boundaries; pure encode/decode/merge + helpers run for real. Object mocks are spec'd so unknown method calls raise.""" + body_holder: Dict[str, Any] = {} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.acreate_batch = AsyncMock(return_value=make_batch()) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + read_body = AsyncMock(side_effect=lambda request: body_holder["body"]) + pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock())) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + is_known_model = MagicMock(return_value=False) + litellm_acreate = AsyncMock(return_value=make_batch()) + + with ExitStack() as stack: + stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context(patch.object(endpoints, "is_known_model", is_known_model)) + stack.enter_context(patch.object(litellm, "acreate_batch", litellm_acreate)) + stack.enter_context( + patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + + h = Harness( + body=body_holder, + read_body=read_body, + pre_call=pre_call, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + is_known_model=is_known_model, + litellm_acreate=litellm_acreate, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + ) + yield h + + +def set_body(harness: Harness, body: Dict[str, Any]) -> None: + harness.body["body"] = body + + +async def call_create( + harness: Harness, + *, + provider: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, +): + return await endpoints.create_batch( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + provider=provider, + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + ) + + +# =========================================================================== # +# SCENARIO 1 - input_file_id encoded with model. The full showcase: every +# assertion type from the design lives here. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__model_encoded_file_id(harness): + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + resp = await call_create(harness) + + # 1. DISPATCH - model-credential path fired via litellm, router did not. + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + + # 2. CREDENTIALS - resolved for the model decoded FROM the file id. + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + # 3. SEAM PAYLOAD - exact, whole dict. A new forwarded key breaks this. + assert harness.acreate_kwargs() == { + "custom_llm_provider": "azure", + "input_file_id": "file-original123", # encoding stripped by this layer + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, # sanitize_openai_provider_metadata(None) + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o-deployment", + } + + # 4. OUTPUT SHAPE - ids re-encoded with the model; input_file_id restored. + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "azure/gpt-4o", id_type="batch" + ) + assert resp.input_file_id == AZURE_FILE_ID + + +@pytest.mark.asyncio +async def test_create__model_encoded_file_id__encodes_output_and_error_ids(harness): + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.litellm_acreate.return_value = make_batch( + id="batch-xyz", + output_file_id="file-out-raw", + error_file_id="file-err-raw", + ) + + resp = await call_create(harness) + + assert resp.output_file_id == encode_file_id_with_model( + "file-out-raw", "azure/gpt-4o" + ) + assert resp.error_file_id == encode_file_id_with_model( + "file-err-raw", "azure/gpt-4o" + ) + + +@pytest.mark.asyncio +async def test_create__model_encoded_file_id__resolver_gets_decoded_model(harness): + """Regression guard: model_id for credential resolution must be derived from + the file id. A hardcode would call the resolver with the wrong model.""" + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# SCENARIO 2 - model from body / header / query. Locks source precedence. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__model_from_body(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "vertex-model", + }, + ) + + resp = await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_called_once_with(model_id="vertex-model") + payload = harness.acreate_kwargs() + assert payload["custom_llm_provider"] == "vertex_ai" + assert payload["input_file_id"] == "file-plain" + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "vertex-model", id_type="batch" + ) + + +@pytest.mark.asyncio +async def test_create__model_from_header(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, headers={"x-litellm-model": "vertex-model"}) + + harness.creds_resolver.assert_called_once_with(model_id="vertex-model") + harness.router_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__model_from_query(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, query={"model": "vertex-model"}) + + harness.creds_resolver.assert_called_once_with(model_id="vertex-model") + harness.router_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__body_model_beats_header_and_query(harness): + """Precedence row: body > header > query (data.get('model') first).""" + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "azure/gpt-4o", + }, + ) + + await call_create( + harness, + headers={"x-litellm-model": "vertex-model"}, + query={"model": "vertex-model"}, + ) + + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# SCENARIO 3 - fallback to custom_llm_provider (env-var creds). MUST NOT touch +# the credential resolver. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__fallback_default_openai(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_not_called() # inverse-bug guard + assert harness.acreate_kwargs()["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_create__fallback_provider_path_param(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, provider="anthropic") + + harness.creds_resolver.assert_not_called() + assert harness.acreate_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_create__fallback_body_custom_llm_provider(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "custom_llm_provider": "bedrock", + }, + ) + + await call_create(harness) + + payload = harness.acreate_kwargs() + assert payload["custom_llm_provider"] == "bedrock" + + +# =========================================================================== # +# Unified file id routing (-> llm_router). Helpers mocked only here because a +# real unified id is opaque base64; the routing contract is what we lock. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__unified_file_id_single_model(harness): + set_body( + harness, + { + "input_file_id": "litellm_proxy_unified_id", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"] + ): + resp = await call_create(harness) + + # DISPATCH - router fired, direct litellm did not. + assert harness.router_acreate.call_count == 1 + harness.litellm_acreate.assert_not_called() + # model injected from the unified id, input_file_id restored, hidden param set + assert harness.router_kwargs()["model"] == "gpt-4o-mini" + assert resp.input_file_id == "litellm_proxy_unified_id" + assert resp._hidden_params["unified_file_id"] == "unified-xyz" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("models", [[], ["m1", "m2"]]) +async def test_create__unified_file_id_not_exactly_one_model_400(harness, models): + set_body( + harness, + { + "input_file_id": "litellm_proxy_unified_id", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=models + ): + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "400" + harness.router_acreate.assert_not_called() + harness.litellm_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__model_encoded_beats_unified(harness): + """Precedence row: a file id that is BOTH model-encoded and (pretend) unified + must take the model-encoded branch (checked first).""" + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=["something-else"] + ): + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# Loadbalancing branch (-> llm_router) and its precedence vs model-encoded. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__loadbalancing_routes_to_router(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "lb-model", + }, + ) + harness.is_known_model.return_value = True + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + await call_create(harness) + + harness.is_known_model.assert_called_once_with( + model="lb-model", llm_router=harness.router + ) + assert harness.router_acreate.call_count == 1 + harness.litellm_acreate.assert_not_called() + harness.creds_resolver.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__model_encoded_beats_loadbalancing(harness): + set_body( + harness, + { + "input_file_id": AZURE_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "lb-model", + }, + ) + harness.is_known_model.return_value = True + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + await call_create(harness) + + assert harness.litellm_acreate.call_count == 1 + harness.router_acreate.assert_not_called() + harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# =========================================================================== # +# Team-level batch expiry enforcement (independent of routing). +# =========================================================================== # + + +def _user_with_expiry(expiry: Any) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + team_metadata={"enforced_batch_output_expires_after": expiry}, + ) + + +@pytest.mark.asyncio +async def test_create__team_expiry_injected(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create( + harness, user=_user_with_expiry({"anchor": "created_at", "seconds": 3600}) + ) + + assert harness.acreate_kwargs()["output_expires_after"] == { + "anchor": "created_at", + "seconds": 3600, + } + + +@pytest.mark.asyncio +async def test_create__no_team_expiry_not_injected(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness, user=UserAPIKeyAuth(api_key="sk-test")) + + assert "output_expires_after" not in harness.acreate_kwargs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "expiry", + [ + {"seconds": 3600}, # missing anchor + {"anchor": "created_at"}, # missing seconds + {"anchor": "completed_at", "seconds": 3600}, # wrong anchor + ], +) +async def test_create__team_expiry_malformed_500(harness, expiry): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + with pytest.raises(ProxyException) as exc: + await call_create(harness, user=_user_with_expiry(expiry)) + + assert exc.value.code == "500" + harness.litellm_acreate.assert_not_called() + harness.router_acreate.assert_not_called() + + +# =========================================================================== # +# Cross-cutting: enrichment route_type, metadata sanitization, failure hook. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__uses_acreate_batch_route_type(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + + await call_create(harness) + + assert harness.pre_call.call_args.kwargs["route_type"] == "acreate_batch" + + +@pytest.mark.asyncio +async def test_create__metadata_sanitized_before_forwarding(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": {"user_key": "user_val", "spend_logs_metadata": {"x": 1}}, + }, + ) + + await call_create(harness) + + # provider-internal-only key dropped, string key kept (real sanitize runs) + assert harness.acreate_kwargs()["metadata"] == {"user_key": "user_val"} + + +@pytest.mark.asyncio +async def test_create__exception_calls_failure_hook(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.litellm_acreate.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_create(harness) + + harness.logging.post_call_failure_hook.assert_called_once() + assert ( + harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# # +# GET /v1/batches/{batch_id} - retrieve_batch routing-contract tests # +# # +# Same discipline as create_batch above. retrieve_batch has more seams: it # +# first consults the ManagedObjectTable (get_batch_from_database) and may # +# short-circuit on a terminal-status row WITHOUT ever calling a provider, # +# then on a miss/non-terminal row routes to one of three downstream seams # +# (litellm.aretrieve_batch via model creds, llm_router.aretrieve_batch, or # +# litellm.aretrieve_batch via env-var provider) and writes the fresh state # +# back via update_batch_in_database. Every test below locks exactly which # +# of those seams fired and asserts the siblings did NOT, so a reordered or # +# negated branch - or a dropped DB short-circuit / write-back - fails loud. # +# # +# The DB seams (get_batch_from_database, update_batch_in_database, # +# resolve_*_to_unified) are true prisma I/O boundaries and are mocked. The # +# encode/decode/credential-merge helpers run for real, so payload and id # +# round-trip assertions reflect production exactly. # +# =========================================================================== # + + +# A real model-encoded BATCH id: decodes to "azure/gpt-4o", strips to +# "batch_orig123". Distinct from AZURE_FILE_ID so retrieve tests can't pass by +# accidentally reusing the create fixture's value. +AZURE_BATCH_ID = encode_file_id_with_model( + "batch_orig123", "azure/gpt-4o", id_type="batch" +) + +# A realistic decoded unified batch id (what _is_base64_encoded_unified_file_id +# returns). model_id / llm_batch_id are parsed out of this by the real helpers. +UNIFIED_BATCH_ID = "litellm_proxy;model_id:gpt-4o-mini;llm_batch_id:batch-raw-xyz" + + +@dataclass +class RetrieveHarness: + """Seams for retrieve_batch. `data['data']` is the dict pre-call enrichment + returns; the routing branches mutate it, so it is reset per call.""" + + data: Dict[str, Any] + pre_call: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + litellm_aretrieve: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + get_batch_from_db: AsyncMock + update_batch_in_db: AsyncMock + resolve_input: AsyncMock + resolve_output: AsyncMock + + @property + def router_aretrieve(self) -> AsyncMock: + return self.router.aretrieve_batch + + def aretrieve_kwargs(self) -> Dict[str, Any]: + """Exact kwargs forwarded to litellm.aretrieve_batch.""" + assert self.litellm_aretrieve.call_count == 1 + return dict(self.litellm_aretrieve.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_aretrieve.call_count == 1 + return dict(self.router_aretrieve.call_args.kwargs) + + +@pytest.fixture +def retrieve_harness(): + data_holder: Dict[str, Any] = {"data": {}} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.aretrieve_batch = AsyncMock(return_value=make_batch()) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock())) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + litellm_aretrieve = AsyncMock(return_value=make_batch()) + # Default: DB miss -> always fall through to provider routing. + get_batch_from_db = AsyncMock(return_value=(None, None)) + update_batch_in_db = AsyncMock(return_value=None) + resolve_input = AsyncMock(return_value=None) + resolve_output = AsyncMock(return_value=None) + + with ExitStack() as stack: + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context( + patch.object(endpoints, "get_batch_from_database", get_batch_from_db) + ) + stack.enter_context( + patch.object(endpoints, "update_batch_in_database", update_batch_in_db) + ) + stack.enter_context( + patch.object(endpoints, "resolve_input_file_id_to_unified", resolve_input) + ) + stack.enter_context( + patch.object( + endpoints, "resolve_output_file_ids_to_unified", resolve_output + ) + ) + stack.enter_context(patch.object(litellm, "aretrieve_batch", litellm_aretrieve)) + stack.enter_context( + patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + stack.enter_context(patch.object(proxy_server, "prisma_client", MagicMock())) + + yield RetrieveHarness( + data=data_holder, + pre_call=pre_call, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + litellm_aretrieve=litellm_aretrieve, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + get_batch_from_db=get_batch_from_db, + update_batch_in_db=update_batch_in_db, + resolve_input=resolve_input, + resolve_output=resolve_output, + ) + + +async def call_retrieve( + harness: RetrieveHarness, + batch_id: str, + *, + provider: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, +): + # Mirror the real flow: data starts as RetrieveBatchRequest(batch_id=...). + harness.data["data"] = {"batch_id": batch_id} + return await endpoints.retrieve_batch( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + provider=provider, + batch_id=batch_id, + ) + + +# --------------------------------------------------------------------------- # +# SCENARIO 1 - batch id encoded with model. litellm.aretrieve_batch via the +# model's resolved credentials; response ids re-encoded for the client. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_id(retrieve_harness): + resp = await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + # 1. DISPATCH - model-credential path fired via litellm, router did not. + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + + # 2. CREDENTIALS - resolved for the model decoded FROM the batch id. + retrieve_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + # 3. SEAM PAYLOAD - exact, whole dict forwarded to the provider call. + # Note `model` is the DECODED model, not the deployment from creds: the + # endpoint overrides it (provider-config providers like bedrock need it). + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "azure", + "batch_id": "batch_orig123", # encoding stripped by this layer + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o", + } + + # 4. OUTPUT SHAPE - ids re-encoded with the model for the round-trip. + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "azure/gpt-4o", id_type="batch" + ) + + # write-back to the managed-object table happened, tagged as a retrieve. + assert retrieve_harness.update_batch_in_db.call_count == 1 + assert retrieve_harness.update_batch_in_db.call_args.kwargs["operation"] == "retrieve" + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_id__forwards_decoded_model_not_deployment( + retrieve_harness, +): + """Regression guard for the line-483 override: the model forwarded to the + provider must be the decoded model id, never the deployment name that the + credential merge pulled in. Dropping the override silently 400s bedrock.""" + await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + assert retrieve_harness.aretrieve_kwargs()["model"] == "azure/gpt-4o" + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_id__encodes_output_and_error_ids( + retrieve_harness, +): + retrieve_harness.litellm_aretrieve.return_value = make_batch( + id="batch-xyz", + output_file_id="file-out-raw", + error_file_id="file-err-raw", + ) + + resp = await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + assert resp.output_file_id == encode_file_id_with_model( + "file-out-raw", "azure/gpt-4o" + ) + assert resp.error_file_id == encode_file_id_with_model( + "file-err-raw", "azure/gpt-4o" + ) + + +@pytest.mark.asyncio +async def test_retrieve__model_encoded_beats_loadbalancing(retrieve_harness): + """Precedence: model-encoded id is checked before the loadbalancing/unified + elif, so it wins even with loadbalancing enabled.""" + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + await call_retrieve(retrieve_harness, AZURE_BATCH_ID) + + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# --------------------------------------------------------------------------- # +# Unified managed batch id -> llm_router.aretrieve_batch. model_id is parsed +# out of the unified id and stamped onto hidden params; raw file ids on the +# response are resolved back to unified ids. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + resp = await call_retrieve(retrieve_harness, "batch-unified-blob") + + # DISPATCH - router fired, direct litellm did not. + assert retrieve_harness.router_aretrieve.call_count == 1 + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.creds_resolver.assert_not_called() + + # router receives the (still-encoded) batch id verbatim - this layer does + # not decode it for the unified path. + assert retrieve_harness.router_kwargs() == {"batch_id": "batch-unified-blob"} + + # hidden params: unified id passed through, model_id parsed from it. + assert resp._hidden_params["unified_batch_id"] == UNIFIED_BATCH_ID + assert resp._hidden_params["model_id"] == "gpt-4o-mini" + + # raw provider file ids on the response are resolved back to unified ids. + retrieve_harness.resolve_input.assert_called_once() + retrieve_harness.resolve_output.assert_called_once() + + +@pytest.mark.asyncio +async def test_retrieve__loadbalancing_raw_id_routes_to_router(retrieve_harness): + """Loadbalancing on + a plain (non-encoded, non-unified) batch id routes to + the router. Locks the current dispatch contract of the shared elif.""" + with patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", True): + resp = await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.router_aretrieve.call_count == 1 + retrieve_harness.litellm_aretrieve.assert_not_called() + assert retrieve_harness.router_kwargs() == {"batch_id": "batch-raw-xyz"} + # not a unified id -> hidden param reflects that, no model_id stamped. + assert resp._hidden_params["unified_batch_id"] is False + assert "model_id" not in resp._hidden_params + # not a unified id -> no file-id resolution. + retrieve_harness.resolve_input.assert_not_called() + retrieve_harness.resolve_output.assert_not_called() + + +# --------------------------------------------------------------------------- # +# SCENARIO 3 - fallback to custom_llm_provider (env-var creds). MUST NOT touch +# the credential resolver or the router. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__fallback_default_openai(retrieve_harness): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.litellm_aretrieve.call_count == 1 + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.creds_resolver.assert_not_called() # inverse-bug guard + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "openai", + "batch_id": "batch-raw-xyz", + } + assert retrieve_harness.update_batch_in_db.call_count == 1 + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_path_param(retrieve_harness): + await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") + + retrieve_harness.creds_resolver.assert_not_called() + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_from_header(retrieve_harness): + retrieve_harness.provider_from_headers.return_value = "bedrock" + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "bedrock" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_from_query(retrieve_harness): + retrieve_harness.provider_from_query.return_value = "vertex_ai" + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_retrieve__fallback_provider_precedence_path_over_header( + retrieve_harness, +): + """provider path param beats the header-derived provider.""" + retrieve_harness.provider_from_headers.return_value = "bedrock" + + await call_retrieve(retrieve_harness, "batch-raw-xyz", provider="anthropic") + + assert retrieve_harness.aretrieve_kwargs()["custom_llm_provider"] == "anthropic" + + +# --------------------------------------------------------------------------- # +# ManagedObjectTable short-circuit. A terminal-status DB row is returned +# immediately - no provider call, no write-back. A non-terminal row falls +# through to a provider sync. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "status", ["completed", "complete", "failed", "cancelled", "expired"] +) +async def test_retrieve__db_terminal_state_short_circuits(retrieve_harness, status): + # "complete" is the DB-normalized alias of "completed"; it is not a valid + # constructor literal but reaches the endpoint via a stored row, so set it + # post-construction to exercise that exact branch. + db_response = make_batch(id="batch-from-db", status="completed") + db_response.status = status + retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response) + + resp = await call_retrieve(retrieve_harness, "batch-raw-xyz") + + # No provider seam fired, and no write-back (the row is already terminal). + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.update_batch_in_db.assert_not_called() + # The DB object is what the client gets back. + assert resp is db_response + + +@pytest.mark.asyncio +async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness): + db_response = make_batch(id="batch-from-db", status="completed") + retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response) + + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + await call_retrieve(retrieve_harness, "batch-unified-blob") + + # Terminal short-circuit still resolves raw provider file ids to unified. + retrieve_harness.resolve_input.assert_called_once() + retrieve_harness.resolve_output.assert_called_once() + retrieve_harness.litellm_aretrieve.assert_not_called() + retrieve_harness.router_aretrieve.assert_not_called() + + +@pytest.mark.asyncio +async def test_retrieve__db_non_terminal_state_syncs_with_provider(retrieve_harness): + """A non-terminal DB row must NOT short-circuit; the endpoint syncs with the + provider to refresh state.""" + db_response = make_batch(id="batch-from-db", status="validating") + retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response) + + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + # Provider sync happened despite the DB hit. + assert retrieve_harness.litellm_aretrieve.call_count == 1 + assert retrieve_harness.update_batch_in_db.call_count == 1 + + +# --------------------------------------------------------------------------- # +# Cross-cutting: enrichment route_type and failure-hook on provider error. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_retrieve__uses_aretrieve_batch_route_type(retrieve_harness): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert ( + retrieve_harness.pre_call.call_args.kwargs["route_type"] == "aretrieve_batch" + ) + + +@pytest.mark.asyncio +async def test_retrieve__exception_calls_failure_hook(retrieve_harness): + retrieve_harness.litellm_aretrieve.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + retrieve_harness.logging.post_call_failure_hook.assert_called_once() + assert ( + retrieve_harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# # +# GET /v1/batches - list_batches routing-contract tests # +# # +# Branch order (first match wins): # +# 1. managed_files hook present -> managed_files_obj.list_user_batches # +# 2. model from body/query/header -> litellm.alist_batches + id encode # +# 3. target_model_names (param or body) -> llm_router.alist_batches # +# 4. fallback -> litellm.alist_batches via env-var custom_llm_provider # +# # +# llm_router is required; absence is a 500 before any branch runs. # +# =========================================================================== # + + +class FakeListPage: + """Stand-in for the SyncCursorPage[Batch] that alist_batches returns. The + endpoint only touches `.data` (to encode ids) and `._hidden_params`.""" + + def __init__(self, data: Any): + self.data = data + self._hidden_params: Dict[str, Any] = {} + + +@dataclass +class ListHarness: + body: Dict[str, Any] + read_body: AsyncMock + pre_call: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + litellm_alist: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + + @property + def router_alist(self) -> AsyncMock: + return self.router.alist_batches + + def set_managed_files(self, page: Any) -> AsyncMock: + """Install a managed_files hook exposing list_user_batches -> page.""" + hook = MagicMock() + hook.list_user_batches = AsyncMock(return_value=page) + self.logging.get_proxy_hook = MagicMock(return_value=hook) + return hook.list_user_batches + + def alist_kwargs(self) -> Dict[str, Any]: + assert self.litellm_alist.call_count == 1 + return dict(self.litellm_alist.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_alist.call_count == 1 + return dict(self.router_alist.call_args.kwargs) + + +@pytest.fixture +def list_harness(): + body_holder: Dict[str, Any] = {"body": {}} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + # Default: no managed_files hook -> branches 2/3/4 are reachable. + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.alist_batches = AsyncMock(return_value=FakeListPage([])) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + read_body = AsyncMock(side_effect=lambda request: body_holder["body"]) + pre_call = AsyncMock(side_effect=lambda **kw: (body_holder["body"], MagicMock())) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + litellm_alist = AsyncMock(return_value=FakeListPage([])) + + with ExitStack() as stack: + stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context(patch.object(litellm, "alist_batches", litellm_alist)) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + + yield ListHarness( + body=body_holder, + read_body=read_body, + pre_call=pre_call, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + litellm_alist=litellm_alist, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + ) + + +async def call_list( + harness: ListHarness, + *, + provider: Optional[str] = None, + limit: Optional[int] = None, + after: Optional[str] = None, + target_model_names: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + body: Optional[Dict[str, Any]] = None, +): + harness.body["body"] = body if body is not None else {} + return await endpoints.list_batches( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + provider=provider, + limit=limit, + after=after, + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + target_model_names=target_model_names, + ) + + +# --------------------------------------------------------------------------- # +# Branch 1 - ManagedObjectTable listing. This is the default production path +# (the managed_files hook is registered) and wins over every other branch. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__managed_files_path(list_harness): + page = FakeListPage([make_batch(id="batch-1")]) + list_user_batches = list_harness.set_managed_files(page) + + user = UserAPIKeyAuth(api_key="sk-test") + resp = await call_list( + list_harness, + user=user, + limit=7, + after="batch-cursor", + provider="openai", + target_model_names="m1,m2", + ) + + # DISPATCH - managed-files seam fired, neither provider seam did. + list_user_batches.assert_called_once_with( + user_api_key_dict=user, + limit=7, + after="batch-cursor", + provider="openai", + target_model_names="m1,m2", + llm_router=list_harness.router, + ) + list_harness.litellm_alist.assert_not_called() + list_harness.router_alist.assert_not_called() + assert resp is page + + +@pytest.mark.asyncio +async def test_list__managed_files_beats_model_param(list_harness): + """Branch 1 is checked before the model branch: a model in the body does not + divert away from managed-files listing.""" + page = FakeListPage([]) + list_user_batches = list_harness.set_managed_files(page) + + await call_list(list_harness, body={"model": "azure/gpt-4o"}) + + list_user_batches.assert_called_once() + list_harness.litellm_alist.assert_not_called() + list_harness.router_alist.assert_not_called() + list_harness.creds_resolver.assert_not_called() + + +# --------------------------------------------------------------------------- # +# Branch 2 - model from body/query/header. CURRENTLY BROKEN: the endpoint +# forwards custom_llm_provider both explicitly and via **data (it calls +# data.update(credentials) but never pops custom_llm_provider the way +# create/retrieve do through prepare_data_with_credentials), so every call +# raises "multiple values for keyword argument 'custom_llm_provider'". +# +# The strict xfail below encodes the INTENDED contract (litellm seam fires, +# creds resolved for the body model, response ids encoded). It xfails today on +# the duplicate-kwarg TypeError; the day that branch is fixed it will XPASS and +# strict-mode turns the green into a failure, forcing whoever fixes it to drop +# the marker and adopt this as a live regression test. +# --------------------------------------------------------------------------- # + + +@pytest.mark.xfail( + strict=True, + raises=ProxyException, + reason="list_batches model branch passes custom_llm_provider twice " + "(explicit kwarg + **data after data.update(credentials)); remove when fixed", +) +@pytest.mark.asyncio +async def test_list__model_from_body_routes_and_encodes(list_harness): + list_harness.litellm_alist.return_value = FakeListPage( + [make_batch(id="batch-1"), make_batch(id="batch-2")] + ) + + resp = await call_list(list_harness, body={"model": "azure/gpt-4o"}) + + assert list_harness.litellm_alist.call_count == 1 + list_harness.router_alist.assert_not_called() + list_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + assert resp.data[0].id == encode_file_id_with_model( + "batch-1", "azure/gpt-4o", id_type="batch" + ) + assert resp.data[1].id == encode_file_id_with_model( + "batch-2", "azure/gpt-4o", id_type="batch" + ) + + +# --------------------------------------------------------------------------- # +# Branch 3 - target_model_names (function param or body) -> llm_router. Routes +# to the FIRST model in the comma list; `model` is stripped from data first. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__target_model_names_param_routes_to_router(list_harness): + await call_list(list_harness, target_model_names="m1,m2", limit=3, after="cur") + + assert list_harness.router_alist.call_count == 1 + list_harness.litellm_alist.assert_not_called() + list_harness.creds_resolver.assert_not_called() + # first model only; after/limit forwarded; nothing else (param not in data). + assert list_harness.router_kwargs() == { + "model": "m1", + "after": "cur", + "limit": 3, + } + + +@pytest.mark.asyncio +async def test_list__target_model_names_from_body(list_harness): + await call_list(list_harness, body={"target_model_names": "m1,m2"}) + + assert list_harness.router_alist.call_count == 1 + list_harness.litellm_alist.assert_not_called() + kwargs = list_harness.router_kwargs() + assert kwargs["model"] == "m1" + # body-sourced target_model_names stays in the forwarded data. + assert kwargs["target_model_names"] == "m1,m2" + + +@pytest.mark.asyncio +async def test_list__target_model_names_takes_first_only(list_harness): + """Locks the current behavior: with multiple target models, only the first + is routed to (silently, unlike create which 400s on >1). A change here - + intentional or not - must update this test.""" + await call_list(list_harness, target_model_names="alpha,beta,gamma") + + assert list_harness.router_kwargs()["model"] == "alpha" + + +# --------------------------------------------------------------------------- # +# Branch 4 - fallback to custom_llm_provider (env-var creds). MUST NOT touch +# the credential resolver or the router. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__fallback_default_openai(list_harness): + await call_list(list_harness) + + assert list_harness.litellm_alist.call_count == 1 + list_harness.router_alist.assert_not_called() + list_harness.creds_resolver.assert_not_called() # inverse-bug guard + assert list_harness.alist_kwargs() == { + "custom_llm_provider": "openai", + "after": None, + "limit": None, + } + + +@pytest.mark.asyncio +async def test_list__fallback_provider_path_param(list_harness): + await call_list(list_harness, provider="anthropic") + + list_harness.creds_resolver.assert_not_called() + assert list_harness.alist_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_list__fallback_provider_from_header(list_harness): + list_harness.provider_from_headers.return_value = "bedrock" + + await call_list(list_harness) + + assert list_harness.alist_kwargs()["custom_llm_provider"] == "bedrock" + + +@pytest.mark.asyncio +async def test_list__fallback_provider_from_query(list_harness): + list_harness.provider_from_query.return_value = "vertex_ai" + + await call_list(list_harness) + + assert list_harness.alist_kwargs()["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_list__fallback_after_and_limit_forwarded(list_harness): + await call_list(list_harness, after="cursor-9", limit=42) + + kwargs = list_harness.alist_kwargs() + assert kwargs["after"] == "cursor-9" + assert kwargs["limit"] == 42 + + +# --------------------------------------------------------------------------- # +# Cross-cutting: router requirement, route_type, failure hook. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_list__no_router_raises_500(list_harness): + with patch.object(proxy_server, "llm_router", None): + with pytest.raises(ProxyException) as exc: + await call_list(list_harness) + + assert exc.value.code == "500" + list_harness.litellm_alist.assert_not_called() + + +@pytest.mark.asyncio +async def test_list__uses_alist_batches_route_type(list_harness): + await call_list(list_harness) + + assert list_harness.pre_call.call_args.kwargs["route_type"] == "alist_batches" + + +@pytest.mark.asyncio +async def test_list__exception_calls_failure_hook(list_harness): + list_harness.litellm_alist.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_list(list_harness) + + list_harness.logging.post_call_failure_hook.assert_called_once() + assert ( + list_harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# # +# POST /v1/batches/{batch_id}/cancel - cancel_batch routing-contract tests # +# # +# Three branches, first match wins: # +# 1. model-encoded batch id -> litellm.acancel_batch via model creds # +# 2. unified batch id -> llm_router.acancel_batch (model+batch_id # +# parsed out of the unified id) # +# 3. fallback -> litellm.acancel_batch via env-var provider # +# Every branch then writes state back via update_batch_in_database( # +# operation="cancel"). There is NO ManagedObjectTable read short-circuit # +# here (unlike retrieve). # +# # +# These tests pin CURRENT behavior so a refactor can't silently change it. # +# Two current-behavior quirks are locked deliberately and noted inline: # +# - SCENARIO 1 forwards the DEPLOYMENT model from creds, not the decoded # +# model (retrieve overrides it; cancel does not). # +# - SCENARIO 3 rebuilds a CancelBatchRequest and forwards only # +# {custom_llm_provider, batch_id}, dropping enrichment keys. # +# =========================================================================== # + + +@dataclass +class CancelHarness: + data: Dict[str, Any] + pre_call: AsyncMock + add_data: AsyncMock + get_headers: MagicMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + litellm_acancel: AsyncMock + router: MagicMock + logging: MagicMock + creds_resolver: MagicMock + update_batch_in_db: AsyncMock + + @property + def router_acancel(self) -> AsyncMock: + return self.router.acancel_batch + + def acancel_kwargs(self) -> Dict[str, Any]: + assert self.litellm_acancel.call_count == 1 + return dict(self.litellm_acancel.call_args.kwargs) + + def router_kwargs(self) -> Dict[str, Any]: + assert self.router_acancel.call_count == 1 + return dict(self.router_acancel.call_args.kwargs) + + +@pytest.fixture +def cancel_harness(): + data_holder: Dict[str, Any] = {"data": {}} + logging = MagicMock(spec=ProxyLogging) + logging.post_call_success_hook = AsyncMock(side_effect=lambda **kw: kw["response"]) + logging.post_call_failure_hook = AsyncMock() + logging.update_request_status = AsyncMock() + logging.get_proxy_hook = MagicMock(return_value=None) + + router = MagicMock(spec=Router) + router.acancel_batch = AsyncMock(return_value=make_batch()) + router.get_deployment_credentials_with_provider = MagicMock( + side_effect=_creds_lookup + ) + + pre_call = AsyncMock(side_effect=lambda **kw: (data_holder["data"], MagicMock())) + # add_litellm_data_to_request is a passthrough that returns the data it got. + add_data = AsyncMock(side_effect=lambda **kw: kw["data"]) + get_headers = MagicMock(return_value={}) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + litellm_acancel = AsyncMock(return_value=make_batch()) + update_batch_in_db = AsyncMock(return_value=None) + + with ExitStack() as stack: + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "common_processing_pre_call_logic", + pre_call, + ) + ) + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, "get_custom_headers", get_headers + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context( + patch.object(endpoints, "update_batch_in_database", update_batch_in_db) + ) + stack.enter_context(patch.object(litellm, "acancel_batch", litellm_acancel)) + stack.enter_context( + patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + stack.enter_context(patch.object(proxy_server, "prisma_client", MagicMock())) + stack.enter_context( + patch.object(proxy_server, "add_litellm_data_to_request", add_data) + ) + + yield CancelHarness( + data=data_holder, + pre_call=pre_call, + add_data=add_data, + get_headers=get_headers, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + litellm_acancel=litellm_acancel, + router=router, + logging=logging, + creds_resolver=router.get_deployment_credentials_with_provider, + update_batch_in_db=update_batch_in_db, + ) + + +async def call_cancel( + harness: CancelHarness, + batch_id: str, + *, + provider: Optional[str] = None, + user: Optional[UserAPIKeyAuth] = None, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + data_extra: Optional[Dict[str, Any]] = None, +): + harness.data["data"] = {"batch_id": batch_id, **(data_extra or {})} + return await endpoints.cancel_batch( + request=FakeRequest(headers=headers, query=query), + batch_id=batch_id, + fastapi_response=Response(), + provider=provider, + user_api_key_dict=user or UserAPIKeyAuth(api_key="sk-test"), + ) + + +# --------------------------------------------------------------------------- # +# SCENARIO 1 - model-encoded batch id -> litellm.acancel_batch via model creds. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__model_encoded_id(cancel_harness): + resp = await call_cancel(cancel_harness, AZURE_BATCH_ID) + + # DISPATCH - model-credential path via litellm; router untouched. + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + + # CREDENTIALS - resolved for the model decoded from the batch id. + cancel_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + # SEAM PAYLOAD - exact dict. NOTE current behavior: `model` is the + # DEPLOYMENT name from creds, NOT the decoded model (cancel, unlike + # retrieve, does not override it). Locking this guards the difference. + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "azure", + "batch_id": "batch_orig123", # decoded/stripped original id + "api_key": "sk-azure", + "api_base": "https://azure.test", + "model": "azure/gpt-4o-deployment", + } + + # OUTPUT SHAPE - response id re-encoded with the DECODED model. + assert resp.id == encode_file_id_with_model( + "batch-provider-id", "azure/gpt-4o", id_type="batch" + ) + + # write-back tagged as a cancel. + assert cancel_harness.update_batch_in_db.call_count == 1 + assert cancel_harness.update_batch_in_db.call_args.kwargs["operation"] == "cancel" + + +@pytest.mark.asyncio +async def test_cancel__model_encoded_id_forwards_deployment_model(cancel_harness): + """Pin the current contract: cancel forwards the creds' deployment model. + If someone adds a decoded-model override (as retrieve has), this flips and + must be reviewed.""" + await call_cancel(cancel_harness, AZURE_BATCH_ID) + + assert cancel_harness.acancel_kwargs()["model"] == "azure/gpt-4o-deployment" + + +@pytest.mark.asyncio +async def test_cancel__model_encoded_beats_unified(cancel_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + await call_cancel(cancel_harness, AZURE_BATCH_ID) + + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + cancel_harness.creds_resolver.assert_called_once_with(model_id="azure/gpt-4o") + + +# --------------------------------------------------------------------------- # +# SCENARIO 2 - unified batch id -> llm_router.acancel_batch. model and batch_id +# are parsed out of the unified id; hidden params stamped. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__unified_batch_id_routes_to_router(cancel_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + resp = await call_cancel(cancel_harness, "batch-unified-blob") + + # DISPATCH - router fired, litellm did not, no creds lookup. + assert cancel_harness.router_acancel.call_count == 1 + cancel_harness.litellm_acancel.assert_not_called() + cancel_harness.creds_resolver.assert_not_called() + + # model + batch_id are extracted from the unified id and forwarded. + assert cancel_harness.router_kwargs() == { + "batch_id": "batch-raw-xyz", + "model": "gpt-4o-mini", + } + + # hidden params: unified id passed through, model_id stamped from data. + assert resp._hidden_params["unified_batch_id"] == UNIFIED_BATCH_ID + assert resp._hidden_params["model_id"] == "gpt-4o-mini" + + assert cancel_harness.update_batch_in_db.call_args.kwargs["operation"] == "cancel" + + +@pytest.mark.asyncio +async def test_cancel__unified_missing_model_id_400(cancel_harness): + # unified id with no model_id segment -> get_model_id returns None -> 400. + with patch.object( + endpoints, + "_is_base64_encoded_unified_file_id", + return_value="litellm_proxy;llm_batch_id:batch-xyz", + ): + with pytest.raises(ProxyException) as exc: + await call_cancel(cancel_harness, "batch-unified-blob") + + assert exc.value.code == "400" + cancel_harness.router_acancel.assert_not_called() + cancel_harness.litellm_acancel.assert_not_called() + + +@pytest.mark.asyncio +async def test_cancel__unified_no_router_500(cancel_harness): + with patch.object(proxy_server, "llm_router", None), patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ): + with pytest.raises(ProxyException) as exc: + await call_cancel(cancel_harness, "batch-unified-blob") + + assert exc.value.code == "500" + + +# --------------------------------------------------------------------------- # +# SCENARIO 3 - fallback to custom_llm_provider. Rebuilds a CancelBatchRequest +# and forwards only {custom_llm_provider, batch_id}. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__fallback_default_openai(cancel_harness): + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.litellm_acancel.call_count == 1 + cancel_harness.router_acancel.assert_not_called() + cancel_harness.creds_resolver.assert_not_called() # inverse-bug guard + # current behavior: enrichment keys dropped; only these two forwarded. + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "openai", + "batch_id": "batch-raw-xyz", + } + assert cancel_harness.update_batch_in_db.call_count == 1 + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_path_param(cancel_harness): + await call_cancel(cancel_harness, "batch-raw-xyz", provider="anthropic") + + cancel_harness.creds_resolver.assert_not_called() + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "anthropic" + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_from_data_body(cancel_harness): + await call_cancel( + cancel_harness, "batch-raw-xyz", data_extra={"custom_llm_provider": "bedrock"} + ) + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "bedrock" + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_from_header(cancel_harness): + cancel_harness.provider_from_headers.return_value = "vertex_ai" + + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_cancel__fallback_provider_from_query(cancel_harness): + cancel_harness.provider_from_query.return_value = "azure" + + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "azure" + + +@pytest.mark.xfail( + strict=True, + raises=ProxyException, + reason="cancel SCENARIO 3: `provider or data.pop('custom_llm_provider')` " + "short-circuits when provider (path param) is set, so a body " + "custom_llm_provider is left in data and forwarded twice -> duplicate-kwarg " + "TypeError. Intended: path param wins cleanly. Remove marker when fixed.", +) +@pytest.mark.asyncio +async def test_cancel__fallback_provider_precedence_path_over_body(cancel_harness): + """Intended contract: provider path param beats a body custom_llm_provider. + CURRENTLY raises because the `or` short-circuit skips the data.pop, leaving + the body value to collide with the explicit kwarg.""" + await call_cancel( + cancel_harness, + "batch-raw-xyz", + provider="anthropic", + data_extra={"custom_llm_provider": "bedrock"}, + ) + + assert cancel_harness.acancel_kwargs()["custom_llm_provider"] == "anthropic" + + +# --------------------------------------------------------------------------- # +# Cross-cutting: enrichment route_type and failure-hook on provider error. +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_cancel__uses_acancel_batch_route_type(cancel_harness): + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.pre_call.call_args.kwargs["route_type"] == "acancel_batch" + + +@pytest.mark.asyncio +async def test_cancel__exception_calls_failure_hook(cancel_harness): + cancel_harness.litellm_acancel.side_effect = ValueError("provider boom") + + with pytest.raises(Exception): + await call_cancel(cancel_harness, "batch-raw-xyz") + + cancel_harness.logging.post_call_failure_hook.assert_called_once() + assert ( + cancel_harness.logging.post_call_failure_hook.call_args.kwargs[ + "original_exception" + ].args[0] + == "provider boom" + ) + + +# =========================================================================== # +# Router-required 500 guards (one per endpoint branch that calls the router). +# These pin the defensive checks that fire when llm_router is unset. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_create__loadbalancing_no_router_500(harness): + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": "lb-model", + }, + ) + harness.is_known_model.return_value = True + with patch.object( + litellm, "enable_loadbalancing_on_batch_endpoints", True + ), patch.object(proxy_server, "llm_router", None): + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "500" + harness.router_acreate.assert_not_called() + harness.litellm_acreate.assert_not_called() + + +@pytest.mark.asyncio +async def test_create__unified_no_router_500(harness): + set_body( + harness, + { + "input_file_id": "litellm_proxy_unified_id", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value="unified-xyz" + ), patch.object( + endpoints, "get_models_from_unified_file_id", return_value=["gpt-4o-mini"] + ), patch.object( + proxy_server, "llm_router", None + ): + with pytest.raises(ProxyException) as exc: + await call_create(harness) + + assert exc.value.code == "500" + + +@pytest.mark.asyncio +async def test_retrieve__unified_no_router_500(retrieve_harness): + with patch.object( + endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID + ), patch.object(proxy_server, "llm_router", None): + with pytest.raises(ProxyException) as exc: + await call_retrieve(retrieve_harness, "batch-unified-blob") + + assert exc.value.code == "500" + retrieve_harness.router_aretrieve.assert_not_called() + retrieve_harness.litellm_aretrieve.assert_not_called()