mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
a04321d2e1
commit
2cf565ae28
23 changed files with 7178 additions and 0 deletions
1
.github/workflows/test-unit-misc.yml
vendored
1
.github/workflows/test-unit-misc.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/test_litellm/batches/__init__.py
Normal file
0
tests/test_litellm/batches/__init__.py
Normal file
735
tests/test_litellm/batches/test_batch_utils.py
Normal file
735
tests/test_litellm/batches/test_batch_utils.py
Normal 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
|
||||
744
tests/test_litellm/batches/test_main.py
Normal file
744
tests/test_litellm/batches/test_main.py
Normal 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
|
||||
0
tests/test_litellm/llms/anthropic/__init__.py
Normal file
0
tests/test_litellm/llms/anthropic/__init__.py
Normal file
0
tests/test_litellm/llms/anthropic/batches/__init__.py
Normal file
0
tests/test_litellm/llms/anthropic/batches/__init__.py
Normal file
286
tests/test_litellm/llms/anthropic/batches/test_handler.py
Normal file
286
tests/test_litellm/llms/anthropic/batches/test_handler.py
Normal 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"
|
||||
650
tests/test_litellm/llms/anthropic/batches/test_transformation.py
Normal file
650
tests/test_litellm/llms/anthropic/batches/test_transformation.py
Normal 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"
|
||||
0
tests/test_litellm/llms/azure/batches/__init__.py
Normal file
0
tests/test_litellm/llms/azure/batches/__init__.py
Normal file
491
tests/test_litellm/llms/azure/batches/test_handler.py
Normal file
491
tests/test_litellm/llms/azure/batches/test_handler.py
Normal 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)
|
||||
0
tests/test_litellm/llms/base_llm/__init__.py
Normal file
0
tests/test_litellm/llms/base_llm/__init__.py
Normal file
0
tests/test_litellm/llms/base_llm/batches/__init__.py
Normal file
0
tests/test_litellm/llms/base_llm/batches/__init__.py
Normal 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
|
||||
231
tests/test_litellm/llms/base_llm/batches/test_transformation.py
Normal file
231
tests/test_litellm/llms/base_llm/batches/test_transformation.py
Normal 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)
|
||||
0
tests/test_litellm/llms/bedrock/__init__.py
Normal file
0
tests/test_litellm/llms/bedrock/__init__.py
Normal file
0
tests/test_litellm/llms/bedrock/batches/__init__.py
Normal file
0
tests/test_litellm/llms/bedrock/batches/__init__.py
Normal file
684
tests/test_litellm/llms/bedrock/batches/test_transformation.py
Normal file
684
tests/test_litellm/llms/bedrock/batches/test_transformation.py
Normal 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"
|
||||
0
tests/test_litellm/llms/vertex_ai/batches/__init__.py
Normal file
0
tests/test_litellm/llms/vertex_ai/batches/__init__.py
Normal file
805
tests/test_litellm/llms/vertex_ai/batches/test_handler.py
Normal file
805
tests/test_litellm/llms/vertex_ai/batches/test_handler.py
Normal 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)
|
||||
396
tests/test_litellm/llms/vertex_ai/batches/test_transformation.py
Normal file
396
tests/test_litellm/llms/vertex_ai/batches/test_transformation.py
Normal 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
|
||||
0
tests/test_litellm/proxy/batches_endpoints/__init__.py
Normal file
0
tests/test_litellm/proxy/batches_endpoints/__init__.py
Normal file
2026
tests/test_litellm/proxy/batches_endpoints/test_endpoints.py
Normal file
2026
tests/test_litellm/proxy/batches_endpoints/test_endpoints.py
Normal file
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue