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 <cursoragent@cursor.com>

* 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 <cursoragent@cursor.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
Sameer Kankute 2026-06-29 09:22:58 +05:30 • committed by GitHub
parent a04321d2e1
commit 2cf565ae28
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
23 changed files with 7178 additions and 0 deletions

View file

@ -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

View file

@ -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

View file

View file

@ -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,<FID>;" - 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

View file

@ -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.<method>.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 <op> 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

View file

@ -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"

View file

@ -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"

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -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="<<<not json>>>")
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"

View file

@ -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 <token>``
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)

View file

@ -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

File diff suppressed because it is too large Load diff