From e95f50a43743f661e11eff3b26cc6cbf15637b6a Mon Sep 17 00:00:00 2001 From: Dawei Gu Date: Wed, 6 May 2026 15:14:26 -0700 Subject: [PATCH] test(bedrock): cover retrieve_batch dispatch for both ARN families MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Codecov flagged 8 uncovered lines on `litellm/batches/main.py` after this PR refactored the Bedrock dispatch into a single guard with two sub-branches (`async-invoke` + `model-invocation-job`). Existing tests exercised the handlers directly but not the dispatch in `main.py`. Adds `tests/test_litellm/batches/test_retrieve_batch_bedrock_dispatch.py` with 6 mocked tests that exercise `litellm.retrieve_batch` end-to-end for the dispatch logic: - async-invoke ARN routes to `_handle_async_invoke_status` - async-invoke ARN with no region falls back to "us-east-1" (preserves prior behavior on this branch) - model-invocation-job ARN routes to the new `_handle_model_invocation_job_status` handler - model-invocation-job ARN with no region forwards None (so the new handler can sniff region from the ARN itself, rather than getting silently routed to us-east-1) - unrelated bedrock ARN family falls through to the generic provider-config retrieve path (neither special handler invoked) - non-bedrock batch ids skip the bedrock dispatch entirely Both handlers are mocked at the import site so the tests don't hit AWS — the focus here is purely the new dispatch logic in main.py. Co-authored-by: Cursor --- tests/test_litellm/batches/__init__.py | 0 .../test_retrieve_batch_bedrock_dispatch.py | 162 ++++++++++++++++++ 2 files changed, 162 insertions(+) create mode 100644 tests/test_litellm/batches/__init__.py create mode 100644 tests/test_litellm/batches/test_retrieve_batch_bedrock_dispatch.py diff --git a/tests/test_litellm/batches/__init__.py b/tests/test_litellm/batches/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/batches/test_retrieve_batch_bedrock_dispatch.py b/tests/test_litellm/batches/test_retrieve_batch_bedrock_dispatch.py new file mode 100644 index 00000000000..1ef789a0615 --- /dev/null +++ b/tests/test_litellm/batches/test_retrieve_batch_bedrock_dispatch.py @@ -0,0 +1,162 @@ +"""Cover the Bedrock-ARN dispatch in ``litellm.batches.main.retrieve_batch``. + +The dispatch picks one of two Bedrock handlers depending on the ARN +family in ``batch_id``: + +* ``:async-invoke/`` -> ``_handle_async_invoke_status`` (data plane) +* ``:model-invocation-job/`` -> ``_handle_model_invocation_job_status`` + (control plane, added in this PR) + +Anything else falls through to the generic ``provider_config`` retrieve +flow. We mock the two handlers so the tests don't hit AWS — the focus +here is purely the dispatch logic that lives in ``main.py``. +""" + +from __future__ import annotations + +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm # noqa: E402 + +ASYNC_INVOKE_ARN = "arn:aws:bedrock:us-west-2:123456789012:async-invoke/abc123def456" +MIJ_ARN = "arn:aws:bedrock:us-west-2:123456789012:model-invocation-job/abc1234567" + + +@pytest.fixture +def mock_handlers(): + """Patch both Bedrock retrieve handlers and yield the mocks. + + We patch at the import site (litellm.batches.main) rather than the + definition site so the ``BedrockBatchesHandler`` reference inside + ``retrieve_batch`` resolves to our mocks. + """ + fake_batch = MagicMock(name="LiteLLMBatch") + with ( + patch( + "litellm.batches.main.BedrockBatchesHandler._handle_async_invoke_status", + return_value=fake_batch, + ) as async_invoke, + patch( + "litellm.batches.main.BedrockBatchesHandler._handle_model_invocation_job_status", + return_value=fake_batch, + ) as mij, + ): + yield async_invoke, mij, fake_batch + + +def test_async_invoke_arn_routes_to_async_invoke_handler(mock_handlers): + """``:async-invoke/`` ARNs go to the data-plane handler.""" + async_invoke, mij, fake_batch = mock_handlers + + result = litellm.retrieve_batch( + batch_id=ASYNC_INVOKE_ARN, + custom_llm_provider="bedrock", + aws_region_name="us-west-2", + ) + + assert result is fake_batch + async_invoke.assert_called_once() + mij.assert_not_called() + call_kwargs = async_invoke.call_args.kwargs + assert call_kwargs["batch_id"] == ASYNC_INVOKE_ARN + assert call_kwargs["aws_region_name"] == "us-west-2" + # Region must be stripped from the forwarded kwargs to avoid TypeError + # (it's already an explicit positional/keyword arg). + assert "aws_region_name" not in { + k + for k in call_kwargs + if k not in {"batch_id", "aws_region_name", "logging_obj"} + } + + +def test_async_invoke_arn_falls_back_to_default_region_when_unset(mock_handlers): + """If no ``aws_region_name`` is passed, the data-plane handler defaults + to ``us-east-1`` (preserving prior behavior on this branch).""" + async_invoke, _mij, _ = mock_handlers + + litellm.retrieve_batch( + batch_id=ASYNC_INVOKE_ARN, + custom_llm_provider="bedrock", + ) + + async_invoke.assert_called_once() + assert async_invoke.call_args.kwargs["aws_region_name"] == "us-east-1" + + +def test_model_invocation_job_arn_routes_to_mij_handler(mock_handlers): + """``:model-invocation-job/`` ARNs go to the new control-plane handler.""" + _async_invoke, mij, fake_batch = mock_handlers + + result = litellm.retrieve_batch( + batch_id=MIJ_ARN, + custom_llm_provider="bedrock", + aws_region_name="us-west-2", + ) + + assert result is fake_batch + mij.assert_called_once() + _async_invoke.assert_not_called() + call_kwargs = mij.call_args.kwargs + assert call_kwargs["batch_id"] == MIJ_ARN + assert call_kwargs["aws_region_name"] == "us-west-2" + + +def test_model_invocation_job_arn_with_no_region_passes_none(mock_handlers): + """The MIJ handler is responsible for sniffing region from the ARN + when none is explicitly provided. Dispatch must forward ``None`` + rather than substituting a default — otherwise per-region jobs in + other AWS regions would silently route to ``us-east-1``.""" + _async_invoke, mij, _ = mock_handlers + + litellm.retrieve_batch( + batch_id=MIJ_ARN, + custom_llm_provider="bedrock", + ) + + mij.assert_called_once() + assert mij.call_args.kwargs["aws_region_name"] is None + + +def test_unrelated_bedrock_arn_falls_through_to_provider_config(mock_handlers): + """Bedrock ARNs that aren't async-invoke or model-invocation-job + must NOT hit either special handler — they should fall through to + the existing generic provider_config path. We don't fully exercise + that path here (it requires a real provider config); we just assert + neither special handler is invoked.""" + async_invoke, mij, _ = mock_handlers + + # Use a plausible-but-unsupported Bedrock ARN family. + unrelated_arn = "arn:aws:bedrock:us-west-2:123456789012:provisioned-model/xyz" + + with pytest.raises(Exception): + # Will raise because no provider_config exists for this path — + # that's fine, we just need to assert neither bedrock handler ran + # before the failure. + litellm.retrieve_batch( + batch_id=unrelated_arn, + custom_llm_provider="bedrock", + ) + + async_invoke.assert_not_called() + mij.assert_not_called() + + +def test_non_bedrock_id_skips_bedrock_dispatch_entirely(mock_handlers): + """Plain (non-ARN) batch ids must not even enter the Bedrock dispatch + block — they belong to other providers' retrieve flows.""" + async_invoke, mij, _ = mock_handlers + + with pytest.raises(Exception): + litellm.retrieve_batch( + batch_id="batch_abc123", + custom_llm_provider="openai", + ) + + async_invoke.assert_not_called() + mij.assert_not_called()